From 5a5cd79c4cb38a2ae3c8c8a51009de433d82ff46 Mon Sep 17 00:00:00 2001 From: Kyle Carberry Date: Tue, 12 May 2026 00:30:38 -0400 Subject: [PATCH] fix: drop buffered chat parts after their durable message commits (#25164) --- coderd/x/chatd/chatd.go | 289 ++++++-------- coderd/x/chatd/chatd_internal_test.go | 482 ++++++++---------------- coderd/x/chatd/chatd_test.go | 70 +++- enterprise/coderd/x/chatd/chatd_test.go | 267 ++----------- 4 files changed, 356 insertions(+), 752 deletions(-) diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 5429e47ba3..e5703be92e 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -86,9 +86,10 @@ const ( maxStreamBufferSize = 10000 // RelaySentinelAfterID is the after_id sentinel used by cross-replica // relay subscribers. It instructs the peer to skip the durable DB - // snapshot and deliver only in-flight buffered parts. The sentinel - // also disables snapshotBufferLocked's redundant-part filter so - // relays receive every part the worker has buffered (see PR #24031). + // snapshot and only deliver buffered message_part events. The + // buffer itself filters committed parts out (see snapshotBufferLocked), + // so the sentinel resolves to "send me any in-progress streaming + // parts you have; I will receive durable messages through pubsub." RelaySentinelAfterID = math.MaxInt64 // maxDurableMessageCacheSize caps the number of recent durable message // events cached per chat for same-replica stream catch-up. @@ -114,10 +115,15 @@ 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 is how long the per-chat stream + // state is kept after processing completes. The retained + // state lets late-connecting cross-replica relay subscribers + // register against the live stream before the next worker + // run starts, preventing a race between cleanupStreamIfIdle + // and subscriber registration. The buffer itself is no + // longer useful at this point: every part has been claimed + // by its durable assistant message and is filtered out of + // the subscriber snapshot. bufferRetainGracePeriod = 5 * time.Second // chatStreamControlFetchTimeout bounds subscriber-owned // control-path DB reads when the caller has no deadline. @@ -1099,21 +1105,22 @@ type SubscribeFnParams struct { Logger slog.Logger } -// bufferedStreamPart is a buffered message_part event tagged with the -// most recently committed assistant message ID at the moment it was -// appended. Subscribers can use the checkpoint to skip parts that -// belong to turns they have already received via durable -// `message` events. +// bufferedStreamPart is a buffered message_part event with its +// committed-message linkage. Parts that have not yet been claimed by +// a durable assistant message carry committedMessageID == 0 and are +// considered "in progress"; when an assistant message is published +// every still-in-progress part is claimed by that durable message +// ID, marking the part as redundant for any subscriber that will +// receive the durable message via REST or pubsub. type bufferedStreamPart struct { event codersdk.ChatStreamEvent - // checkpoint is the chatStreamState.lastCommittedAssistantMessageID - // value at the time this part was buffered. A subscriber whose - // cursor is past this checkpoint already has the durable assistant - // message for the turn this part belongs to, so the part is - // redundant. The cursor is clamped to the current checkpoint at - // snapshot time, so tool/user message IDs in the cursor cannot - // over-drop parts from an in-progress assistant turn. - checkpoint int64 + // committedMessageID is the durable assistant message ID that + // claimed this part, or 0 while the part belongs to the + // in-progress turn. snapshotBufferLocked drops parts with + // committedMessageID != 0 because the subscriber will receive + // the durable message through a different channel (REST snapshot, + // initial DB query in SubscribeAuthorized, or pubsub). + committedMessageID int64 } type chatStreamState struct { @@ -1132,18 +1139,16 @@ type chatStreamState struct { // to retry. currentRetry *codersdk.ChatStreamRetry // bufferRetainedAt records when processing completed and - // the buffer was retained for late-connecting relay - // subscribers. Zero while buffering is active. When + // the per-chat stream state entered the post-completion + // grace window. 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. + // period expires so cross-replica relay subscribers can + // register without racing state deletion. The buffer + // itself does not deliver content here: every part is + // claimed by a durable assistant message before + // bufferRetainedAt is set, so snapshotBufferLocked + // returns no parts during the grace window. bufferRetainedAt time.Time - // lastCommittedAssistantMessageID tracks the highest assistant - // durable message ID published for this chat on this replica. - // publishToStream tags each buffered message_part with this - // value so subscribeToStream can filter out parts belonging to - // already-committed turns. - lastCommittedAssistantMessageID int64 } // heartbeatEntry tracks a single chat's cancel function and workspace @@ -4178,8 +4183,10 @@ func (p *Server) publishToStream(chatID uuid.UUID, event codersdk.ChatStreamEven state.buffer = state.buffer[1:] } state.buffer = append(state.buffer, bufferedStreamPart{ - event: event, - checkpoint: state.lastCommittedAssistantMessageID, + event: event, + // committedMessageID stays 0 here: the part belongs to + // the in-progress turn until publishMessage claims it + // with the committed assistant message ID. }) } subscribers := make([]chan codersdk.ChatStreamEvent, 0, len(state.subscribers)) @@ -4264,50 +4271,30 @@ func (p *Server) getCachedDurableMessages( } // snapshotBufferLocked returns the buffered message_part events that -// the caller should receive in their initial snapshot, filtered by -// the requested cursor. +// the caller should receive in their initial snapshot. // -// The cursor is clamped to lastCommittedAssistantMessageID so a -// cursor that points at a tool or user message ID past the last -// committed assistant turn cannot drop parts from the in-progress -// turn. The filter then drops parts whose checkpoint is below the -// clamped cursor; those parts belong to assistant turns the -// subscriber already has via durable `message` events. +// Parts whose committedMessageID != 0 are dropped: those parts were +// claimed by a durable assistant message that the subscriber will +// receive through a different channel (REST snapshot, the initial DB +// query in SubscribeAuthorized, or pubsub catch-up). Delivering them +// here would render the same content twice on the client, once in the +// streaming UI and once as a durable message. // -// The caller must hold the per-chat stream state lock. See -// subscribeToStream for the documented afterMessageID semantics. -func snapshotBufferLocked( - buffer []bufferedStreamPart, - afterMessageID int64, - lastCommittedAssistantMessageID int64, -) []codersdk.ChatStreamEvent { +// Every caller receives the same view: in-progress parts are always +// delivered and committed parts are always dropped, regardless of +// cursor or relay sentinel. This keeps the buffer free of duplicate +// work for every subscriber, including cross-replica relay +// subscribers whose user-facing peers receive the durable message +// via pubsub. +// +// The caller must hold the per-chat stream state lock. +func snapshotBufferLocked(buffer []bufferedStreamPart) []codersdk.ChatStreamEvent { if len(buffer) == 0 { return nil } - // Compute the effective cursor used to drop redundant parts. - // - afterMessageID <= 0 ("no cursor; deliver everything") and - // the RelaySentinelAfterID both disable filtering. - // - Otherwise clamp the cursor to lastCommittedAssistantMessageID - // so a tool/user cursor past the last assistant turn cannot - // over-drop parts from the in-progress assistant turn. We can - // only be confident a buffered part is redundant when the - // cursor is at or past its checkpoint AND the checkpoint maps - // to a turn the subscriber already has via durable messages. - // - If lastCommittedAssistantMessageID is still zero (e.g. - // fresh state after cleanup), no buffered part can be proven - // redundant, so deliver everything. - var effectiveCursor int64 - switch { - case afterMessageID <= 0, afterMessageID == RelaySentinelAfterID: - effectiveCursor = 0 - case lastCommittedAssistantMessageID < afterMessageID: - effectiveCursor = lastCommittedAssistantMessageID - default: - effectiveCursor = afterMessageID - } snapshot := make([]codersdk.ChatStreamEvent, 0, len(buffer)) for _, part := range buffer { - if effectiveCursor > 0 && part.checkpoint < effectiveCursor { + if part.committedMessageID != 0 { continue } snapshot = append(snapshot, part.event) @@ -4316,25 +4303,17 @@ func snapshotBufferLocked( } // subscribeToStream registers a subscriber to the per-chat in-memory -// stream and returns a filtered snapshot of currently-buffered -// message_part events plus the current retry phase, the live -// subscriber channel, and a cancel func. +// stream and returns a snapshot of currently in-progress message_part +// events plus the current retry phase, the live subscriber channel, +// and a cancel func. // -// afterMessageID semantics: -// - 0: no filter; the full buffer snapshot is returned. -// New browser sessions use this and only see parts for the -// currently-streaming turn (the buffer is cleared at the -// start of each processChat run). -// - RelaySentinelAfterID: no filter; cross-replica relays pass -// this sentinel to skip the durable DB snapshot while still -// receiving all in-flight buffered parts. -// - 0 < afterMessageID < RelaySentinelAfterID: parts whose -// checkpoint is less than the cursor are dropped from the -// snapshot. The cursor is clamped to the per-chat -// lastCommittedAssistantMessageID before filtering so cursors -// that point at tool/user message IDs past the last committed -// assistant turn cannot over-drop in-progress parts. -func (p *Server) subscribeToStream(chatID uuid.UUID, afterMessageID int64) ( +// Parts that were claimed by a committed durable assistant message +// (committedMessageID != 0) are excluded from the snapshot. The +// subscriber will receive those durable messages through the REST +// snapshot, the initial DB query in SubscribeAuthorized, or pubsub, +// so re-delivering their constituent parts here would render the +// same content twice. +func (p *Server) subscribeToStream(chatID uuid.UUID) ( []codersdk.ChatStreamEvent, *codersdk.ChatStreamRetry, <-chan codersdk.ChatStreamEvent, @@ -4342,7 +4321,7 @@ func (p *Server) subscribeToStream(chatID uuid.UUID, afterMessageID int64) ( ) { state := p.getOrCreateStreamState(chatID) state.mu.Lock() - snapshot := snapshotBufferLocked(state.buffer, afterMessageID, state.lastCommittedAssistantMessageID) + snapshot := snapshotBufferLocked(state.buffer) var currentRetry *codersdk.ChatStreamRetry if state.currentRetry != nil { retryCopy := *state.currentRetry @@ -4400,8 +4379,8 @@ func (p *Server) cleanupStreamIfIdle(chatID uuid.UUID, state *chatStreamState) b return false } // Keep stream state alive during the grace period so - // late-connecting relay subscribers can snapshot the - // buffer after the worker finishes processing. + // late-connecting cross-replica relay subscribers can + // register against this chat before GC. if !state.bufferRetainedAt.IsZero() && p.clock.Now().Before(state.bufferRetainedAt.Add(bufferRetainGracePeriod)) { return false @@ -4645,7 +4624,7 @@ func (p *Server) SubscribeAuthorized( // persisted messages. Capture the current retry phase under the same // lock so the transient snapshot and subscriber registration reflect // a single moment in time. - localSnapshot, localRetry, localParts, localCancel := p.subscribeToStream(chatID, afterMessageID) + localSnapshot, localRetry, localParts, localCancel := p.subscribeToStream(chatID) // Merge all event sources. mergedCtx, mergedCancel := context.WithCancel(ctx) @@ -5326,84 +5305,48 @@ func (p *Server) publishMessage(chatID uuid.UUID, message database.ChatMessage) Message: &sdkMessage, } p.cacheDurableMessage(chatID, event) - p.advanceAssistantCheckpoint(chatID, message) + // Claim every still-in-progress buffered message_part for this + // durable assistant message BEFORE publishing it, so any new + // subscriber that races publishEvent below takes a buffer + // snapshot in which the parts for this turn are already + // suppressed. Existing subscribers already received the + // constituent parts on the live channel; the frontend + // dedupes those against the durable message via + // clearStreamState in the same batch. + p.claimCommittedParts(chatID, message) p.publishEvent(chatID, event) p.publishChatStreamNotify(chatID, coderdpubsub.ChatStreamNotifyMessage{ AfterMessageID: message.ID - 1, }) } -// seedAssistantCheckpoint initializes the per-chat checkpoint from -// the last durable assistant message ID before any parts are -// buffered for this run. This closes the cleanup-and-recreate race -// where a freshly stored chatStreamState would start with -// lastCommittedAssistantMessageID = 0, tagging every part with -// checkpoint 0 and forcing snapshotBufferLocked to deliver the -// entire buffer to every reconnecting subscriber. -// -// On lookup error or when there is no prior assistant message, the -// checkpoint stays at its current value (either zero for a brand-new -// state, or the value carried forward from a prior run on this -// replica). This is safe: a zero checkpoint produces an over- -// inclusive snapshot, not data loss. -func (p *Server) seedAssistantCheckpoint( - ctx context.Context, - chatID uuid.UUID, - state *chatStreamState, - logger slog.Logger, -) { - // Use a short timeout so a stalled DB does not block the - // start of a chat run. The seed is best-effort: missing it - // only degrades the snapshot filter to "deliver everything", - // which is the prior behavior. - // - // The seed reads the last assistant message ID to bound the - // in-memory checkpoint; it never returns user data. The - // system context is required because processChat runs without - // an actor and the durable read is part of the chat worker - // loop. There is no authorization to skip; the chat row was - // already authorized before processChat was scheduled. - //nolint:gocritic // chatd worker reads its own durable state to seed the in-memory checkpoint; no user context exists here. - lookupCtx, cancel := context.WithTimeout( - dbauthz.AsSystemRestricted(ctx), - 5*time.Second, - ) - defer cancel() - last, err := p.db.GetLastChatMessageByRole(lookupCtx, database.GetLastChatMessageByRoleParams{ - ChatID: chatID, - Role: database.ChatMessageRoleAssistant, - }) - if errors.Is(err, sql.ErrNoRows) { - return - } - if err != nil { - logger.Warn(ctx, "failed to seed assistant checkpoint", slog.Error(err)) - return - } - state.mu.Lock() - defer state.mu.Unlock() - if last.ID > state.lastCommittedAssistantMessageID { - state.lastCommittedAssistantMessageID = last.ID - } -} - -// advanceAssistantCheckpoint bumps the per-chat checkpoint when an -// assistant durable message is published. Subsequent buffered -// message_part events are tagged with the new checkpoint so -// subscribeToStream can filter parts belonging to already-committed -// turns when the subscriber's cursor is past the checkpoint. +// claimCommittedParts walks the chat's buffered message_part events +// and assigns every in-progress part (committedMessageID == 0) to +// the supplied assistant message ID. Subsequent subscriber snapshots +// drop those parts so a reconnecting client does not re-render the +// content of an assistant turn that has already been delivered as a +// durable message via REST or pubsub. // // Tool and user messages do not end an assistant streaming turn, so -// the checkpoint is only advanced for assistant-role messages. -func (p *Server) advanceAssistantCheckpoint(chatID uuid.UUID, message database.ChatMessage) { +// only assistant-role messages claim parts. +func (p *Server) claimCommittedParts(chatID uuid.UUID, message database.ChatMessage) { if message.Role != database.ChatMessageRoleAssistant { return } - state := p.getOrCreateStreamState(chatID) + val, ok := p.chatStreams.Load(chatID) + if !ok { + return + } + state, ok := val.(*chatStreamState) + if !ok { + return + } state.mu.Lock() defer state.mu.Unlock() - if message.ID > state.lastCommittedAssistantMessageID { - state.lastCommittedAssistantMessageID = message.ID + for i := range state.buffer { + if state.buffer[i].committedMessageID == 0 { + state.buffer[i].committedMessageID = message.ID + } } } @@ -5857,19 +5800,7 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { streamState.bufferRetainedAt = time.Time{} streamState.resetDropCounters() streamState.buffering = true - // lastCommittedAssistantMessageID is intentionally NOT reset - // here: the checkpoint is lifetime-scoped across runs so that - // after a state was reaped and a new run starts, reconnecting - // subscribers can still filter parts from prior turns once we - // re-seed it below. streamState.mu.Unlock() - // Seed the checkpoint from the durable store so that after a - // cleanupStreamIfIdle reaped the previous state, the very - // first parts buffered by this run are not tagged with - // checkpoint=0 (which would make snapshotBufferLocked deliver - // the full buffer to every reconnecting subscriber regardless - // of their cursor). - p.seedAssistantCheckpoint(ctx, chat.ID, streamState, logger) defer func() { streamState.mu.Lock() // Fallback cleanup for exit paths that return before a @@ -5877,11 +5808,18 @@ func (p *Server) processChat(ctx context.Context, chat database.Chat) { streamState.currentRetry = nil streamState.resetDropCounters() streamState.buffering = false - // Retain the buffer for a grace period so - // cross-replica relay subscribers can still snapshot - // it after processing completes. The buffer is + // Retain the per-chat stream state for a grace period + // so cross-replica relay subscribers can register + // against this chat after processing completes, + // without racing cleanupStreamIfIdle. The buffer is // cleared when the next processChat starts or when - // cleanupStreamIfIdle runs after the grace period. + // cleanupStreamIfIdle runs after the grace period; on + // the normal-completion path every part has been + // claimed by its durable assistant message, so the + // snapshot is empty. On error or panic exit some parts + // may still be in-progress; those are likewise + // discarded when the buffer is cleared, and the + // frontend recovers via the next REST snapshot. streamState.bufferRetainedAt = p.clock.Now() streamState.mu.Unlock() }() @@ -7264,9 +7202,10 @@ func (p *Server) runChat( } } - // Do NOT clear the stream buffer here. Cross-replica - // relay subscribers may still need to snapshot buffered - // message_parts after processing completes. The buffer + // Do NOT clear the stream buffer here. The per-chat + // stream state must remain alive for the post-completion + // grace window so cross-replica relay subscribers can + // register without racing cleanupStreamIfIdle. 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. diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index f4a804c475..bbe63b8613 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -4,7 +4,6 @@ import ( "context" "database/sql" "encoding/json" - "math" "sync" "testing" "time" @@ -3503,9 +3502,6 @@ func TestProcessChat_IgnoresStaleControlNotification(t *testing.T) { database.ChatUsageLimitConfig{}, sql.ErrNoRows, ).AnyTimes() db.EXPECT().GetChatMessagesForPromptByChatID(gomock.Any(), chatID).Return(nil, nil).AnyTimes() - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), gomock.Any()).Return( - database.ChatMessage{}, sql.ErrNoRows, - ).AnyTimes() chat := database.Chat{ID: chatID, LastModelConfigID: uuid.New()} done := make(chan struct{}) @@ -3758,7 +3754,7 @@ func TestSubscribeCancelDuringGrace_ReapedBySweep(t *testing.T) { // Real subscribeToStream cancel path: the WS subscriber detach // that leaks in prod. - snapshot, currentRetry, events, cancelSub := server.subscribeToStream(chatID, 0) + snapshot, currentRetry, events, cancelSub := server.subscribeToStream(chatID) require.Len(t, snapshot, 1) require.Nil(t, currentRetry) require.NotNil(t, events) @@ -5280,9 +5276,6 @@ func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) { database.ChatUsageLimitConfig{}, sql.ErrNoRows, ).AnyTimes() db.EXPECT().GetChatMessagesForPromptByChatID(gomock.Any(), chatID).Return(nil, nil).AnyTimes() - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), gomock.Any()).Return( - database.ChatMessage{}, sql.ErrNoRows, - ).AnyTimes() // The deferred cleanup transaction: InsertChatMessages fails, // so UpdateChatStatus must NOT be called. @@ -5367,11 +5360,12 @@ func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) { } } -// makeBufferedPart is a small constructor for buffered message_part +// makeInProgressPart is a small constructor for buffered message_part // fixtures used by snapshotBufferLocked / subscribeToStream tests. It -// embeds the checkpoint and a recognizable text body so failing -// assertions can identify which part survived the filter. -func makeBufferedPart(checkpoint int64, text string) bufferedStreamPart { +// builds an in-progress part (committedMessageID == 0) with a +// recognizable text body so failing assertions can identify which +// part survived the filter. +func makeInProgressPart(text string) bufferedStreamPart { return bufferedStreamPart{ event: codersdk.ChatStreamEvent{ Type: codersdk.ChatStreamEventTypeMessagePart, @@ -5380,10 +5374,17 @@ func makeBufferedPart(checkpoint int64, text string) bufferedStreamPart { Part: codersdk.ChatMessageText(text), }, }, - checkpoint: checkpoint, } } +// makeCommittedPart builds a part already claimed by the given +// durable assistant message ID. +func makeCommittedPart(committedID int64, text string) bufferedStreamPart { + p := makeInProgressPart(text) + p.committedMessageID = committedID + return p +} + func partText(event codersdk.ChatStreamEvent) string { if event.MessagePart == nil { return "" @@ -5391,126 +5392,46 @@ func partText(event codersdk.ChatStreamEvent) string { return event.MessagePart.Part.Text } -// TestSnapshotBufferLocked_FiltersStaleParts is the core contract: -// when a subscriber passes a real cursor, parts whose checkpoint is -// less than the cursor are dropped from the snapshot. Parts at or -// past the cursor are delivered. -func TestSnapshotBufferLocked_FiltersStaleParts(t *testing.T) { +// TestSnapshotBufferLocked_DropsCommittedParts asserts the core +// dedup contract: parts that were claimed by a durable assistant +// message (committedMessageID != 0) are dropped from the snapshot +// because the subscriber will receive that durable message through +// the REST snapshot, the initial DB query, or pubsub. +func TestSnapshotBufferLocked_DropsCommittedParts(t *testing.T) { t.Parallel() buffer := []bufferedStreamPart{ - makeBufferedPart(10, "stale-1"), - makeBufferedPart(10, "stale-2"), - makeBufferedPart(20, "boundary-1"), - makeBufferedPart(20, "boundary-2"), - makeBufferedPart(30, "fresh-1"), + makeCommittedPart(100, "turnA-1"), + makeCommittedPart(100, "turnA-2"), + makeCommittedPart(200, "turnB-1"), + makeInProgressPart("in-progress-1"), + makeInProgressPart("in-progress-2"), } - // Cursor matches a real assistant checkpoint, so the effective - // cursor is the requested cursor unchanged. - snapshot := snapshotBufferLocked(buffer, 20, 30) - - require.Len(t, snapshot, 3, - "only parts checkpointed at >= afterMessageID should be kept") - require.Equal(t, "boundary-1", partText(snapshot[0])) - require.Equal(t, "boundary-2", partText(snapshot[1])) - require.Equal(t, "fresh-1", partText(snapshot[2])) -} - -// TestSnapshotBufferLocked_ClampsCursorToLastCommittedCheckpoint -// guards against DEREM-1: a subscriber whose cursor points at a -// tool or user message ID past the most recent committed assistant -// turn must not over-drop parts from the in-progress assistant -// turn. The filter clamps the cursor down to the latest assistant -// checkpoint so those in-progress parts survive. -func TestSnapshotBufferLocked_ClampsCursorToLastCommittedCheckpoint(t *testing.T) { - t.Parallel() - - // Turn A committed at assistant message 100, then tool - // messages 101..103 followed. Turn B is now streaming and its - // parts are tagged with checkpoint=100 (no new assistant turn - // has been committed yet on this replica). - buffer := []bufferedStreamPart{ - makeBufferedPart(100, "turnB-part-1"), - makeBufferedPart(100, "turnB-part-2"), - } - - // Client reloaded chat via REST and saw the latest message - // (a tool result at id=103), then reconnected with cursor=103. - // Without clamping, the filter would drop every turn B part - // because checkpoint (100) < afterMessageID (103). - snapshot := snapshotBufferLocked(buffer, 103, 100) + snapshot := snapshotBufferLocked(buffer) require.Len(t, snapshot, 2, - "cursor past the last assistant checkpoint must be clamped down so in-progress parts survive") - require.Equal(t, "turnB-part-1", partText(snapshot[0])) - require.Equal(t, "turnB-part-2", partText(snapshot[1])) + "only in-progress (committedMessageID == 0) parts should be kept") + require.Equal(t, "in-progress-1", partText(snapshot[0])) + require.Equal(t, "in-progress-2", partText(snapshot[1])) } -// TestSnapshotBufferLocked_ZeroCheckpointReturnsAll guards against -// DEREM-2: a freshly created chatStreamState (after -// cleanupStreamIfIdle reaped the previous state and seeding from DB -// has not yet run) has lastCommittedAssistantMessageID = 0. With a -// zero checkpoint, no buffered part can be proven redundant, so the -// full buffer must be returned regardless of the requested cursor. -func TestSnapshotBufferLocked_ZeroCheckpointReturnsAll(t *testing.T) { +// TestSnapshotBufferLocked_AllInProgressReturnsAll covers the +// fresh-load convention: when no assistant message has committed +// yet, every buffered part is in-progress and must be delivered. +func TestSnapshotBufferLocked_AllInProgressReturnsAll(t *testing.T) { t.Parallel() buffer := []bufferedStreamPart{ - makeBufferedPart(0, "a"), - makeBufferedPart(0, "b"), - makeBufferedPart(0, "c"), + makeInProgressPart("a"), + makeInProgressPart("b"), + makeInProgressPart("c"), } - snapshot := snapshotBufferLocked(buffer, 999, 0) + snapshot := snapshotBufferLocked(buffer) require.Len(t, snapshot, 3, - "lastCommittedAssistantMessageID==0 must disable the filter to avoid losing the entire in-progress turn") - require.Equal(t, "a", partText(snapshot[0])) - require.Equal(t, "b", partText(snapshot[1])) - require.Equal(t, "c", partText(snapshot[2])) -} - -// TestSnapshotBufferLocked_ZeroCursorReturnsAll covers the -// fresh-load convention: callers without a cursor get the full -// buffer. Buffering is reset at the start of every processChat run, -// so the buffer only ever contains parts from the current turn in -// this path. -func TestSnapshotBufferLocked_ZeroCursorReturnsAll(t *testing.T) { - t.Parallel() - - buffer := []bufferedStreamPart{ - makeBufferedPart(10, "a"), - makeBufferedPart(20, "b"), - makeBufferedPart(30, "c"), - } - - snapshot := snapshotBufferLocked(buffer, 0, 30) - - require.Len(t, snapshot, 3, - "afterMessageID == 0 means 'no cursor'; the full buffer must be returned") - require.Equal(t, "a", partText(snapshot[0])) - require.Equal(t, "b", partText(snapshot[1])) - require.Equal(t, "c", partText(snapshot[2])) -} - -// TestSnapshotBufferLocked_RelaySentinelReturnsAll: cross-replica -// relay dials with after_id=RelaySentinelAfterID to skip the durable -// DB snapshot. The buffer snapshot must NOT be filtered for that -// sentinel; otherwise the relay race PR #24031 fixed comes back. -func TestSnapshotBufferLocked_RelaySentinelReturnsAll(t *testing.T) { - t.Parallel() - - buffer := []bufferedStreamPart{ - makeBufferedPart(10, "a"), - makeBufferedPart(20, "b"), - makeBufferedPart(30, "c"), - } - - snapshot := snapshotBufferLocked(buffer, RelaySentinelAfterID, 30) - - require.Len(t, snapshot, 3, - "the relay sentinel must NOT filter the buffer") + "all in-progress parts must be delivered to the subscriber") require.Equal(t, "a", partText(snapshot[0])) require.Equal(t, "b", partText(snapshot[1])) require.Equal(t, "c", partText(snapshot[2])) @@ -5522,16 +5443,33 @@ func TestSnapshotBufferLocked_RelaySentinelReturnsAll(t *testing.T) { func TestSnapshotBufferLocked_EmptyBufferReturnsNil(t *testing.T) { t.Parallel() - require.Nil(t, snapshotBufferLocked(nil, 0, 0)) - require.Nil(t, snapshotBufferLocked(nil, 42, 30)) - require.Nil(t, snapshotBufferLocked([]bufferedStreamPart{}, 42, 30)) + require.Nil(t, snapshotBufferLocked(nil)) + require.Nil(t, snapshotBufferLocked([]bufferedStreamPart{})) } -// TestPublishToStream_TagsPartsWithCurrentCheckpoint verifies that -// parts buffered while the chat is streaming carry the current -// committed-assistant-message-ID checkpoint. Subscribers can then -// filter against this value. -func TestPublishToStream_TagsPartsWithCurrentCheckpoint(t *testing.T) { +// TestSnapshotBufferLocked_AllCommittedReturnsEmpty covers the +// natural resting point after an assistant turn commits and before +// the next turn starts streaming: every buffered part has been +// claimed and must be filtered out. The snapshot must be empty so +// reconnecting subscribers do not re-render content that is already +// available as a durable message. +func TestSnapshotBufferLocked_AllCommittedReturnsEmpty(t *testing.T) { + t.Parallel() + + buffer := []bufferedStreamPart{ + makeCommittedPart(100, "a"), + makeCommittedPart(100, "b"), + makeCommittedPart(200, "c"), + } + + require.Empty(t, snapshotBufferLocked(buffer)) +} + +// TestPublishToStream_AppendsAsInProgress verifies that parts +// buffered while the chat is streaming are tagged as in-progress +// (committedMessageID == 0) until publishMessage claims them via a +// committed assistant message. +func TestPublishToStream_AppendsAsInProgress(t *testing.T) { t.Parallel() mClock := quartz.NewMock(t) @@ -5542,9 +5480,8 @@ func TestPublishToStream_TagsPartsWithCurrentCheckpoint(t *testing.T) { chatID := uuid.New() state := &chatStreamState{ - buffering: true, - subscribers: map[uuid.UUID]chan codersdk.ChatStreamEvent{}, - lastCommittedAssistantMessageID: 100, + buffering: true, + subscribers: map[uuid.UUID]chan codersdk.ChatStreamEvent{}, } server.chatStreams.Store(chatID, state) @@ -5559,206 +5496,138 @@ func TestPublishToStream_TagsPartsWithCurrentCheckpoint(t *testing.T) { state.mu.Lock() defer state.mu.Unlock() require.Len(t, state.buffer, 1) - require.Equal(t, int64(100), state.buffer[0].checkpoint, - "part must be tagged with the current checkpoint at append time") + require.Equal(t, int64(0), state.buffer[0].committedMessageID, + "newly buffered parts must be in-progress until publishMessage claims them") require.Equal(t, "hello", partText(state.buffer[0].event)) } -// TestAdvanceAssistantCheckpoint covers the per-role behavior of -// advanceAssistantCheckpoint: -// - assistant messages advance the checkpoint monotonically. -// - tool / user messages leave the checkpoint untouched. -// - older assistant IDs (out-of-order publication) do not move -// the checkpoint backwards. -func TestAdvanceAssistantCheckpoint(t *testing.T) { +// TestClaimCommittedParts covers the per-role behavior of +// claimCommittedParts: +// - assistant messages claim every in-progress part with the +// committed message ID. +// - tool / user messages do not claim parts. +// - parts already claimed by an earlier assistant message are not +// re-claimed. +// - a chat with no live state is a no-op (does not panic). +func TestClaimCommittedParts(t *testing.T) { t.Parallel() - server := &Server{ - logger: slogtest.Make(t, nil), - clock: quartz.NewMock(t), - } - - chatID := uuid.New() - state := server.getOrCreateStreamState(chatID) - - requireCheckpoint := func(want int64) { - t.Helper() - state.mu.Lock() - got := state.lastCommittedAssistantMessageID - state.mu.Unlock() - require.Equal(t, want, got) - } - - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ - ID: 100, - Role: database.ChatMessageRoleAssistant, - }) - requireCheckpoint(100) - - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ - ID: 200, - Role: database.ChatMessageRoleAssistant, - }) - requireCheckpoint(200) - - // Out-of-order: an older ID must not move the checkpoint - // backwards (defends against publish reordering). - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ - ID: 150, - Role: database.ChatMessageRoleAssistant, - }) - requireCheckpoint(200) - - // Tool messages do not end an assistant streaming turn. - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ - ID: 300, - Role: database.ChatMessageRoleTool, - }) - requireCheckpoint(200) - - // User messages do not end an assistant streaming turn either. - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ - ID: 400, - Role: database.ChatMessageRoleUser, - }) - requireCheckpoint(200) -} - -// TestSeedAssistantCheckpoint covers the three behaviors of -// seedAssistantCheckpoint: -// - success: a durable assistant message exists and its ID is -// installed as the checkpoint. -// - monotonic guard: an older ID does not move the checkpoint -// backwards (defends against concurrent advance from another -// publish path racing with the seed). -// - db error: a non sql.ErrNoRows failure must not change the -// checkpoint and must not panic. -func TestSeedAssistantCheckpoint(t *testing.T) { - t.Parallel() - - t.Run("InstallsLatestAssistantID", func(t *testing.T) { + t.Run("AssistantClaimsAllInProgressParts", func(t *testing.T) { t.Parallel() - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) server := &Server{ - db: db, logger: slogtest.Make(t, nil), clock: quartz.NewMock(t), } chatID := uuid.New() state := server.getOrCreateStreamState(chatID) + state.mu.Lock() + state.buffer = []bufferedStreamPart{ + makeCommittedPart(100, "old-1"), + makeInProgressPart("new-1"), + makeInProgressPart("new-2"), + } + state.mu.Unlock() - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), database.GetLastChatMessageByRoleParams{ - ChatID: chatID, - Role: database.ChatMessageRoleAssistant, - }).Return(database.ChatMessage{ - ID: 500, + server.claimCommittedParts(chatID, database.ChatMessage{ + ID: 200, Role: database.ChatMessageRoleAssistant, - }, nil) - - server.seedAssistantCheckpoint(ctx, chatID, state, server.logger) + }) state.mu.Lock() defer state.mu.Unlock() - require.Equal(t, int64(500), state.lastCommittedAssistantMessageID, - "seed must install the latest durable assistant message ID as the checkpoint") + require.Equal(t, int64(100), state.buffer[0].committedMessageID, + "already-claimed parts must keep their original message ID") + require.Equal(t, int64(200), state.buffer[1].committedMessageID, + "in-progress parts must be claimed by the new message ID") + require.Equal(t, int64(200), state.buffer[2].committedMessageID, + "in-progress parts must be claimed by the new message ID") }) - t.Run("DoesNotMoveCheckpointBackwards", func(t *testing.T) { + t.Run("ToolMessageIsNoOp", func(t *testing.T) { t.Parallel() - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) server := &Server{ - db: db, logger: slogtest.Make(t, nil), clock: quartz.NewMock(t), } chatID := uuid.New() state := server.getOrCreateStreamState(chatID) state.mu.Lock() - state.lastCommittedAssistantMessageID = 1000 - state.mu.Unlock() - - // DB reports an older assistant message ID. The monotonic - // guard must keep the existing higher checkpoint. - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), gomock.Any()).Return( - database.ChatMessage{ID: 500, Role: database.ChatMessageRoleAssistant}, - nil, - ) - - server.seedAssistantCheckpoint(ctx, chatID, state, server.logger) - - state.mu.Lock() - defer state.mu.Unlock() - require.Equal(t, int64(1000), state.lastCommittedAssistantMessageID, - "seed must not move the checkpoint backwards") - }) - - t.Run("DBErrorLeavesCheckpointUntouched", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - server := &Server{ - db: db, - logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), - clock: quartz.NewMock(t), + state.buffer = []bufferedStreamPart{ + makeInProgressPart("a"), + makeInProgressPart("b"), } - chatID := uuid.New() - state := server.getOrCreateStreamState(chatID) - state.mu.Lock() - state.lastCommittedAssistantMessageID = 42 state.mu.Unlock() - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), gomock.Any()).Return( - database.ChatMessage{}, xerrors.New("database explode"), - ) - - server.seedAssistantCheckpoint(ctx, chatID, state, server.logger) + server.claimCommittedParts(chatID, database.ChatMessage{ + ID: 300, + Role: database.ChatMessageRoleTool, + }) state.mu.Lock() defer state.mu.Unlock() - require.Equal(t, int64(42), state.lastCommittedAssistantMessageID, - "a non-ErrNoRows DB error must not change the checkpoint") + require.Equal(t, int64(0), state.buffer[0].committedMessageID, + "tool messages must not claim buffered parts") + require.Equal(t, int64(0), state.buffer[1].committedMessageID, + "tool messages must not claim buffered parts") }) - t.Run("NoRowsLeavesCheckpointAtZero", func(t *testing.T) { + t.Run("UserMessageIsNoOp", func(t *testing.T) { t.Parallel() - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) server := &Server{ - db: db, logger: slogtest.Make(t, nil), clock: quartz.NewMock(t), } chatID := uuid.New() state := server.getOrCreateStreamState(chatID) + state.mu.Lock() + state.buffer = []bufferedStreamPart{ + makeInProgressPart("a"), + } + state.mu.Unlock() - db.EXPECT().GetLastChatMessageByRole(gomock.Any(), gomock.Any()).Return( - database.ChatMessage{}, sql.ErrNoRows, - ) - - server.seedAssistantCheckpoint(ctx, chatID, state, server.logger) + server.claimCommittedParts(chatID, database.ChatMessage{ + ID: 400, + Role: database.ChatMessageRoleUser, + }) state.mu.Lock() defer state.mu.Unlock() - require.Equal(t, int64(0), state.lastCommittedAssistantMessageID, - "a fresh chat with no prior assistant messages must leave the checkpoint at zero") + require.Equal(t, int64(0), state.buffer[0].committedMessageID, + "user messages must not claim buffered parts") + }) + + t.Run("NoLiveStateIsNoOp", func(t *testing.T) { + t.Parallel() + + server := &Server{ + logger: slogtest.Make(t, nil), + clock: quartz.NewMock(t), + } + chatID := uuid.New() + + // No state stored: claimCommittedParts must not panic and + // must not allocate a new state for an unknown chat. + require.NotPanics(t, func() { + server.claimCommittedParts(chatID, database.ChatMessage{ + ID: 500, + Role: database.ChatMessageRoleAssistant, + }) + }) + _, ok := server.chatStreams.Load(chatID) + require.False(t, ok, + "claimCommittedParts must not create stream state for a chat that has none") }) } // TestSubscribeToStream_FiltersBufferedParts_Integration wires -// publishToStream, advanceAssistantCheckpoint, and subscribeToStream -// together to confirm the end-to-end contract: a subscriber with a -// known cursor only receives parts from turns the cursor does not -// already cover. +// publishToStream, claimCommittedParts (via publishMessage), and +// subscribeToStream together to confirm the end-to-end contract: a +// reconnecting subscriber only receives parts that belong to the +// current in-progress turn, not parts that were already committed +// to durable assistant messages. func TestSubscribeToStream_FiltersBufferedParts_Integration(t *testing.T) { t.Parallel() @@ -5769,18 +5638,18 @@ func TestSubscribeToStream_FiltersBufferedParts_Integration(t *testing.T) { } chatID := uuid.New() - // Start buffering, then simulate the lifecycle: - // 1. Stream parts of turn A (checkpoint = 0, no commit yet). - // 2. Commit turn A's durable message with ID 100. - // 3. Stream parts of turn B (checkpoint now = 100). - // 4. Commit turn B's durable message with ID 200. - // 5. Stream parts of turn C (checkpoint now = 200). + // Simulate the lifecycle: + // 1. Stream parts of turn A (still in-progress, no commit yet). + // 2. Commit turn A; its parts are claimed by message 100. + // 3. Stream parts of turn B (in-progress). + // 4. Commit turn B; its parts are claimed by message 200. + // 5. Stream parts of turn C (in-progress, never committed). state := server.getOrCreateStreamState(chatID) state.mu.Lock() state.buffering = true state.mu.Unlock() - publish := func(text string) { + publishPart := func(text string) { server.publishToStream(chatID, codersdk.ChatStreamEvent{ Type: codersdk.ChatStreamEventTypeMessagePart, MessagePart: &codersdk.ChatStreamMessagePart{ @@ -5790,52 +5659,31 @@ func TestSubscribeToStream_FiltersBufferedParts_Integration(t *testing.T) { }) } - publish("A-1") - publish("A-2") - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ + publishPart("A-1") + publishPart("A-2") + server.claimCommittedParts(chatID, database.ChatMessage{ ID: 100, Role: database.ChatMessageRoleAssistant, }) - publish("B-1") - publish("B-2") - server.advanceAssistantCheckpoint(chatID, database.ChatMessage{ + publishPart("B-1") + publishPart("B-2") + server.claimCommittedParts(chatID, database.ChatMessage{ ID: 200, Role: database.ChatMessageRoleAssistant, }) - publish("C-1") + publishPart("C-1") - // Subscriber that already has turn A (cursor = 100) should - // receive only turn B and turn C parts. - snapshot, _, _, cancel := server.subscribeToStream(chatID, 100) + // Reconnecting subscriber: only the currently in-progress turn + // (turn C) survives the filter, no matter what cursor the + // client passes through SubscribeAuthorized (the filter no + // longer depends on the cursor). + snapshot, _, _, cancel := server.subscribeToStream(chatID) defer cancel() texts := make([]string, 0, len(snapshot)) for _, ev := range snapshot { texts = append(texts, partText(ev)) } - require.Equal(t, []string{"B-1", "B-2", "C-1"}, texts, - "subscriber past turn A must not receive turn A parts") - - // Subscriber that already has both A and B (cursor = 200) - // should receive only turn C parts. - snapshot2, _, _, cancel2 := server.subscribeToStream(chatID, 200) - defer cancel2() - texts2 := make([]string, 0, len(snapshot2)) - for _, ev := range snapshot2 { - texts2 = append(texts2, partText(ev)) - } - require.Equal(t, []string{"C-1"}, texts2, - "subscriber past turn B must not receive turn A or B parts") - - // Fresh subscriber (cursor = 0) receives the entire buffer. - snapshot3, _, _, cancel3 := server.subscribeToStream(chatID, 0) - defer cancel3() - require.Len(t, snapshot3, 5, - "fresh subscriber must receive every buffered part") - - // Relay subscriber (sentinel) receives the entire buffer. - snapshot4, _, _, cancel4 := server.subscribeToStream(chatID, math.MaxInt64) - defer cancel4() - require.Len(t, snapshot4, 5, - "relay sentinel must receive every buffered part") + require.Equal(t, []string{"C-1"}, texts, + "only in-progress (un-claimed) buffered parts must survive the filter") } diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index 9d1512707a..6ec2e33920 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -9560,6 +9560,48 @@ func TestAdvisorHappyPath_RootChat(t *testing.T) { }) require.NoError(t, err) + // Subscribe before the worker commits any durable messages so we + // observe the advisor tool-result deltas live. Buffered parts are + // claimed by their committed durable message ID at publishMessage + // time and dropped from snapshots of late-connecting subscribers, so + // a post-completion Subscribe() would no longer see streaming + // deltas. Collecting events from the live channel covers the + // streaming UX contract this test exists to verify. + _, liveEvents, cancelLive, ok := server.Subscribe(ctx, chat.ID, nil, 0) + require.True(t, ok) + var ( + livePartsMu sync.Mutex + liveAdvisorDeltas []string + liveCollectorDone = make(chan struct{}) + ) + go func() { + defer close(liveCollectorDone) + for { + select { + case <-ctx.Done(): + return + case event, eventsOK := <-liveEvents: + if !eventsOK { + return + } + if event.Type != codersdk.ChatStreamEventTypeMessagePart || + event.MessagePart == nil { + continue + } + part := event.MessagePart.Part + if event.MessagePart.Role != codersdk.ChatMessageRoleTool || + part.Type != codersdk.ChatMessagePartTypeToolResult || + part.ToolName != chatadvisor.ToolName || + part.ResultDelta == "" { + continue + } + livePartsMu.Lock() + liveAdvisorDeltas = append(liveAdvisorDeltas, part.ResultDelta) + livePartsMu.Unlock() + } + } + }() + require.Eventually(t, func() bool { got, getErr := db.GetChatByID(ctx, chat.ID) if getErr != nil { @@ -9614,24 +9656,16 @@ func TestAdvisorHappyPath_RootChat(t *testing.T) { require.True(t, parentSawAdvisorResult, "parent must see the advisor reply in its continuation call") - snapshot, _, cancelStream, ok := server.Subscribe(ctx, chat.ID, nil, 0) - require.True(t, ok) - cancelStream() - - var streamedAdvisorDeltas []string - for _, event := range snapshot { - if event.Type != codersdk.ChatStreamEventTypeMessagePart || event.MessagePart == nil { - continue - } - part := event.MessagePart.Part - if event.MessagePart.Role == codersdk.ChatMessageRoleTool && - part.Type == codersdk.ChatMessagePartTypeToolResult && - part.ToolName == chatadvisor.ToolName && - part.ResultDelta != "" { - streamedAdvisorDeltas = append(streamedAdvisorDeltas, part.ResultDelta) - } - } - require.Equal(t, advisorDeltas, streamedAdvisorDeltas, + // Stop the live collector and assert it captured the streaming + // advisor deltas during processing. Late subscribers no longer + // see committed parts because publishMessage claims them out of + // new snapshots, so the assertion must use the live collector. + cancelLive() + <-liveCollectorDone + livePartsMu.Lock() + collectedAdvisorDeltas := append([]string(nil), liveAdvisorDeltas...) + livePartsMu.Unlock() + require.Equal(t, advisorDeltas, collectedAdvisorDeltas, "advisor nested text deltas must stream into the parent tool card") persisted, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ diff --git a/enterprise/coderd/x/chatd/chatd_test.go b/enterprise/coderd/x/chatd/chatd_test.go index 6d66cc9179..97b63cba03 100644 --- a/enterprise/coderd/x/chatd/chatd_test.go +++ b/enterprise/coderd/x/chatd/chatd_test.go @@ -1257,11 +1257,9 @@ func TestSubscribeRelayMultipleReconnects(t *testing.T) { 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. +// TestSubscribeRelayDialCanceledOnFastCompletion verifies that a +// subscriber on a remote replica still sees the committed assistant +// response when the worker completes faster than the relay dial. // // Scenario: // 1. Subscriber subscribes to a chat while it's in waiting state (no relay). @@ -1269,12 +1267,15 @@ func TestSubscribeRelayMultipleReconnects(t *testing.T) { // 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. +// 6. Even though the relay never delivered streaming parts, the +// committed assistant message arrives via pubsub so the user +// does not need to refresh to see the response. // -// 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. +// Streaming parts for committed turns are intentionally NOT replayed +// via the relay: they would duplicate the durable message on the +// user's screen. The buffer retains in-progress parts only; once an +// assistant turn commits, the parts that built it are claimed by +// the durable message ID and dropped from new buffer snapshots. func TestSubscribeRelayDialCanceledOnFastCompletion(t *testing.T) { t.Parallel() @@ -1336,8 +1337,10 @@ func TestSubscribeRelayDialCanceledOnFastCompletion(t *testing.T) { 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. + // grace period after processing, so the relay session + // can complete (control events, status updates) even + // though every part has been claimed by its durable + // message and the snapshot is empty. snapshot, relayEvents, cancel, ok := worker.Subscribe(ctx, chatID, requestHeader, math.MaxInt64) if !ok { return nil, nil, nil, xerrors.New("worker subscribe failed") @@ -1381,27 +1384,22 @@ func TestSubscribeRelayDialCanceledOnFastCompletion(t *testing.T) { // Release the relay dial now that the worker is done. close(workerDone) - // Collect all events that arrived at the subscriber. - var messageParts []string + // Collect events that arrived at the subscriber. The committed + // assistant message is guaranteed to arrive via pubsub even when + // the relay dial races worker completion; streaming parts are + // best-effort and are not asserted here because the buffer drops + // already-committed parts to prevent duplicate UI rendering. 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++ - } + if event.Type == codersdk.ChatStreamEventTypeMessage && + event.Message != nil && + event.Message.Role == codersdk.ChatMessageRoleAssistant { + committedAssistantMsgs++ } - return committedAssistantMsgs > 0 && len(messageParts) > 0 + return committedAssistantMsgs > 0 default: return false } @@ -1415,221 +1413,6 @@ func TestSubscribeRelayDialCanceledOnFastCompletion(t *testing.T) { // 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") -} - -// TestSubscribeRelayDrainWithinGraceLeavesBufferRetained characterizes -// the multi-replica trigger for the retained-buffer leak: an enterprise -// relay drain (relayDrainTimeout = 200ms) always fires inside the -// worker's 5s grace window, so the worker-side subscriber-detach hits -// cleanupStreamIfIdle's early-return and the buffer stays mapped. -// streamJanitorLoop is the timer-driven backstop. -// -// The assertion is behavioral (a fresh worker.Subscribe sees the -// retained message_parts) rather than a chatStreams-size check because -// _test.go identifiers in coderd/x/chatd do not link into the -// enterprise test binary, and adding a production accessor for this -// isn't justified. The matching reap assertion lives in the OSS unit -// tests in coderd/x/chatd/chatd_internal_test.go. -func TestSubscribeRelayDrainWithinGraceLeavesBufferRetained(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - workerID := uuid.New() - subscriberID := uuid.New() - ctx := testutil.Context(t, testutil.WaitLong) - - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("relay-drain-characterization") - } - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("hello ", "from ", "worker")..., - ) - }) - - workerLogger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) - // Freeze the worker's clock so streamJanitorLoop cannot race the - // buffer-retained assertion on slow CI. - workerClock := quartz.NewMock(t) - trapAcquire := workerClock.Trap().NewTicker("chatd", "acquire") - defer trapAcquire.Close() - worker := osschatd.New(osschatd.Config{ - Logger: workerLogger, - Database: db, - ReplicaID: workerID, - Pubsub: ps, - PendingChatAcquireInterval: time.Millisecond, - InFlightChatStaleAfter: testutil.WaitSuperLong, - Clock: workerClock, - }) - worker.Start() - trapAcquire.MustWait(ctx).MustRelease(ctx) - t.Cleanup(func() { - require.NoError(t, worker.Close()) - }) - - // Use a mock clock for the subscriber so the relay drain - // timer never fires until we explicitly advance it. This - // removes the nondeterministic 200ms race between the drain - // timer and the multi-hop snapshot forwarding pipeline. - subscriberClock := quartz.NewMock(t) - trapDrain := subscriberClock.Trap().NewTimer("drain") - defer trapDrain.Close() - - // Subscriber dials through to the worker. On cancel the relay - // drain fires well inside the worker's 5s grace, exercising the - // cleanupStreamIfIdle early-return path. - 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, - ) { - 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 - }, subscriberClock) - - user, org, model := seedChatDependencies(t, db) - setOpenAIProviderBaseURL(ctx, t, db, openAIURL) - - chat := seedWaitingChat(t, db, org.ID, user, model, "relay-drain-characterization") - - // Seed the pending turn directly instead of using SendMessage. - // SendMessage publishes a pending control notification that is - // irrelevant to this relay-retention case. Under CI that - // notification can arrive after processChat arms its control - // subscription and interrupt the worker before it emits parts. - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Content: pqtype.NullRawMessage{ - RawMessage: json.RawMessage(`[{"type":"text","text":"hello"}]`), - Valid: true, - }, - }) - _, err := db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{ - ID: chat.ID, - Status: database.ChatStatusPending, - }) - require.NoError(t, err) - - // Attach before processing so the relay opens as soon as - // status=running arrives. - _, events, subCancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0) - require.True(t, ok) - - // Wake the worker with the acquire ticker. This keeps the - // setup free of pending control notifications while still - // exercising the normal processing loop. - workerClock.Advance(time.Millisecond).MustWait(ctx) - - // Drain events until processing has clearly completed: we need - // the assistant message and at least one message_part so we know - // processChat's defer has flipped buffering=false and populated - // bufferRetainedAt before the subscriber detaches. - // - // Each Eventually gets its own context so one slow assertion - // cannot starve subsequent ones of their deadline. - var committedAssistantMsgs int - var messagePartsSeen int - evCtx1 := testutil.Context(t, testutil.WaitLong) - testutil.Eventually(evCtx1, t, func(context.Context) bool { - select { - case event := <-events: - switch event.Type { - case codersdk.ChatStreamEventTypeMessagePart: - messagePartsSeen++ - case codersdk.ChatStreamEventTypeMessage: - if event.Message != nil && event.Message.Role == codersdk.ChatMessageRoleAssistant { - committedAssistantMsgs++ - } - } - return committedAssistantMsgs > 0 && messagePartsSeen > 0 - default: - return false - } - }, testutil.IntervalFast) - - // Drain all NewTimer("drain") calls in a background goroutine. - // The merge loop may create one or two drain timers depending - // on the relative ordering of the status=WAITING pubsub - // notification and the async relay dial completion. Each - // trapped call must be released so the production goroutine - // is unblocked, and the clock must be advanced past the - // 200ms drain timeout to fire the timer. - var drainsFired atomic.Int32 - go func() { - for { - call, err := trapDrain.Wait(ctx) - if err != nil { - return - } - if err := call.Release(ctx); err != nil { - return - } - subscriberClock.Advance(200 * time.Millisecond) - drainsFired.Add(1) - } - }() - - // Wait for DB status=waiting AND at least one drain timer to - // have fired. Checking drainsFired proves the relay was torn - // down by the drain path, not by context cancellation. - evCtx2 := testutil.Context(t, testutil.WaitLong) - testutil.Eventually(evCtx2, t, func(ctx context.Context) bool { - if drainsFired.Load() == 0 { - return false - } - fromDB, dbErr := db.GetChatByID(ctx, chat.ID) - if dbErr != nil { - return false - } - return fromDB.Status == database.ChatStatusWaiting - }, testutil.IntervalFast) - - // Tear the subscriber down inside the worker's grace window. - subCancel() - - // A fresh worker.Subscribe still sees the retained - // message_parts: the buffer was not reaped when the relay - // drained. Eventually absorbs the short window before the - // worker observes the teardown. The retry itself re-enters - // cleanupStreamIfIdle via its own cancel defer but still - // early-returns because grace is still open. - evCtx3 := testutil.Context(t, testutil.WaitLong) - testutil.Eventually(evCtx3, t, func(ctx context.Context) bool { - snap, _, snapCancel, ok := worker.Subscribe(ctx, chat.ID, nil, math.MaxInt64) - if !ok { - return false - } - defer snapCancel() - for _, e := range snap { - if e.Type == codersdk.ChatStreamEventTypeMessagePart { - return true - } - } - return false - }, testutil.IntervalFast, - "retained buffer must still contain message_parts after the "+ - "relay drains within grace") } // TestSubscribeRelayEstablishedMidStream demonstrates that when the