From 2eb10c5e0d27a811f7c7d3f60b435081b037d91b Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 12:41:46 +0100 Subject: [PATCH 01/11] feat: add manager to prevent exhausting connections --- .../internal/sandbox/storage/redis/main.go | 4 +- .../sandbox/storage/redis/state_change.go | 63 ++++++++-- .../storage/redis/subscription_manager.go | 113 ++++++++++++++++++ .../internal/sandbox/storage/redis/utils.go | 17 ++- 4 files changed, 183 insertions(+), 14 deletions(-) create mode 100644 packages/api/internal/sandbox/storage/redis/subscription_manager.go diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index 7d7d399d33..637ccaca8d 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -13,7 +13,7 @@ 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 + retryInterval = 1 * time.Second // fallback polling interval; PubSub is the primary notification mechanism ) var _ sandbox.Storage = (*Storage)(nil) @@ -22,6 +22,7 @@ type Storage struct { redisClient redis.UniversalClient lockService *redislock.Client lockOption *redislock.Options + subManager *subscriptionManager } func (s *Storage) Name() string { return sandbox.StorageNameRedis } @@ -35,6 +36,7 @@ func NewStorage( lockOption: &redislock.Options{ RetryStrategy: newConstantBackoff(retryInterval), }, + subManager: newSubscriptionManager(redisClient), } } diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index 05477eeb57..d8c4c7a253 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -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) + pubErr := s.redisClient.Publish(cbCtx, getGlobalTransitionNotifyChannel(), routingKey).Err() + if pubErr != nil { + logger.L().Warn(cbCtx, "Failed to publish transition notification", + logger.WithSandboxID(sandboxID), + zap.String("transitionID", transitionID), + zap.Error(pubErr), + ) + } } } @@ -226,30 +240,55 @@ func (s *Storage) WaitForStateChange(ctx context.Context, teamID uuid.UUID, sand } // waitForTransition waits for a specific transition to complete. +// It registers with the central subscriptionManager to receive an immediate +// signal via the fan-out PubSub channel when the transition completes. +// A 1-second ticker acts as a safety-net fallback in case a PubSub message +// is missed (e.g. during a Redis failover or transient connection drop). func (s *Storage) waitForTransition( ctx context.Context, teamID uuid.UUID, sandboxID, transitionID string, ) error { + routingKey := getTransitionRoutingKey(teamID.String(), sandboxID) 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) - } + // Register with the central subscription manager before checking the + // transition key so we cannot miss a publish that fires between the + // check and the select below. + 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(retryInterval) + defer ticker.Stop() + + for { select { case <-ctx.Done(): return ctx.Err() - case <-time.After(retryInterval): - // Continue polling + case <-ch: + // Signalled by the central subscription manager — check the result. + return s.checkTransitionResult(ctx, resultKey) + 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) + } } } } diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go new file mode 100644 index 0000000000..be9a5e737f --- /dev/null +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -0,0 +1,113 @@ +package redis + +import ( + "context" + "sync" + + "github.com/redis/go-redis/v9" +) + +// subscriptionManager maintains a single Redis PubSub connection subscribed to +// globalTransitionNotifyChannel and fans out transition-complete signals to +// registered in-process waiters. +// +// One connection per API pod is sufficient regardless of how many sandboxes are +// concurrently transitioning. Each published message carries the per-sandbox +// routing key as its payload; the manager uses that to wake only the goroutines +// waiting on that specific sandbox. +type subscriptionManager struct { + mu sync.Mutex + waiters map[string][]chan struct{} // routingKey → registered waiters + + ps *redis.PubSub + + ctx context.Context + cancel context.CancelFunc +} + +func newSubscriptionManager(redisClient redis.UniversalClient) *subscriptionManager { + ctx, cancel := context.WithCancel(context.Background()) + + m := &subscriptionManager{ + waiters: make(map[string][]chan struct{}), + ps: redisClient.Subscribe(ctx, getGlobalTransitionNotifyChannel()), + ctx: ctx, + cancel: cancel, + } + + go m.run() + + return m +} + +// subscribe registers a waiter for the given routingKey (per-sandbox). +// The returned channel receives a signal when a transition-complete message +// arrives for that sandbox. The caller MUST invoke the returned cleanup function +// when done to avoid a memory leak. +func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, func()) { + ch := make(chan struct{}, 1) // buffered so dispatch never blocks + + m.mu.Lock() + m.waiters[routingKey] = append(m.waiters[routingKey], ch) + m.mu.Unlock() + + cleanup := func() { + m.mu.Lock() + defer m.mu.Unlock() + + waiters := m.waiters[routingKey] + for i, w := range waiters { + if w == ch { + m.waiters[routingKey] = append(waiters[:i], waiters[i+1:]...) + break + } + } + + if len(m.waiters[routingKey]) == 0 { + delete(m.waiters, routingKey) + } + } + + return ch, cleanup +} + +// run reads from the single global PubSub channel and dispatches signals to +// all waiters whose routing key matches the message payload. +// Runs for the lifetime of the subscriptionManager. +func (m *subscriptionManager) run() { + ch := m.ps.Channel() + for { + select { + case <-m.ctx.Done(): + return + case msg, ok := <-ch: + if !ok { + return + } + m.dispatch(msg.Payload) + } + } +} + +// dispatch signals all waiters registered for the given routing key. +func (m *subscriptionManager) dispatch(routingKey string) { + m.mu.Lock() + waiters := m.waiters[routingKey] + snapshot := make([]chan struct{}, len(waiters)) + copy(snapshot, waiters) + m.mu.Unlock() + + for _, w := range snapshot { + select { + case w <- struct{}{}: + default: + // Waiter already has a pending signal; skip. + } + } +} + +// close shuts down the subscription manager and its Redis PubSub connection. +func (m *subscriptionManager) close() { + m.cancel() + _ = m.ps.Close() +} diff --git a/packages/api/internal/sandbox/storage/redis/utils.go b/packages/api/internal/sandbox/storage/redis/utils.go index 31452da8cb..9207ec7a35 100644 --- a/packages/api/internal/sandbox/storage/redis/utils.go +++ b/packages/api/internal/sandbox/storage/redis/utils.go @@ -8,11 +8,19 @@ import ( const ( sandboxKeyPrefix = "sandbox:storage" - transitionKeyPrefix = "transition:" + transitionKeyPrefix = "transition" + notifySuffix = "notify" sandboxesKey = "sandboxes" indexKey = "index" ) +// getGlobalTransitionNotifyChannel is the single Redis PubSub channel used by +// all sandboxes across all teams. The per-sandbox routing key is embedded in +// the message payload so one connection per API pod is sufficient. +func getGlobalTransitionNotifyChannel() string { + return redis_utils.CreateKey(sandboxKeyPrefix, transitionKeyPrefix, notifySuffix) +} + var ( globalTeamsSet = redis_utils.CreateKey(sandboxKeyPrefix, "global:teams") globalExpirationSet = redis_utils.CreateKey(sandboxKeyPrefix, "global:expiration") @@ -49,3 +57,10 @@ func getTransitionKey(teamID, sandboxID string) string { func getTransitionResultKey(teamID, sandboxID, transitionID string) string { return redis_utils.CreateKey(getTransitionKey(teamID, sandboxID), transitionID) } + +// getTransitionRoutingKey returns the per-sandbox routing key embedded in the +// payload of messages published to globalTransitionNotifyChannel. The +// subscriptionManager uses it to fan out signals to the correct in-process waiters. +func getTransitionRoutingKey(teamID, sandboxID string) string { + return redis_utils.CreateKey(GetTeamPrefix(teamID), transitionKeyPrefix, sandboxID, transitionNotifySuffix) +} From 651aa06de8ec6b53982eaa2de9adb77c781702b7 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 16:33:41 +0100 Subject: [PATCH 02/11] chore: use rwmutex --- .../internal/sandbox/storage/redis/subscription_manager.go | 6 +++--- packages/api/internal/sandbox/storage/redis/utils.go | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index be9a5e737f..80b763108f 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -16,7 +16,7 @@ import ( // routing key as its payload; the manager uses that to wake only the goroutines // waiting on that specific sandbox. type subscriptionManager struct { - mu sync.Mutex + mu sync.RWMutex waiters map[string][]chan struct{} // routingKey → registered waiters ps *redis.PubSub @@ -91,11 +91,11 @@ func (m *subscriptionManager) run() { // dispatch signals all waiters registered for the given routing key. func (m *subscriptionManager) dispatch(routingKey string) { - m.mu.Lock() + m.mu.RLock() waiters := m.waiters[routingKey] snapshot := make([]chan struct{}, len(waiters)) copy(snapshot, waiters) - m.mu.Unlock() + m.mu.RUnlock() for _, w := range snapshot { select { diff --git a/packages/api/internal/sandbox/storage/redis/utils.go b/packages/api/internal/sandbox/storage/redis/utils.go index 9207ec7a35..fef51c24d6 100644 --- a/packages/api/internal/sandbox/storage/redis/utils.go +++ b/packages/api/internal/sandbox/storage/redis/utils.go @@ -62,5 +62,5 @@ func getTransitionResultKey(teamID, sandboxID, transitionID string) string { // payload of messages published to globalTransitionNotifyChannel. The // subscriptionManager uses it to fan out signals to the correct in-process waiters. func getTransitionRoutingKey(teamID, sandboxID string) string { - return redis_utils.CreateKey(GetTeamPrefix(teamID), transitionKeyPrefix, sandboxID, transitionNotifySuffix) + return redis_utils.CreateKey(GetTeamPrefix(teamID), transitionKeyPrefix, sandboxID, notifySuffix) } From 1e279fba545cac55de67c41147e0d76a04363ead Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 16:44:14 +0100 Subject: [PATCH 03/11] chore: simplify to map --- .../storage/redis/subscription_manager.go | 31 +++++++------------ 1 file changed, 12 insertions(+), 19 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index 80b763108f..1519c8129f 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -17,7 +17,7 @@ import ( // waiting on that specific sandbox. type subscriptionManager struct { mu sync.RWMutex - waiters map[string][]chan struct{} // routingKey → registered waiters + waiters map[string]map[chan struct{}]struct{} // routingKey → registered waiters ps *redis.PubSub @@ -29,7 +29,7 @@ func newSubscriptionManager(redisClient redis.UniversalClient) *subscriptionMana ctx, cancel := context.WithCancel(context.Background()) m := &subscriptionManager{ - waiters: make(map[string][]chan struct{}), + waiters: make(map[string]map[chan struct{}]struct{}), ps: redisClient.Subscribe(ctx, getGlobalTransitionNotifyChannel()), ctx: ctx, cancel: cancel, @@ -45,30 +45,26 @@ func newSubscriptionManager(redisClient redis.UniversalClient) *subscriptionMana // arrives for that sandbox. The caller MUST invoke the returned cleanup function // when done to avoid a memory leak. func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, func()) { - ch := make(chan struct{}, 1) // buffered so dispatch never blocks + channel := make(chan struct{}, 1) // buffered so dispatch never blocks m.mu.Lock() - m.waiters[routingKey] = append(m.waiters[routingKey], ch) + if m.waiters[routingKey] == nil { + m.waiters[routingKey] = make(map[chan struct{}]struct{}) + } + m.waiters[routingKey][channel] = struct{}{} m.mu.Unlock() cleanup := func() { m.mu.Lock() defer m.mu.Unlock() - waiters := m.waiters[routingKey] - for i, w := range waiters { - if w == ch { - m.waiters[routingKey] = append(waiters[:i], waiters[i+1:]...) - break - } - } - + delete(m.waiters[routingKey], channel) if len(m.waiters[routingKey]) == 0 { delete(m.waiters, routingKey) } } - return ch, cleanup + return channel, cleanup } // run reads from the single global PubSub channel and dispatches signals to @@ -92,14 +88,11 @@ func (m *subscriptionManager) run() { // dispatch signals all waiters registered for the given routing key. func (m *subscriptionManager) dispatch(routingKey string) { m.mu.RLock() - waiters := m.waiters[routingKey] - snapshot := make([]chan struct{}, len(waiters)) - copy(snapshot, waiters) - m.mu.RUnlock() + defer m.mu.RUnlock() - for _, w := range snapshot { + for waiter := range m.waiters[routingKey] { select { - case w <- struct{}{}: + case waiter <- struct{}{}: default: // Waiter already has a pending signal; skip. } From 840e38545273a1f88781916113f645c26b1be822 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 16:56:52 +0100 Subject: [PATCH 04/11] fix: lock interval --- packages/api/internal/sandbox/storage/redis/main.go | 5 +++-- packages/api/internal/sandbox/storage/redis/state_change.go | 2 +- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index 637ccaca8d..be272ff4d0 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -13,7 +13,8 @@ const ( lockTimeout = time.Minute transitionKeyTTL = 70 * time.Second // Should be longer than the longest expected state transition time transitionResultKeyTTL = 30 * time.Second - retryInterval = 1 * time.Second // fallback polling interval; PubSub is the primary notification mechanism + lockRetryInterval = 20 * time.Millisecond + pollInterval = 1 * time.Second // fallback polling interval; PubSub is the primary notification mechanism ) var _ sandbox.Storage = (*Storage)(nil) @@ -34,7 +35,7 @@ func NewStorage( redisClient: redisClient, lockService: redislock.New(redisClient), lockOption: &redislock.Options{ - RetryStrategy: newConstantBackoff(retryInterval), + RetryStrategy: newConstantBackoff(lockRetryInterval), }, subManager: newSubscriptionManager(redisClient), } diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index d8c4c7a253..6ec482fc90 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -270,7 +270,7 @@ func (s *Storage) waitForTransition( } // 1-second fallback ticker in case a PubSub message is missed. - ticker := time.NewTicker(retryInterval) + ticker := time.NewTicker(pollInterval) defer ticker.Stop() for { From 382314a080d2d8e5617d0e8b57d81c2620822d71 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 17:02:15 +0100 Subject: [PATCH 05/11] chore: close the storage --- packages/api/internal/orchestrator/orchestrator.go | 4 ++++ packages/api/internal/sandbox/storage/redis/main.go | 5 +++++ 2 files changed, 9 insertions(+) diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index b195c4b153..ca66ff6bb9 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -66,6 +66,7 @@ type Orchestrator struct { snapshotCache SnapshotCacheInvalidator snapshotUpsertSem *utils.AdjustableSemaphore + redisStorage *redisbackend.Storage } func New( @@ -147,6 +148,7 @@ func New( var reservationStorage sandbox.ReservationStorage var sandboxStorage sandbox.Storage redisStorage := redisbackend.NewStorage(redisClient) + o.redisStorage = redisStorage switch config.SandboxStorageBackend { case cfg.SandboxStorageBackendMemory: @@ -268,6 +270,8 @@ func (o *Orchestrator) Close(ctx context.Context) error { errs = append(errs, err) } + o.redisStorage.Close() + return errors.Join(errs...) } diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index be272ff4d0..c2964867ee 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -41,6 +41,11 @@ func NewStorage( } } +// 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 From 7b9e5ab78989b63730851b9c832f7d4605186389 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Wed, 11 Mar 2026 17:17:58 +0100 Subject: [PATCH 06/11] chore: lint --- .../api/internal/orchestrator/orchestrator.go | 2 +- .../api/internal/sandbox/storage/redis/main.go | 4 +++- .../sandbox/storage/redis/state_change_test.go | 3 ++- .../sandbox/storage/redis/subscription_manager.go | 15 ++++++--------- 4 files changed, 12 insertions(+), 12 deletions(-) diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index ca66ff6bb9..1d9cb36c69 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -147,7 +147,7 @@ func New( var reservationStorage sandbox.ReservationStorage var sandboxStorage sandbox.Storage - redisStorage := redisbackend.NewStorage(redisClient) + redisStorage := redisbackend.NewStorage(ctx, redisClient) o.redisStorage = redisStorage switch config.SandboxStorageBackend { diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index c2964867ee..110b2b3b16 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -1,6 +1,7 @@ package redis import ( + "context" "time" "github.com/bsm/redislock" @@ -29,6 +30,7 @@ type Storage struct { func (s *Storage) Name() string { return sandbox.StorageNameRedis } func NewStorage( + ctx context.Context, redisClient redis.UniversalClient, ) *Storage { return &Storage{ @@ -37,7 +39,7 @@ func NewStorage( lockOption: &redislock.Options{ RetryStrategy: newConstantBackoff(lockRetryInterval), }, - subManager: newSubscriptionManager(redisClient), + subManager: newSubscriptionManager(ctx, redisClient), } } diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index 0cc9cbeb3e..8cca06338c 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -20,7 +20,8 @@ func setupTestStorage(t *testing.T) (*Storage, redis.UniversalClient) { t.Helper() client := redis_utils.SetupInstance(t) - storage := NewStorage(client) + storage := NewStorage(t.Context(), client) + t.Cleanup(storage.Close) return storage, client } diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index 1519c8129f..ec4cfbd13d 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -19,23 +19,20 @@ type subscriptionManager struct { mu sync.RWMutex waiters map[string]map[chan struct{}]struct{} // routingKey → registered waiters - ps *redis.PubSub - - ctx context.Context + ps *redis.PubSub cancel context.CancelFunc } -func newSubscriptionManager(redisClient redis.UniversalClient) *subscriptionManager { - ctx, cancel := context.WithCancel(context.Background()) +func newSubscriptionManager(ctx context.Context, redisClient redis.UniversalClient) *subscriptionManager { + ctx, cancel := context.WithCancel(ctx) m := &subscriptionManager{ waiters: make(map[string]map[chan struct{}]struct{}), ps: redisClient.Subscribe(ctx, getGlobalTransitionNotifyChannel()), - ctx: ctx, cancel: cancel, } - go m.run() + go m.run(ctx) return m } @@ -70,11 +67,11 @@ func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, fun // run reads from the single global PubSub channel and dispatches signals to // all waiters whose routing key matches the message payload. // Runs for the lifetime of the subscriptionManager. -func (m *subscriptionManager) run() { +func (m *subscriptionManager) run(ctx context.Context) { ch := m.ps.Channel() for { select { - case <-m.ctx.Done(): + case <-ctx.Done(): return case msg, ok := <-ch: if !ok { From d9f2e199d8c5731c8bb72c106682e0255a64275a Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Mon, 23 Mar 2026 18:57:07 +0100 Subject: [PATCH 07/11] chore: add tests --- .../api/internal/orchestrator/orchestrator.go | 5 +- .../sandbox/storage/redis/state_change.go | 9 +- .../storage/redis/state_change_test.go | 188 ++++++++++++ .../storage/redis/subscription_manager.go | 10 +- .../redis/subscription_manager_test.go | 271 ++++++++++++++++++ 5 files changed, 466 insertions(+), 17 deletions(-) create mode 100644 packages/api/internal/sandbox/storage/redis/subscription_manager_test.go diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index 1d9cb36c69..00ef99edca 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -124,6 +124,8 @@ func New( bestOfKAlgorithm := placement.NewBestOfK(getBestOfKConfig(ctx, featureFlags)).(*placement.BestOfK) + redisStorage := redisbackend.NewStorage(ctx, redisClient) + o := Orchestrator{ httpClient: httpClient, analytics: analyticsInstance, @@ -138,6 +140,7 @@ func New( snapshotCache: snapshotCache, tel: tel, clusters: clusters, + redisStorage: redisStorage, sandboxCounter: sandboxCounter, createdCounter: createdCounter, @@ -147,8 +150,6 @@ func New( var reservationStorage sandbox.ReservationStorage var sandboxStorage sandbox.Storage - redisStorage := redisbackend.NewStorage(ctx, redisClient) - o.redisStorage = redisStorage switch config.SandboxStorageBackend { case cfg.SandboxStorageBackendMemory: diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index 6ec482fc90..d5d3d0c210 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -240,10 +240,7 @@ func (s *Storage) WaitForStateChange(ctx context.Context, teamID uuid.UUID, sand } // waitForTransition waits for a specific transition to complete. -// It registers with the central subscriptionManager to receive an immediate -// signal via the fan-out PubSub channel when the transition completes. -// A 1-second ticker acts as a safety-net fallback in case a PubSub message -// is missed (e.g. during a Redis failover or transient connection drop). +// 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, @@ -254,9 +251,7 @@ func (s *Storage) waitForTransition( transitionKey := getTransitionKey(teamID.String(), sandboxID) resultKey := getTransitionResultKey(teamID.String(), sandboxID, transitionID) - // Register with the central subscription manager before checking the - // transition key so we cannot miss a publish that fires between the - // check and the select below. + // Register to pubsub ch, cleanup := s.subManager.subscribe(routingKey) defer cleanup() diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index 8cca06338c..87f29acce6 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -965,6 +965,194 @@ func TestStartRemoving_Eviction(t *testing.T) { }) } +// TestWaitForStateChange_PubSubWakesWaiterFast verifies that the PubSub notification +// path (rather than the fallback 1-second ticker) wakes up the waiter promptly. +func TestWaitForStateChange_PubSubWakesWaiterFast(t *testing.T) { + t.Parallel() + + storage, _ := setupTestStorage(t) + ctx := context.Background() + + sbx := createTestSandbox("pubsub-fast-wake") + err := storage.Add(ctx, sbx) + require.NoError(t, err) + + // Start a transition + _, alreadyDone, callback, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + assert.False(t, alreadyDone) + require.NotNil(t, callback) + + // Start a waiter + var waitErr error + waitDone := make(chan struct{}) + waitStarted := make(chan struct{}) + go func() { + close(waitStarted) + waitErr = storage.WaitForStateChange(ctx, sbx.TeamID, sbx.SandboxID) + close(waitDone) + }() + + // Ensure the waiter is subscribed before completing the transition + <-waitStarted + time.Sleep(50 * time.Millisecond) + + // Complete the transition — this publishes a PubSub notification + start := time.Now() + callback(ctx, nil) + + // The waiter should complete well before the 1-second poll interval + select { + case <-waitDone: + elapsed := time.Since(start) + require.NoError(t, waitErr) + assert.Less(t, elapsed, 500*time.Millisecond, + "waiter should be woken by PubSub much faster than the 1s poll interval") + case <-time.After(2 * time.Second): + require.FailNow(t, "WaitForStateChange did not complete in time") + } +} + +// TestWaitForStateChange_MultipleWaitersPubSub verifies that multiple concurrent +// waiters are all woken promptly via the PubSub notification path. +func TestWaitForStateChange_MultipleWaitersPubSub(t *testing.T) { + t.Parallel() + + storage, _ := setupTestStorage(t) + ctx := context.Background() + + sbx := createTestSandbox("pubsub-multi-waiters") + err := storage.Add(ctx, sbx) + require.NoError(t, err) + + // Start a transition + _, alreadyDone, callback, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + assert.False(t, alreadyDone) + require.NotNil(t, callback) + + // Start multiple waiters + numWaiters := 5 + errs := make([]error, numWaiters) + completionTimes := make([]time.Duration, numWaiters) + var wg sync.WaitGroup + + for i := range numWaiters { + wg.Add(1) + go func(idx int) { + defer wg.Done() + errs[idx] = storage.WaitForStateChange(ctx, sbx.TeamID, sbx.SandboxID) + }(i) + } + + // Let all waiters subscribe + time.Sleep(100 * time.Millisecond) + + // Complete the transition + callbackTime := time.Now() + callback(ctx, nil) + + // Wait for all + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() + + select { + case <-done: + elapsed := time.Since(callbackTime) + for i := range numWaiters { + require.NoError(t, errs[i], "waiter %d should complete without error", i) + } + _ = completionTimes // used for timing assertion via elapsed + assert.Less(t, elapsed, 500*time.Millisecond, + "all waiters should be woken by PubSub much faster than the 1s poll interval") + case <-time.After(3 * time.Second): + require.FailNow(t, "not all waiters completed in time") + } +} + +// TestCallback_PublishesNotification verifies that the transition callback publishes +// a notification to the global PubSub channel with the correct routing key. +func TestCallback_PublishesNotification(t *testing.T) { + t.Parallel() + + storage, client := setupTestStorage(t) + ctx := context.Background() + + sbx := createTestSandbox("callback-publishes") + err := storage.Add(ctx, sbx) + require.NoError(t, err) + + // Subscribe to the global notification channel directly + pubsub := client.Subscribe(ctx, getGlobalTransitionNotifyChannel()) + defer pubsub.Close() + + // Wait for the subscription to be ready + time.Sleep(100 * time.Millisecond) + + // Start a transition + _, _, callback, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + require.NotNil(t, callback) + + // Complete the transition + callback(ctx, nil) + + // Read the published message + msg, err := pubsub.ReceiveMessage(ctx) + require.NoError(t, err) + + expectedRoutingKey := getTransitionRoutingKey(sbx.TeamID.String(), sbx.SandboxID) + assert.Equal(t, expectedRoutingKey, msg.Payload, "published payload should be the per-sandbox routing key") +} + +// TestStartRemoving_PauseThenKill_PubSubFastWake verifies that the PubSub path +// makes the waiting kill complete faster than it would with only polling. +func TestStartRemoving_PauseThenKill_PubSubFastWake(t *testing.T) { + t.Parallel() + + storage, _ := setupTestStorage(t) + ctx := context.Background() + + sbx := createTestSandbox("pubsub-pause-kill") + err := storage.Add(ctx, sbx) + require.NoError(t, err) + + // Start pause + _, _, pauseCallback, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + + // Concurrently start a kill (will wait for pause to finish) + killDone := make(chan struct{}) + var killErr error + var killCallback func(context.Context, error) + go func() { + _, _, killCallback, killErr = storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionKill}) + close(killDone) + }() + + // Let the kill request start waiting + time.Sleep(50 * time.Millisecond) + + // Complete the pause — PubSub should wake the kill waiter immediately + start := time.Now() + pauseCallback(ctx, nil) + + select { + case <-killDone: + elapsed := time.Since(start) + require.NoError(t, killErr) + require.NotNil(t, killCallback) + assert.Less(t, elapsed, 500*time.Millisecond, + "kill should be woken by PubSub notification, not 1s poll") + killCallback(ctx, nil) + case <-time.After(3 * time.Second): + require.FailNow(t, "kill did not complete in time") + } +} + // TestStartRemoving_CompletedTransitionAllowsNewTransition tests that a completed // transition doesn't block a new transition to a different state. func TestStartRemoving_CompletedTransitionAllowsNewTransition(t *testing.T) { diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index ec4cfbd13d..9a5516a603 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -7,14 +7,8 @@ import ( "github.com/redis/go-redis/v9" ) -// subscriptionManager maintains a single Redis PubSub connection subscribed to -// globalTransitionNotifyChannel and fans out transition-complete signals to -// registered in-process waiters. -// -// One connection per API pod is sufficient regardless of how many sandboxes are -// concurrently transitioning. Each published message carries the per-sandbox -// routing key as its payload; the manager uses that to wake only the goroutines -// waiting on that specific sandbox. +// subscriptionManager maintains a Redis PubSub connection and +// fans out transition-complete signals to registered in-process waiters. type subscriptionManager struct { mu sync.RWMutex waiters map[string]map[chan struct{}]struct{} // routingKey → registered waiters diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go new file mode 100644 index 0000000000..a2c30fe28f --- /dev/null +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go @@ -0,0 +1,271 @@ +package redis + +import ( + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" +) + +func setupTestManager(t *testing.T) (*subscriptionManager, *Storage) { + t.Helper() + + client := redis_utils.SetupInstance(t) + storage := NewStorage(t.Context(), client) + t.Cleanup(storage.Close) + + return storage.subManager, storage +} + +func TestSubscriptionManager_SubscribeAndDispatch(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + ch, cleanup := m.subscribe("key1") + defer cleanup() + + m.dispatch("key1") + + select { + case <-ch: + // OK + case <-time.After(time.Second): + require.FailNow(t, "expected signal on channel") + } +} + +func TestSubscriptionManager_DispatchOnlyMatchingKey(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + ch1, cleanup1 := m.subscribe("key1") + defer cleanup1() + ch2, cleanup2 := m.subscribe("key2") + defer cleanup2() + + // Dispatch only to key2 + m.dispatch("key2") + + select { + case <-ch2: + // OK — key2 was signalled + case <-time.After(time.Second): + require.FailNow(t, "expected signal on ch2") + } + + // ch1 should NOT have been signalled + select { + case <-ch1: + require.FailNow(t, "ch1 should not have received a signal") + case <-time.After(50 * time.Millisecond): + // OK — no signal + } +} + +func TestSubscriptionManager_MultipleWaitersForSameKey(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + const numWaiters = 5 + channels := make([]<-chan struct{}, numWaiters) + cleanups := make([]func(), numWaiters) + for i := range numWaiters { + channels[i], cleanups[i] = m.subscribe("shared-key") + defer cleanups[i]() + } + + m.dispatch("shared-key") + + for i, ch := range channels { + select { + case <-ch: + // OK + case <-time.After(time.Second): + require.FailNowf(t, "waiter did not receive signal", "waiter %d", i) + } + } +} + +func TestSubscriptionManager_CleanupRemovesWaiter(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + ch, cleanup := m.subscribe("key-cleanup") + cleanup() + + // After cleanup, dispatch should not send to the removed channel + m.dispatch("key-cleanup") + + select { + case <-ch: + require.FailNow(t, "should not receive signal after cleanup") + case <-time.After(10 * time.Millisecond): + // OK + } + + // Verify the routing key entry was fully removed + m.mu.RLock() + _, exists := m.waiters["key-cleanup"] + m.mu.RUnlock() + assert.False(t, exists, "routing key should be removed from waiters map after last subscriber cleans up") +} + +func TestSubscriptionManager_CleanupPartialRemoval(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + ch1, cleanup1 := m.subscribe("key-partial") + ch2, cleanup2 := m.subscribe("key-partial") + defer cleanup2() + + // Remove only the first subscriber + cleanup1() + + m.dispatch("key-partial") + + // ch1 should NOT receive (it was cleaned up) + select { + case <-ch1: + require.FailNow(t, "ch1 should not receive after cleanup") + case <-time.After(10 * time.Millisecond): + // OK + } + + // ch2 should still receive + select { + case <-ch2: + // OK + case <-time.After(time.Second): + require.FailNow(t, "ch2 should still receive signal") + } + + // The routing key should still exist (ch2 is still subscribed) + m.mu.RLock() + _, exists := m.waiters["key-partial"] + m.mu.RUnlock() + assert.True(t, exists, "routing key should still exist with remaining subscriber") +} + +func TestSubscriptionManager_DoubleDispatchDoesNotBlock(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + ch, cleanup := m.subscribe("key-double") + defer cleanup() + + // Dispatch twice — the channel is buffered(1), so the second dispatch + // should be silently dropped (not block). + m.dispatch("key-double") + m.dispatch("key-double") + + select { + case <-ch: + // OK — first signal consumed + case <-time.After(time.Second): + require.FailNow(t, "expected signal on channel") + } + + // No second signal should be available + select { + case <-ch: + require.FailNow(t, "should not have a second signal") + case <-time.After(50 * time.Millisecond): + // OK + } +} + +func TestSubscriptionManager_ConcurrentSubscribeDispatchCleanup(t *testing.T) { + t.Parallel() + + m, _ := setupTestManager(t) + + const goroutines = 20 + var wg sync.WaitGroup + + for i := range goroutines { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + ch, cleanup := m.subscribe("concurrent-key") + defer cleanup() + + // Dispatch from every goroutine + m.dispatch("concurrent-key") + + // Drain the channel (may or may not have a signal depending on timing) + select { + case <-ch: + case <-time.After(100 * time.Millisecond): + } + }(i) + } + + wg.Wait() + + // After all goroutines clean up, the waiters map should be empty for this key + m.mu.RLock() + _, exists := m.waiters["concurrent-key"] + m.mu.RUnlock() + assert.False(t, exists, "all waiters should be cleaned up") +} + +func TestSubscriptionManager_PubSubEndToEnd(t *testing.T) { + t.Parallel() + + client := redis_utils.SetupInstance(t) + storage := NewStorage(t.Context(), client) + t.Cleanup(storage.Close) + + routingKey := "test:routing:key" + ch, cleanup := storage.subManager.subscribe(routingKey) + defer cleanup() + + // Allow time for the PubSub subscription to be established + time.Sleep(50 * time.Millisecond) + + // Publish via Redis (simulating what the callback does) + err := client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), routingKey).Err() + require.NoError(t, err) + + select { + case <-ch: + // OK — received the PubSub notification + case <-time.After(3 * time.Second): + require.FailNow(t, "did not receive PubSub notification") + } +} + +func TestSubscriptionManager_PubSubIgnoresUnrelatedKeys(t *testing.T) { + t.Parallel() + + client := redis_utils.SetupInstance(t) + storage := NewStorage(t.Context(), client) + t.Cleanup(storage.Close) + + ch, cleanup := storage.subManager.subscribe("my:sandbox:key") + defer cleanup() + + time.Sleep(50 * time.Millisecond) + + // Publish a message with a different routing key + err := client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), "other:sandbox:key").Err() + require.NoError(t, err) + + select { + case <-ch: + require.FailNow(t, "should not have received signal for a different routing key") + case <-time.After(200 * time.Millisecond): + // OK + } +} From 755fb9462ec112fb422a104a88eff8d1b23c6470 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Mon, 23 Mar 2026 20:06:33 +0100 Subject: [PATCH 08/11] chore: use transaction id in the subscription key --- .../sandbox/storage/redis/state_change.go | 8 +- .../storage/redis/state_change_test.go | 79 ++++++++++++++++++- .../internal/sandbox/storage/redis/utils.go | 11 +-- 3 files changed, 87 insertions(+), 11 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index d5d3d0c210..893bf4a1e6 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -197,7 +197,7 @@ func (s *Storage) createCallback(teamID uuid.UUID, sandboxID, transitionKey, res // 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) + routingKey := getTransitionRoutingKey(teamID.String(), sandboxID, transitionID) pubErr := s.redisClient.Publish(cbCtx, getGlobalTransitionNotifyChannel(), routingKey).Err() if pubErr != nil { logger.L().Warn(cbCtx, "Failed to publish transition notification", @@ -247,11 +247,12 @@ func (s *Storage) waitForTransition( sandboxID, transitionID string, ) error { - routingKey := getTransitionRoutingKey(teamID.String(), sandboxID) + routingKey := getTransitionRoutingKey(teamID.String(), sandboxID, transitionID) transitionKey := getTransitionKey(teamID.String(), sandboxID) resultKey := getTransitionResultKey(teamID.String(), sandboxID, transitionID) - // Register to pubsub + // 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() @@ -273,7 +274,6 @@ func (s *Storage) waitForTransition( case <-ctx.Done(): return ctx.Err() case <-ch: - // Signalled by the central subscription manager — check the result. return s.checkTransitionResult(ctx, resultKey) case <-ticker.C: // Fallback poll: check whether the transition key is still present. diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index 87f29acce6..bef94bc709 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -1097,6 +1097,11 @@ func TestCallback_PublishesNotification(t *testing.T) { require.NoError(t, err) require.NotNil(t, callback) + // Read the transitionID so we can build the expected routing key + transitionKey := getTransitionKey(sbx.TeamID.String(), sbx.SandboxID) + transitionID, err := client.Get(ctx, transitionKey).Result() + require.NoError(t, err) + // Complete the transition callback(ctx, nil) @@ -1104,8 +1109,8 @@ func TestCallback_PublishesNotification(t *testing.T) { msg, err := pubsub.ReceiveMessage(ctx) require.NoError(t, err) - expectedRoutingKey := getTransitionRoutingKey(sbx.TeamID.String(), sbx.SandboxID) - assert.Equal(t, expectedRoutingKey, msg.Payload, "published payload should be the per-sandbox routing key") + expectedRoutingKey := getTransitionRoutingKey(sbx.TeamID.String(), sbx.SandboxID, transitionID) + assert.Equal(t, expectedRoutingKey, msg.Payload, "published payload should be the per-transition routing key") } // TestStartRemoving_PauseThenKill_PubSubFastWake verifies that the PubSub path @@ -1153,6 +1158,76 @@ func TestStartRemoving_PauseThenKill_PubSubFastWake(t *testing.T) { } } +// TestWaitForTransition_StalePubSubNotification verifies that a PubSub notification +// from a previous transition does not wake a waiter for the current transition. +func TestWaitForTransition_StalePubSubNotification(t *testing.T) { + t.Parallel() + + storage, client := setupTestStorage(t) + + sbx := createTestSandbox("stale-pubsub") + require.NoError(t, storage.Add(t.Context(), sbx)) + + // --- Transition A: pause --- + _, _, callbackA, err := storage.StartRemoving(t.Context(), sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + require.NotNil(t, callbackA) + + // Get transition A's ID so we can craft its routing key later. + transitionKey := getTransitionKey(sbx.TeamID.String(), sbx.SandboxID) + transitionIDA, err := client.Get(t.Context(), transitionKey).Result() + require.NoError(t, err) + + // Complete transition A. + callbackA(t.Context(), nil) + + // Restore to Running so we can start a new transition. + _, err = storage.Update(t.Context(), sbx.TeamID, sbx.SandboxID, func(s sandbox.Sandbox) (sandbox.Sandbox, error) { + s.State = sandbox.StateRunning + return s, nil + }) + require.NoError(t, err) + + // --- Transition B: pause again --- + _, _, callbackB, err := storage.StartRemoving(t.Context(), sbx.TeamID, sbx.SandboxID, sandbox.RemoveOpts{Action: sandbox.StateActionPause}) + require.NoError(t, err) + require.NotNil(t, callbackB) + + // Start a waiter for transition B. + waiterDone := make(chan error, 1) + go func() { + waiterDone <- storage.WaitForStateChange(t.Context(), sbx.TeamID, sbx.SandboxID) + }() + + // Let the waiter subscribe. + time.Sleep(100 * time.Millisecond) + + // Publish a notification using transition A's routing key (stale). + staleRoutingKey := getTransitionRoutingKey(sbx.TeamID.String(), sbx.SandboxID, transitionIDA) + require.NoError(t, client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), staleRoutingKey).Err()) + + // Give time for the stale notification to be (not) delivered. + time.Sleep(200 * time.Millisecond) + + // The waiter should still be blocking — stale key doesn't match. + select { + case err := <-waiterDone: + require.FailNow(t, "waiter returned prematurely on stale PubSub notification", "err: %v", err) + default: + // OK — still waiting + } + + // Now complete transition B. + callbackB(t.Context(), nil) + + select { + case err := <-waiterDone: + require.NoError(t, err) + case <-time.After(3 * time.Second): + require.FailNow(t, "waiter did not complete after transition B finished") + } +} + // TestStartRemoving_CompletedTransitionAllowsNewTransition tests that a completed // transition doesn't block a new transition to a different state. func TestStartRemoving_CompletedTransitionAllowsNewTransition(t *testing.T) { diff --git a/packages/api/internal/sandbox/storage/redis/utils.go b/packages/api/internal/sandbox/storage/redis/utils.go index fef51c24d6..486e18246b 100644 --- a/packages/api/internal/sandbox/storage/redis/utils.go +++ b/packages/api/internal/sandbox/storage/redis/utils.go @@ -58,9 +58,10 @@ func getTransitionResultKey(teamID, sandboxID, transitionID string) string { return redis_utils.CreateKey(getTransitionKey(teamID, sandboxID), transitionID) } -// getTransitionRoutingKey returns the per-sandbox routing key embedded in the -// payload of messages published to globalTransitionNotifyChannel. The -// subscriptionManager uses it to fan out signals to the correct in-process waiters. -func getTransitionRoutingKey(teamID, sandboxID string) string { - return redis_utils.CreateKey(GetTeamPrefix(teamID), transitionKeyPrefix, sandboxID, notifySuffix) +// getTransitionRoutingKey returns the per-transition routing key embedded in the +// payload of messages published to globalTransitionNotifyChannel. Including +// transitionID ensures that notifications for one transition can never wake +// waiters subscribed to a different transition for the same sandbox. +func getTransitionRoutingKey(teamID, sandboxID, transitionID string) string { + return redis_utils.CreateKey(GetTeamPrefix(teamID), transitionKeyPrefix, sandboxID, transitionID, notifySuffix) } From 13183257ed7967414ae21c6700d52971cda07045 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Mon, 23 Mar 2026 20:09:08 +0100 Subject: [PATCH 09/11] chore: lint --- .../storage/redis/state_change_test.go | 1 + .../redis/subscription_manager_test.go | 33 +++++++++---------- 2 files changed, 16 insertions(+), 18 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index bef94bc709..1640a6e649 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -1184,6 +1184,7 @@ func TestWaitForTransition_StalePubSubNotification(t *testing.T) { // Restore to Running so we can start a new transition. _, err = storage.Update(t.Context(), sbx.TeamID, sbx.SandboxID, func(s sandbox.Sandbox) (sandbox.Sandbox, error) { s.State = sandbox.StateRunning + return s, nil }) require.NoError(t, err) diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go index a2c30fe28f..9fc676d6e4 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go @@ -11,20 +11,20 @@ import ( redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" ) -func setupTestManager(t *testing.T) (*subscriptionManager, *Storage) { +func setupTestManager(t *testing.T) *subscriptionManager { t.Helper() client := redis_utils.SetupInstance(t) storage := NewStorage(t.Context(), client) t.Cleanup(storage.Close) - return storage.subManager, storage + return storage.subManager } func TestSubscriptionManager_SubscribeAndDispatch(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) ch, cleanup := m.subscribe("key1") defer cleanup() @@ -42,7 +42,7 @@ func TestSubscriptionManager_SubscribeAndDispatch(t *testing.T) { func TestSubscriptionManager_DispatchOnlyMatchingKey(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) ch1, cleanup1 := m.subscribe("key1") defer cleanup1() @@ -71,14 +71,14 @@ func TestSubscriptionManager_DispatchOnlyMatchingKey(t *testing.T) { func TestSubscriptionManager_MultipleWaitersForSameKey(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) const numWaiters = 5 channels := make([]<-chan struct{}, numWaiters) - cleanups := make([]func(), numWaiters) for i := range numWaiters { - channels[i], cleanups[i] = m.subscribe("shared-key") - defer cleanups[i]() + ch, cleanup := m.subscribe("shared-key") + t.Cleanup(cleanup) + channels[i] = ch } m.dispatch("shared-key") @@ -96,7 +96,7 @@ func TestSubscriptionManager_MultipleWaitersForSameKey(t *testing.T) { func TestSubscriptionManager_CleanupRemovesWaiter(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) ch, cleanup := m.subscribe("key-cleanup") cleanup() @@ -121,7 +121,7 @@ func TestSubscriptionManager_CleanupRemovesWaiter(t *testing.T) { func TestSubscriptionManager_CleanupPartialRemoval(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) ch1, cleanup1 := m.subscribe("key-partial") ch2, cleanup2 := m.subscribe("key-partial") @@ -158,7 +158,7 @@ func TestSubscriptionManager_CleanupPartialRemoval(t *testing.T) { func TestSubscriptionManager_DoubleDispatchDoesNotBlock(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) ch, cleanup := m.subscribe("key-double") defer cleanup() @@ -187,16 +187,13 @@ func TestSubscriptionManager_DoubleDispatchDoesNotBlock(t *testing.T) { func TestSubscriptionManager_ConcurrentSubscribeDispatchCleanup(t *testing.T) { t.Parallel() - m, _ := setupTestManager(t) + m := setupTestManager(t) const goroutines = 20 var wg sync.WaitGroup - for i := range goroutines { - wg.Add(1) - go func(idx int) { - defer wg.Done() - + for range goroutines { + wg.Go(func() { ch, cleanup := m.subscribe("concurrent-key") defer cleanup() @@ -208,7 +205,7 @@ func TestSubscriptionManager_ConcurrentSubscribeDispatchCleanup(t *testing.T) { case <-ch: case <-time.After(100 * time.Millisecond): } - }(i) + }) } wg.Wait() From 461f42ef4b4f2b6790b50677fa38721fe02a7a56 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Tue, 24 Mar 2026 09:54:54 +0100 Subject: [PATCH 10/11] chore: simplify the sub transition key --- .../sandbox/storage/redis/state_change.go | 2 +- .../sandbox/storage/redis/state_change_test.go | 4 ++-- .../storage/redis/subscription_manager.go | 12 +++++------- .../storage/redis/subscription_manager_test.go | 18 +++++++++--------- .../internal/sandbox/storage/redis/utils.go | 12 +++++------- 5 files changed, 22 insertions(+), 26 deletions(-) diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index 893bf4a1e6..d7e381ab0e 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -198,7 +198,7 @@ func (s *Storage) createCallback(teamID uuid.UUID, sandboxID, transitionKey, res // 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, getGlobalTransitionNotifyChannel(), routingKey).Err() + pubErr := s.redisClient.Publish(cbCtx, globalTransitionNotifyChannel, routingKey).Err() if pubErr != nil { logger.L().Warn(cbCtx, "Failed to publish transition notification", logger.WithSandboxID(sandboxID), diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index 1640a6e649..f201c2cbc4 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -1086,7 +1086,7 @@ func TestCallback_PublishesNotification(t *testing.T) { require.NoError(t, err) // Subscribe to the global notification channel directly - pubsub := client.Subscribe(ctx, getGlobalTransitionNotifyChannel()) + pubsub := client.Subscribe(ctx, globalTransitionNotifyChannel) defer pubsub.Close() // Wait for the subscription to be ready @@ -1205,7 +1205,7 @@ func TestWaitForTransition_StalePubSubNotification(t *testing.T) { // Publish a notification using transition A's routing key (stale). staleRoutingKey := getTransitionRoutingKey(sbx.TeamID.String(), sbx.SandboxID, transitionIDA) - require.NoError(t, client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), staleRoutingKey).Err()) + require.NoError(t, client.Publish(t.Context(), globalTransitionNotifyChannel, staleRoutingKey).Err()) // Give time for the stale notification to be (not) delivered. time.Sleep(200 * time.Millisecond) diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index 9a5516a603..99b145b4ec 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -22,7 +22,7 @@ func newSubscriptionManager(ctx context.Context, redisClient redis.UniversalClie m := &subscriptionManager{ waiters: make(map[string]map[chan struct{}]struct{}), - ps: redisClient.Subscribe(ctx, getGlobalTransitionNotifyChannel()), + ps: redisClient.Subscribe(ctx, globalTransitionNotifyChannel), cancel: cancel, } @@ -31,10 +31,10 @@ func newSubscriptionManager(ctx context.Context, redisClient redis.UniversalClie return m } -// subscribe registers a waiter for the given routingKey (per-sandbox). +// subscribe registers a waiter for the given routingKey (per-transition). // The returned channel receives a signal when a transition-complete message -// arrives for that sandbox. The caller MUST invoke the returned cleanup function -// when done to avoid a memory leak. +// arrives for that transition. The caller MUST invoke the returned cleanup +// function when done to avoid a memory leak. func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, func()) { channel := make(chan struct{}, 1) // buffered so dispatch never blocks @@ -58,9 +58,7 @@ func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, fun return channel, cleanup } -// run reads from the single global PubSub channel and dispatches signals to -// all waiters whose routing key matches the message payload. -// Runs for the lifetime of the subscriptionManager. +// run reads from the PubSub channel and dispatches signals to matching waiters. func (m *subscriptionManager) run(ctx context.Context) { ch := m.ps.Channel() for { diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go index 9fc676d6e4..0937ea2255 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go @@ -27,7 +27,7 @@ func TestSubscriptionManager_SubscribeAndDispatch(t *testing.T) { m := setupTestManager(t) ch, cleanup := m.subscribe("key1") - defer cleanup() + t.Cleanup(cleanup) m.dispatch("key1") @@ -45,9 +45,9 @@ func TestSubscriptionManager_DispatchOnlyMatchingKey(t *testing.T) { m := setupTestManager(t) ch1, cleanup1 := m.subscribe("key1") - defer cleanup1() + t.Cleanup(cleanup1) ch2, cleanup2 := m.subscribe("key2") - defer cleanup2() + t.Cleanup(cleanup2) // Dispatch only to key2 m.dispatch("key2") @@ -125,7 +125,7 @@ func TestSubscriptionManager_CleanupPartialRemoval(t *testing.T) { ch1, cleanup1 := m.subscribe("key-partial") ch2, cleanup2 := m.subscribe("key-partial") - defer cleanup2() + t.Cleanup(cleanup2) // Remove only the first subscriber cleanup1() @@ -161,7 +161,7 @@ func TestSubscriptionManager_DoubleDispatchDoesNotBlock(t *testing.T) { m := setupTestManager(t) ch, cleanup := m.subscribe("key-double") - defer cleanup() + t.Cleanup(cleanup) // Dispatch twice — the channel is buffered(1), so the second dispatch // should be silently dropped (not block). @@ -226,13 +226,13 @@ func TestSubscriptionManager_PubSubEndToEnd(t *testing.T) { routingKey := "test:routing:key" ch, cleanup := storage.subManager.subscribe(routingKey) - defer cleanup() + t.Cleanup(cleanup) // Allow time for the PubSub subscription to be established time.Sleep(50 * time.Millisecond) // Publish via Redis (simulating what the callback does) - err := client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), routingKey).Err() + err := client.Publish(t.Context(), globalTransitionNotifyChannel, routingKey).Err() require.NoError(t, err) select { @@ -251,12 +251,12 @@ func TestSubscriptionManager_PubSubIgnoresUnrelatedKeys(t *testing.T) { t.Cleanup(storage.Close) ch, cleanup := storage.subManager.subscribe("my:sandbox:key") - defer cleanup() + t.Cleanup(cleanup) time.Sleep(50 * time.Millisecond) // Publish a message with a different routing key - err := client.Publish(t.Context(), getGlobalTransitionNotifyChannel(), "other:sandbox:key").Err() + err := client.Publish(t.Context(), globalTransitionNotifyChannel, "other:sandbox:key").Err() require.NoError(t, err) select { diff --git a/packages/api/internal/sandbox/storage/redis/utils.go b/packages/api/internal/sandbox/storage/redis/utils.go index 486e18246b..b3b77e0c95 100644 --- a/packages/api/internal/sandbox/storage/redis/utils.go +++ b/packages/api/internal/sandbox/storage/redis/utils.go @@ -14,14 +14,12 @@ const ( indexKey = "index" ) -// getGlobalTransitionNotifyChannel is the single Redis PubSub channel used by -// all sandboxes across all teams. The per-sandbox routing key is embedded in -// the message payload so one connection per API pod is sufficient. -func getGlobalTransitionNotifyChannel() string { - return redis_utils.CreateKey(sandboxKeyPrefix, transitionKeyPrefix, notifySuffix) -} - var ( + // globalTransitionNotifyChannel is the single Redis PubSub channel used by + // all transitions. The per-transition routing key is embedded in the message + // payload so one connection per API pod is sufficient. + globalTransitionNotifyChannel = redis_utils.CreateKey(sandboxKeyPrefix, transitionKeyPrefix, notifySuffix) + globalTeamsSet = redis_utils.CreateKey(sandboxKeyPrefix, "global:teams") globalExpirationSet = redis_utils.CreateKey(sandboxKeyPrefix, "global:expiration") ) From a80e6b735c669f9876d75f84b38de3c70b589863 Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Tue, 24 Mar 2026 10:04:40 +0100 Subject: [PATCH 11/11] chore: refactor start and close --- .../api/internal/orchestrator/orchestrator.go | 3 +- .../internal/sandbox/storage/redis/main.go | 9 ++- .../storage/redis/state_change_test.go | 3 +- .../storage/redis/subscription_manager.go | 69 +++++++++++-------- .../redis/subscription_manager_test.go | 9 ++- 5 files changed, 58 insertions(+), 35 deletions(-) diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index 00ef99edca..65e9c67857 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -124,7 +124,8 @@ func New( bestOfKAlgorithm := placement.NewBestOfK(getBestOfKConfig(ctx, featureFlags)).(*placement.BestOfK) - redisStorage := redisbackend.NewStorage(ctx, redisClient) + redisStorage := redisbackend.NewStorage(redisClient) + go redisStorage.Start(ctx) o := Orchestrator{ httpClient: httpClient, diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index 110b2b3b16..291881089d 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -30,7 +30,6 @@ type Storage struct { func (s *Storage) Name() string { return sandbox.StorageNameRedis } func NewStorage( - ctx context.Context, redisClient redis.UniversalClient, ) *Storage { return &Storage{ @@ -39,10 +38,16 @@ func NewStorage( lockOption: &redislock.Options{ RetryStrategy: newConstantBackoff(lockRetryInterval), }, - subManager: newSubscriptionManager(ctx, redisClient), + 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() diff --git a/packages/api/internal/sandbox/storage/redis/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index f201c2cbc4..b895a373a3 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -20,7 +20,8 @@ func setupTestStorage(t *testing.T) (*Storage, redis.UniversalClient) { t.Helper() client := redis_utils.SetupInstance(t) - storage := NewStorage(t.Context(), client) + storage := NewStorage(client) + go storage.Start(t.Context()) t.Cleanup(storage.Close) return storage, client diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go index 99b145b4ec..5105a9866b 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -13,22 +13,50 @@ type subscriptionManager struct { mu sync.RWMutex waiters map[string]map[chan struct{}]struct{} // routingKey → registered waiters - ps *redis.PubSub - cancel context.CancelFunc + redisClient redis.UniversalClient + stop chan struct{} + once sync.Once } -func newSubscriptionManager(ctx context.Context, redisClient redis.UniversalClient) *subscriptionManager { +func newSubscriptionManager(redisClient redis.UniversalClient) *subscriptionManager { + return &subscriptionManager{ + waiters: make(map[string]map[chan struct{}]struct{}), + redisClient: redisClient, + stop: make(chan struct{}), + } +} + +// start subscribes to the global PubSub channel and dispatches signals +// to registered waiters. It blocks until the context is cancelled or +// close is called. It is intended to be called in a goroutine. +func (m *subscriptionManager) start(ctx context.Context) { ctx, cancel := context.WithCancel(ctx) + defer cancel() - m := &subscriptionManager{ - waiters: make(map[string]map[chan struct{}]struct{}), - ps: redisClient.Subscribe(ctx, globalTransitionNotifyChannel), - cancel: cancel, - } + // Cancel the context when close is called. + go func() { + select { + case <-m.stop: + cancel() + case <-ctx.Done(): + } + }() - go m.run(ctx) + ps := m.redisClient.Subscribe(ctx, globalTransitionNotifyChannel) + defer ps.Close() - return m + ch := ps.Channel() + for { + select { + case <-ctx.Done(): + return + case msg, ok := <-ch: + if !ok { + return + } + m.dispatch(msg.Payload) + } + } } // subscribe registers a waiter for the given routingKey (per-transition). @@ -58,22 +86,6 @@ func (m *subscriptionManager) subscribe(routingKey string) (<-chan struct{}, fun return channel, cleanup } -// run reads from the PubSub channel and dispatches signals to matching waiters. -func (m *subscriptionManager) run(ctx context.Context) { - ch := m.ps.Channel() - for { - select { - case <-ctx.Done(): - return - case msg, ok := <-ch: - if !ok { - return - } - m.dispatch(msg.Payload) - } - } -} - // dispatch signals all waiters registered for the given routing key. func (m *subscriptionManager) dispatch(routingKey string) { m.mu.RLock() @@ -90,6 +102,7 @@ func (m *subscriptionManager) dispatch(routingKey string) { // close shuts down the subscription manager and its Redis PubSub connection. func (m *subscriptionManager) close() { - m.cancel() - _ = m.ps.Close() + m.once.Do(func() { + close(m.stop) + }) } diff --git a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go index 0937ea2255..50cf10af59 100644 --- a/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager_test.go @@ -15,7 +15,8 @@ func setupTestManager(t *testing.T) *subscriptionManager { t.Helper() client := redis_utils.SetupInstance(t) - storage := NewStorage(t.Context(), client) + storage := NewStorage(client) + go storage.Start(t.Context()) t.Cleanup(storage.Close) return storage.subManager @@ -221,7 +222,8 @@ func TestSubscriptionManager_PubSubEndToEnd(t *testing.T) { t.Parallel() client := redis_utils.SetupInstance(t) - storage := NewStorage(t.Context(), client) + storage := NewStorage(client) + go storage.Start(t.Context()) t.Cleanup(storage.Close) routingKey := "test:routing:key" @@ -247,7 +249,8 @@ func TestSubscriptionManager_PubSubIgnoresUnrelatedKeys(t *testing.T) { t.Parallel() client := redis_utils.SetupInstance(t) - storage := NewStorage(t.Context(), client) + storage := NewStorage(client) + go storage.Start(t.Context()) t.Cleanup(storage.Close) ch, cleanup := storage.subManager.subscribe("my:sandbox:key")