Skip to content
8 changes: 7 additions & 1 deletion packages/api/internal/orchestrator/orchestrator.go
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ type Orchestrator struct {
snapshotCache SnapshotCacheInvalidator

snapshotUpsertSem *utils.AdjustableSemaphore
redisStorage *redisbackend.Storage

// connectGroup deduplicates concurrent dial+register attempts for the same
// physical node. It is keyed by NomadNodeShortID (Nomad-managed nodes) or
Expand Down Expand Up @@ -139,6 +140,9 @@ func New(

bestOfKAlgorithm := placement.NewBestOfK(getBestOfKConfig(ctx, featureFlags)).(*placement.BestOfK)

redisStorage := redisbackend.NewStorage(redisClient)
go redisStorage.Start(ctx)

o := Orchestrator{
httpClient: httpClient,
analytics: analyticsInstance,
Expand All @@ -153,6 +157,7 @@ func New(
snapshotCache: snapshotCache,
tel: tel,
clusters: clusters,
redisStorage: redisStorage,

sandboxCounter: sandboxCounter,
createdCounter: createdCounter,
Expand All @@ -162,7 +167,6 @@ func New(

var reservationStorage sandbox.ReservationStorage
var sandboxStorage sandbox.Storage
redisStorage := redisbackend.NewStorage(redisClient)

switch config.SandboxStorageBackend {
case cfg.SandboxStorageBackendMemory:
Expand Down Expand Up @@ -284,6 +288,8 @@ func (o *Orchestrator) Close(ctx context.Context) error {
errs = append(errs, err)
}

o.redisStorage.Close()

return errors.Join(errs...)
}

Expand Down
19 changes: 17 additions & 2 deletions packages/api/internal/sandbox/storage/redis/main.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package redis

import (
"context"
"time"

"github.com/bsm/redislock"
Expand All @@ -13,7 +14,8 @@ const (
lockTimeout = time.Minute
transitionKeyTTL = 70 * time.Second // Should be longer than the longest expected state transition time
transitionResultKeyTTL = 30 * time.Second
retryInterval = 20 * time.Millisecond
lockRetryInterval = 20 * time.Millisecond
pollInterval = 1 * time.Second // fallback polling interval; PubSub is the primary notification mechanism
)

var _ sandbox.Storage = (*Storage)(nil)
Expand All @@ -22,6 +24,7 @@ type Storage struct {
redisClient redis.UniversalClient
lockService *redislock.Client
lockOption *redislock.Options
subManager *subscriptionManager
}

func (s *Storage) Name() string { return sandbox.StorageNameRedis }
Expand All @@ -33,11 +36,23 @@ func NewStorage(
redisClient: redisClient,
lockService: redislock.New(redisClient),
lockOption: &redislock.Options{
RetryStrategy: newConstantBackoff(retryInterval),
RetryStrategy: newConstantBackoff(lockRetryInterval),
},
subManager: newSubscriptionManager(redisClient),
}
}

// Start subscribes to the global PubSub channel and blocks until the context
// is cancelled or Close is called. It is intended to be called in a goroutine.
func (s *Storage) Start(ctx context.Context) {
s.subManager.start(ctx)
}

// Close shuts down the subscription manager and its background goroutine.
func (s *Storage) Close() {
s.subManager.close()
}

// Sync is here only for legacy reasons, redis backend doesn't need any sync
func (s *Storage) Sync(_ []sandbox.Sandbox, _ string) []sandbox.Sandbox {
return nil
Expand Down
58 changes: 46 additions & 12 deletions packages/api/internal/sandbox/storage/redis/state_change.go
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,20 @@ func (s *Storage) createCallback(teamID uuid.UUID, sandboxID, transitionKey, res
if delErr != nil {
logger.L().Warn(cbCtx, "Failed to delete transition key", logger.WithSandboxID(sandboxID), zap.Error(delErr))
}

// Notify subscribers that the transition is complete so waitForTransition
// goroutines wake up immediately rather than waiting for the next poll tick.
// The routing key is published as the payload so the single global channel
// can serve all sandboxes across all teams.
routingKey := getTransitionRoutingKey(teamID.String(), sandboxID, transitionID)
pubErr := s.redisClient.Publish(cbCtx, globalTransitionNotifyChannel, routingKey).Err()
if pubErr != nil {
logger.L().Warn(cbCtx, "Failed to publish transition notification",
logger.WithSandboxID(sandboxID),
zap.String("transitionID", transitionID),
zap.Error(pubErr),
)
}
}
}

Expand Down Expand Up @@ -226,30 +240,50 @@ func (s *Storage) WaitForStateChange(ctx context.Context, teamID uuid.UUID, sand
}

// waitForTransition waits for a specific transition to complete.
// It should receive a signal via the fan-out PubSub channel or fallback to a 1-second ticker
func (s *Storage) waitForTransition(
ctx context.Context,
teamID uuid.UUID,
sandboxID,
transitionID string,
) error {
routingKey := getTransitionRoutingKey(teamID.String(), sandboxID, transitionID)
transitionKey := getTransitionKey(teamID.String(), sandboxID)
resultKey := getTransitionResultKey(teamID.String(), sandboxID, transitionID)

for {
currentTransitionID, err := s.redisClient.Get(ctx, transitionKey).Result()
if errors.Is(err, redis.Nil) || transitionID != currentTransitionID {
// Transition key gone or new transition started - check the result
return s.checkTransitionResult(ctx, getTransitionResultKey(teamID.String(), sandboxID, transitionID))
}
if err != nil {
return fmt.Errorf("failed to check transition key: %w", err)
}
// Subscribe to this specific transition's routing key so notifications
// from other transitions for the same sandbox cannot wake us.
ch, cleanup := s.subManager.subscribe(routingKey)
defer cleanup()

// Wait before the next poll
// Initial check: the transition may have completed before we subscribed.
currentID, err := s.redisClient.Get(ctx, transitionKey).Result()
if errors.Is(err, redis.Nil) || currentID != transitionID {
return s.checkTransitionResult(ctx, resultKey)
}
if err != nil {
return fmt.Errorf("failed to check transition key: %w", err)
}

// 1-second fallback ticker in case a PubSub message is missed.
ticker := time.NewTicker(pollInterval)
defer ticker.Stop()

for {
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(retryInterval):
// Continue polling
case <-ch:
return s.checkTransitionResult(ctx, resultKey)
Comment thread
jakubno marked this conversation as resolved.
case <-ticker.C:
// Fallback poll: check whether the transition key is still present.
currentID, err := s.redisClient.Get(ctx, transitionKey).Result()
if errors.Is(err, redis.Nil) || currentID != transitionID {
return s.checkTransitionResult(ctx, resultKey)
}
if err != nil {
return fmt.Errorf("failed to check transition key: %w", err)
}
}
}
}
Expand Down
Loading
Loading