diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 6133d29db1..5faf43a6db 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -90,6 +90,12 @@ const ( // goroutines and lifecycle management. streamDropWarnInterval = 10 * time.Second + // bufferRetainGracePeriod is how long the message_part + // buffer is kept after processing completes. This gives + // cross-replica relay subscribers time to connect and + // snapshot the buffer before it is garbage-collected. + bufferRetainGracePeriod = 5 * time.Second + // DefaultMaxChatsPerAcquire is the maximum number of chats to // acquire in a single processOnce call. Batching avoids // waiting a full polling interval between acquisitions @@ -691,6 +697,13 @@ type chatStreamState struct { bufferLastWarnAt time.Time subscriberDropCount int64 subscriberLastWarnAt time.Time + // bufferRetainedAt records when processing completed and + // the buffer was retained for late-connecting relay + // subscribers. Zero while buffering is active. When + // non-zero, cleanupStreamIfIdle skips GC until the grace + // period expires so cross-replica relays can still + // snapshot the buffer. + bufferRetainedAt time.Time } // resetDropCounters zeroes the rate-limiting state for both buffer @@ -2681,12 +2694,23 @@ func (p *Server) getOrCreateStreamState(chatID uuid.UUID) *chatStreamState { // cleanupStreamIfIdle removes the chat entry from the sync.Map // when there are no subscribers and the stream is not buffering. +// When bufferRetainedAt is set, cleanup is deferred until the +// grace period expires so cross-replica relay subscribers can +// still snapshot the buffer. // The caller must hold state.mu. func (p *Server) cleanupStreamIfIdle(chatID uuid.UUID, state *chatStreamState) { - if !state.buffering && len(state.subscribers) == 0 { - p.chatStreams.Delete(chatID) - p.workspaceMCPToolsCache.Delete(chatID) + if state.buffering || len(state.subscribers) > 0 { + return } + // Keep stream state alive during the grace period so + // late-connecting relay subscribers can snapshot the + // buffer after the worker finishes processing. + if !state.bufferRetainedAt.IsZero() && + p.clock.Now().Before(state.bufferRetainedAt.Add(bufferRetainGracePeriod)) { + return + } + p.chatStreams.Delete(chatID) + p.workspaceMCPToolsCache.Delete(chatID) } func (p *Server) Subscribe( @@ -3567,15 +3591,20 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { streamState := p.getOrCreateStreamState(chat.ID) streamState.mu.Lock() streamState.buffer = nil + streamState.bufferRetainedAt = time.Time{} streamState.resetDropCounters() streamState.buffering = true streamState.mu.Unlock() defer func() { streamState.mu.Lock() - streamState.buffer = nil streamState.resetDropCounters() streamState.buffering = false - p.cleanupStreamIfIdle(chat.ID, streamState) + // Retain the buffer for a grace period so + // cross-replica relay subscribers can still snapshot + // it after processing completes. The buffer is + // cleared when the next processChat starts or when + // cleanupStreamIfIdle runs after the grace period. + streamState.bufferRetainedAt = p.clock.Now() streamState.mu.Unlock() }() @@ -4302,17 +4331,12 @@ func (p *Server) runChat( p.publishMessage(chat.ID, msg) } - // Clear the stream buffer now that the step is - // persisted. Late-joining subscribers will load - // these messages from the database instead. - if val, ok := p.chatStreams.Load(chat.ID); ok { - if ss, ok := val.(*chatStreamState); ok { - ss.mu.Lock() - ss.buffer = nil - ss.resetDropCounters() - ss.mu.Unlock() - } - } + // Do NOT clear the stream buffer here. Cross-replica + // relay subscribers may still need to snapshot buffered + // message_parts after processing completes. The buffer + // is bounded by maxStreamBufferSize and is cleared when + // the next processChat starts or when the stream state + // is garbage-collected after the retention grace period. return nil } diff --git a/enterprise/coderd/x/chatd/chatd.go b/enterprise/coderd/x/chatd/chatd.go index d16a6b0fff..22c63d3722 100644 --- a/enterprise/coderd/x/chatd/chatd.go +++ b/enterprise/coderd/x/chatd/chatd.go @@ -26,6 +26,12 @@ const RelaySourceHeader = "X-Coder-Relay-Source-Replica" const ( authorizationHeader = "Authorization" cookieHeader = "Cookie" + + // relayDrainTimeout is how long an established relay is + // kept open after the chat leaves running state, giving + // buffered snapshot events time to be forwarded before + // the relay is torn down. + relayDrainTimeout = 200 * time.Millisecond ) // MultiReplicaSubscribeConfig holds the dependencies for multi-replica chat @@ -169,6 +175,21 @@ func NewMultiReplicaSubscribeFn( var reconnectTimer *quartz.Timer var reconnectCh <-chan time.Time + // drainAndClose is set when the chat transitions away + // from running while a relay dial is still in progress. + // Instead of canceling the dial immediately, we let it + // complete so the snapshot of buffered message_parts + // can be forwarded to the subscriber. + var drainAndClose bool + + // Drain timer state. When the relay connects in + // drain-and-close mode, a short timer is started. + // During this window the normal relayPartsCh case + // forwards buffered snapshot events. When the timer + // fires the relay is torn down. + var drainTimer *quartz.Timer + var drainTimerCh <-chan time.Time + // Helper to close relay and stop any pending reconnect // timer. closeRelay := func() { @@ -200,6 +221,12 @@ func NewMultiReplicaSubscribeFn( reconnectTimer = nil reconnectCh = nil } + if drainTimer != nil { + drainTimer.Stop() + drainTimer = nil + drainTimerCh = nil + } + drainAndClose = false } // openRelayAsync dials the remote replica in a background @@ -335,16 +362,52 @@ func NewMultiReplicaSubscribeFn( // A nil parts channel signals the dial // failed — schedule a retry. if result.parts == nil { - scheduleRelayReconnect() + if drainAndClose { + // Dial failed and we were only + // waiting to drain — nothing to do. + drainAndClose = false + } else { + scheduleRelayReconnect() + } continue - } - // An async relay dial completed; swap + } // An async relay dial completed; swap // in the new relay channel. if relayCancel != nil { relayCancel() } relayParts = result.parts relayCancel = result.cancel + if drainAndClose { + // The chat is no longer running on + // the remote worker, but the dial + // completed. Verify no new worker + // has claimed the chat before we + // drain stale parts. + currentChat, dbErr := params.DB.GetChatByID(ctx, chatID) + if dbErr != nil { + logger.Warn(ctx, "failed to check chat status for relay drain", + slog.F("chat_id", chatID), + slog.Error(dbErr), + ) + } + if dbErr == nil && currentChat.Status == database.ChatStatusRunning && + currentChat.WorkerID.Valid && + currentChat.WorkerID.UUID != params.WorkerID { + // A new worker picked up the chat; + // discard the stale relay and let + // openRelayAsync handle the new one. + closeRelay() + } else { + // Chat is still idle — drain the + // buffered snapshot before closing. + if drainTimer != nil { + drainTimer.Stop() + } + drainTimer = cfg.clock().NewTimer(relayDrainTimeout, "drain") + drainTimerCh = drainTimer.C + drainAndClose = false + } + } case <-reconnectCh: reconnectCh = nil // Re-check whether the chat is still @@ -374,8 +437,31 @@ func NewMultiReplicaSubscribeFn( if sn.Status == database.ChatStatusRunning && sn.WorkerID != uuid.Nil && sn.WorkerID != params.WorkerID { openRelayAsync(sn.WorkerID) } else { - closeRelay() + switch { + case dialCancel != nil && relayParts == nil: + // In-progress dial: let it complete + // so its snapshot can be forwarded. + drainAndClose = true + case relayParts != nil: + // Active relay: give it a short + // window to deliver any remaining + // buffered parts before closing. + if drainTimer != nil { + drainTimer.Stop() + } + drainTimer = cfg.clock().NewTimer(relayDrainTimeout, "drain") + drainTimerCh = drainTimer.C + default: + closeRelay() + } } + case <-drainTimerCh: + drainTimerCh = nil + drainTimer = nil + closeRelay() + drainTimerCh = nil + drainTimer = nil + closeRelay() case event, ok := <-relayPartsCh: if !ok { if relayCancel != nil { diff --git a/enterprise/coderd/x/chatd/chatd_test.go b/enterprise/coderd/x/chatd/chatd_test.go index d1861fd561..c79f1de62f 100644 --- a/enterprise/coderd/x/chatd/chatd_test.go +++ b/enterprise/coderd/x/chatd/chatd_test.go @@ -1245,3 +1245,334 @@ func TestSubscribeRelayMultipleReconnects(t *testing.T) { consumePart("relay-3") require.GreaterOrEqual(t, int(callCount.Load()), 3) } + +// TestSubscribeRelayDialCanceledOnFastCompletion demonstrates a race +// condition in multi-replica chat streaming where the relay connection +// from the subscriber replica to the worker replica is canceled before +// it can be established because the worker completes processing before +// the async relay dial finishes. +// +// Scenario: +// 1. Subscriber subscribes to a chat while it's in waiting state (no relay). +// 2. User sends a message → chat becomes pending → worker picks it up. +// 3. Subscriber receives status=running via pubsub → enterprise opens relay async. +// 4. Worker completes quickly → publishes committed message + status=waiting. +// 5. Subscriber receives status=waiting → enterprise cancels the in-progress relay dial. +// 6. The relay was never established, so no message_part events were delivered. +// 7. The committed message arrives via pubsub (durable path), but streaming is lost. +// +// This reproduces the user-facing issue where refreshing the page is needed +// to see a response: the streaming tokens never arrive via the relay, and +// the response only appears after the full committed message is delivered. +func TestSubscribeRelayDialCanceledOnFastCompletion(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + workerID := uuid.New() + subscriberID := uuid.New() + + var dialAttempted atomic.Bool + + // Gate: closed when the worker finishes processing. + workerDone := make(chan struct{}) + + openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + if !req.Stream { + return chattest.OpenAINonStreamingResponse("fast-completion-relay-race") + } + return chattest.OpenAIStreamingResponse( + chattest.OpenAITextChunks("hello ", "world ", "from ", "the ", "worker")..., + ) + }) + + // Worker server with a 1-hour acquire interval so it only processes + // when explicitly woken by SendMessage's signalWake. + workerLogger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + worker := osschatd.New(osschatd.Config{ + Logger: workerLogger, + Database: db, + ReplicaID: workerID, + Pubsub: ps, + PendingChatAcquireInterval: time.Hour, + InFlightChatStaleAfter: testutil.WaitSuperLong, + }) + t.Cleanup(func() { + require.NoError(t, worker.Close()) + }) + + // Subscriber's relay dialer blocks until the worker finishes, + // simulating a slow relay dial (network latency between replicas). + // After the worker completes, the dialer connects to the worker + // to retrieve buffered parts from the retained buffer. + subscriber := newTestServer(t, db, ps, subscriberID, func( + ctx context.Context, + chatID uuid.UUID, + targetWorkerID uuid.UUID, + requestHeader http.Header, + ) ( + []codersdk.ChatStreamEvent, + <-chan codersdk.ChatStreamEvent, + func(), + error, + ) { + dialAttempted.Store(true) + // Block until the worker finishes processing, simulating + // a slow relay dial. + select { + case <-workerDone: + case <-ctx.Done(): + return nil, nil, nil, ctx.Err() + } + // Connect to the worker. The buffer is retained for a + // grace period after processing, so the relay still gets + // the message_part snapshot. + snapshot, relayEvents, cancel, ok := worker.Subscribe(ctx, chatID, requestHeader, math.MaxInt64) + if !ok { + return nil, nil, nil, xerrors.New("worker subscribe failed") + } + return snapshot, relayEvents, cancel, nil + }, nil) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + setOpenAIProviderBaseURL(ctx, t, db, openAIURL) + + // Create the chat in waiting state so the subscriber sees it + // before the worker picks it up (avoids the synchronous relay + // path in Subscribe). + chat := seedWaitingChat(ctx, t, db, user, model, "fast-completion-relay-race") + + // Subscribe from the subscriber replica while the chat is idle. + // No relay is opened because the chat is in waiting state. + _, events, subCancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0) + require.True(t, ok) + defer subCancel() + + // Send a message via the worker server to transition the chat to + // pending and wake the worker's processing loop. + _, err := worker.SendMessage(ctx, osschatd.SendMessageOptions{ + ChatID: chat.ID, + CreatedBy: user.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + // Wait for the worker to fully process the chat. + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusWaiting + }, testutil.WaitMedium, testutil.IntervalFast) + + // Release the relay dial now that the worker is done. + close(workerDone) + + // Collect all events that arrived at the subscriber. + var messageParts []string + var committedAssistantMsgs int + + // Drain events until we see both the committed message (via + // pubsub) and at least one streaming part (via relay + // drain-and-close). + require.Eventually(t, func() bool { + select { + case event := <-events: + switch event.Type { + case codersdk.ChatStreamEventTypeMessagePart: + if event.MessagePart != nil { + messageParts = append(messageParts, event.MessagePart.Part.Text) + } + case codersdk.ChatStreamEventTypeMessage: + if event.Message != nil && event.Message.Role == codersdk.ChatMessageRoleAssistant { + committedAssistantMsgs++ + } + } + return committedAssistantMsgs > 0 && len(messageParts) > 0 + default: + return false + } + }, testutil.WaitLong, testutil.IntervalFast) + + // The committed assistant message arrives via pubsub → DB query + // (durable path). + require.Equal(t, 1, committedAssistantMsgs, + "committed assistant message should arrive via pubsub durable path") + + // The relay dial was attempted when status=running arrived. + require.True(t, dialAttempted.Load(), + "relay dial should have been attempted when status changed to running") + + // Streaming parts are now received even though the relay was + // slower than the worker: the OSS buffer retention grace period + // keeps parts available, and the enterprise relay completes the + // dial (drain-and-close) instead of canceling it immediately. + require.NotEmpty(t, messageParts, + "streaming parts should be received via the relay even when the "+ + "worker completes before the relay is established") +} + +// TestSubscribeRelayEstablishedMidStream demonstrates that when the +// relay is established while the worker is still streaming, the +// subscriber receives buffered parts via the relay snapshot and live +// parts through the relay channel. +// +// This is the complementary test to TestSubscribeRelayDialCanceledOnFastCompletion: +// it shows the relay mechanism works correctly when timing is favorable +// (relay connects before the worker finishes), contrasting with the race +// condition where the relay is too slow. +func TestSubscribeRelayEstablishedMidStream(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + workerID := uuid.New() + subscriberID := uuid.New() + + // Gate: worker blocks after first streaming request until we + // release it. This gives the relay time to establish. + firstChunkEmitted := make(chan struct{}) + continueStreaming := make(chan struct{}) + + openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + if !req.Stream { + return chattest.OpenAINonStreamingResponse("mid-stream-relay") + } + // Signal that the first streaming request was received, + // then block until released. + select { + case <-firstChunkEmitted: + default: + close(firstChunkEmitted) + } + <-continueStreaming + return chattest.OpenAIStreamingResponse( + chattest.OpenAITextChunks("continued ", "response")..., + ) + }) + + // Worker with a 1-hour acquire interval; only processes when + // explicitly woken. + workerLogger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + worker := osschatd.New(osschatd.Config{ + Logger: workerLogger, + Database: db, + ReplicaID: workerID, + Pubsub: ps, + PendingChatAcquireInterval: time.Hour, + InFlightChatStaleAfter: testutil.WaitSuperLong, + }) + t.Cleanup(func() { + require.NoError(t, worker.Close()) + }) + + // Subscriber's dialer connects to the worker with no delay. + // This simulates a relay that succeeds promptly. + subscriber := newTestServer(t, db, ps, subscriberID, func( + ctx context.Context, + chatID uuid.UUID, + targetWorkerID uuid.UUID, + requestHeader http.Header, + ) ( + []codersdk.ChatStreamEvent, + <-chan codersdk.ChatStreamEvent, + func(), + error, + ) { + if targetWorkerID != workerID { + return nil, nil, nil, xerrors.Errorf("unexpected relay target %s", targetWorkerID) + } + snapshot, relayEvents, cancel, ok := worker.Subscribe(ctx, chatID, requestHeader, math.MaxInt64) + if !ok { + return nil, nil, nil, xerrors.New("worker subscribe failed") + } + return snapshot, relayEvents, cancel, nil + }, nil) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + setOpenAIProviderBaseURL(ctx, t, db, openAIURL) + + // Create the chat in waiting state. + chat := seedWaitingChat(ctx, t, db, user, model, "mid-stream-relay") + + // Subscribe from the subscriber replica while the chat is idle. + _, events, subCancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0) + require.True(t, ok) + defer subCancel() + + // Send a message to make the chat pending and wake the worker. + _, err := worker.SendMessage(ctx, osschatd.SendMessageOptions{ + ChatID: chat.ID, + CreatedBy: user.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + // Wait for the worker to reach the LLM (first streaming request). + select { + case <-firstChunkEmitted: + case <-ctx.Done(): + t.Fatal("timed out waiting for worker to start streaming") + } + + // Wait for the subscriber to receive the running status, which + // triggers the relay. Because the dialer is non-blocking, the + // relay establishes promptly. + require.Eventually(t, func() bool { + select { + case event := <-events: + return event.Type == codersdk.ChatStreamEventTypeStatus && + event.Status != nil && + event.Status.Status == codersdk.ChatStatusRunning + default: + return false + } + }, testutil.WaitMedium, testutil.IntervalFast) + + // Now release the worker to continue streaming. + close(continueStreaming) + + // Wait for the worker to complete. + require.Eventually(t, func() bool { + fromDB, dbErr := db.GetChatByID(ctx, chat.ID) + if dbErr != nil { + return false + } + return fromDB.Status == database.ChatStatusWaiting + }, testutil.WaitMedium, testutil.IntervalFast) + + // Collect remaining events. + var messageParts []string + var hasCommittedMsg bool + + require.Eventually(t, func() bool { + select { + case event := <-events: + switch event.Type { + case codersdk.ChatStreamEventTypeMessagePart: + if event.MessagePart != nil { + messageParts = append(messageParts, event.MessagePart.Part.Text) + } + case codersdk.ChatStreamEventTypeMessage: + if event.Message != nil && event.Message.Role == codersdk.ChatMessageRoleAssistant { + hasCommittedMsg = true + } + } + return hasCommittedMsg + default: + return false + } + }, testutil.WaitLong, testutil.IntervalFast) + + // The committed message arrives via pubsub. + require.True(t, hasCommittedMsg, + "committed assistant message should arrive") + + // When the relay is established mid-stream, streaming parts + // SHOULD be received through the relay. This contrasts with + // TestSubscribeRelayDialCanceledOnFastCompletion where no parts + // arrive because the relay is never established. + require.NotEmpty(t, messageParts, + "streaming parts should be received when relay establishes while worker is still streaming") +}