diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index 3bdc28c88b..6e13963920 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -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 @@ -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, @@ -153,6 +157,7 @@ func New( snapshotCache: snapshotCache, tel: tel, clusters: clusters, + redisStorage: redisStorage, sandboxCounter: sandboxCounter, createdCounter: createdCounter, @@ -162,7 +167,6 @@ func New( var reservationStorage sandbox.ReservationStorage var sandboxStorage sandbox.Storage - redisStorage := redisbackend.NewStorage(redisClient) switch config.SandboxStorageBackend { case cfg.SandboxStorageBackendMemory: @@ -284,6 +288,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 7d7d399d33..291881089d 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" @@ -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) @@ -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 } @@ -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 diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index 05477eeb57..d7e381ab0e 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, 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), + ) + } } } @@ -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) + 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/state_change_test.go b/packages/api/internal/sandbox/storage/redis/state_change_test.go index 0cc9cbeb3e..b895a373a3 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change_test.go +++ b/packages/api/internal/sandbox/storage/redis/state_change_test.go @@ -21,6 +21,8 @@ func setupTestStorage(t *testing.T) (*Storage, redis.UniversalClient) { client := redis_utils.SetupInstance(t) storage := NewStorage(client) + go storage.Start(t.Context()) + t.Cleanup(storage.Close) return storage, client } @@ -964,6 +966,270 @@ 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, globalTransitionNotifyChannel) + 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) + + // 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) + + // Read the published message + msg, err := pubsub.ReceiveMessage(ctx) + require.NoError(t, err) + + 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 +// 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") + } +} + +// 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(), globalTransitionNotifyChannel, 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/subscription_manager.go b/packages/api/internal/sandbox/storage/redis/subscription_manager.go new file mode 100644 index 0000000000..5105a9866b --- /dev/null +++ b/packages/api/internal/sandbox/storage/redis/subscription_manager.go @@ -0,0 +1,108 @@ +package redis + +import ( + "context" + "sync" + + "github.com/redis/go-redis/v9" +) + +// 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 + + redisClient redis.UniversalClient + stop chan struct{} + once sync.Once +} + +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() + + // Cancel the context when close is called. + go func() { + select { + case <-m.stop: + cancel() + case <-ctx.Done(): + } + }() + + ps := m.redisClient.Subscribe(ctx, globalTransitionNotifyChannel) + defer ps.Close() + + 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). +// The returned channel receives a signal when a transition-complete message +// 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 + + m.mu.Lock() + 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() + + delete(m.waiters[routingKey], channel) + if len(m.waiters[routingKey]) == 0 { + delete(m.waiters, routingKey) + } + } + + return channel, cleanup +} + +// dispatch signals all waiters registered for the given routing key. +func (m *subscriptionManager) dispatch(routingKey string) { + m.mu.RLock() + defer m.mu.RUnlock() + + for waiter := range m.waiters[routingKey] { + select { + case waiter <- 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.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 new file mode 100644 index 0000000000..50cf10af59 --- /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 { + t.Helper() + + client := redis_utils.SetupInstance(t) + storage := NewStorage(client) + go storage.Start(t.Context()) + t.Cleanup(storage.Close) + + return storage.subManager +} + +func TestSubscriptionManager_SubscribeAndDispatch(t *testing.T) { + t.Parallel() + + m := setupTestManager(t) + + ch, cleanup := m.subscribe("key1") + t.Cleanup(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") + t.Cleanup(cleanup1) + ch2, cleanup2 := m.subscribe("key2") + t.Cleanup(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) + for i := range numWaiters { + ch, cleanup := m.subscribe("shared-key") + t.Cleanup(cleanup) + channels[i] = ch + } + + 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") + t.Cleanup(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") + t.Cleanup(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 range goroutines { + wg.Go(func() { + 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): + } + }) + } + + 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(client) + go storage.Start(t.Context()) + t.Cleanup(storage.Close) + + routingKey := "test:routing:key" + ch, cleanup := storage.subManager.subscribe(routingKey) + 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(), globalTransitionNotifyChannel, 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(client) + go storage.Start(t.Context()) + t.Cleanup(storage.Close) + + ch, cleanup := storage.subManager.subscribe("my:sandbox:key") + t.Cleanup(cleanup) + + time.Sleep(50 * time.Millisecond) + + // Publish a message with a different routing key + err := client.Publish(t.Context(), globalTransitionNotifyChannel, "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 + } +} diff --git a/packages/api/internal/sandbox/storage/redis/utils.go b/packages/api/internal/sandbox/storage/redis/utils.go index 31452da8cb..b3b77e0c95 100644 --- a/packages/api/internal/sandbox/storage/redis/utils.go +++ b/packages/api/internal/sandbox/storage/redis/utils.go @@ -8,12 +8,18 @@ import ( const ( sandboxKeyPrefix = "sandbox:storage" - transitionKeyPrefix = "transition:" + transitionKeyPrefix = "transition" + notifySuffix = "notify" sandboxesKey = "sandboxes" indexKey = "index" ) 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") ) @@ -49,3 +55,11 @@ 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-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) +}