From 782b7166a4ba4ce68259f95b3da63478e6f311ef Mon Sep 17 00:00:00 2001 From: Mathias Fredriksson Date: Wed, 29 Apr 2026 14:08:35 +0300 Subject: [PATCH] fix: preserve stream state on interrupt, fix auto-promote error handling (#24314) When tryAutoPromoteQueuedMessage's insert fails, return the error instead of swallowing it so the transaction rolls back and the queued message survives. Previously the POP DELETE committed while the INSERT silently failed, permanently losing the message. Remove clearStreamState() from the pending/waiting status handler in the frontend. The durable message event clears stream state via the existing needsStreamReset path, eliminating the visual gap where content vanishes before the persisted message arrives. Fixes CODAGT-61 --- coderd/x/chatd/chatd.go | 7 +- coderd/x/chatd/chatd_internal_test.go | 225 +++++++++++++++++- .../ChatConversation/chatStore.test.tsx | 191 ++++++++++++++- .../ChatConversation/useChatStore.ts | 7 +- 4 files changed, 419 insertions(+), 11 deletions(-) diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 6483b18559..7393722e6d 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -5040,9 +5040,7 @@ func (p *Server) tryAutoPromoteQueuedMessage( ).withCreatedBy(chat.OwnerID)) msgs, err := insertChatMessageWithStore(ctx, tx, msgParams) if err != nil { - logger.Error(ctx, "failed to promote queued message", - slog.F("queued_message_id", nextQueued.ID), slog.Error(err)) - return nil, nil, false, nil + return nil, nil, false, xerrors.Errorf("insert promoted message: %w", err) } msg := msgs[0] @@ -5148,7 +5146,8 @@ func (p *Server) finishActiveChat( var promoteErr error result.promotedMessage, result.remainingQueuedMessages, result.shouldPublishQueueUpdate, promoteErr = p.tryAutoPromoteQueuedMessage(ctx, tx, latestChat) if promoteErr != nil { - logger.Error(ctx, "failed to auto-promote queued message", slog.Error(promoteErr)) + logger.Error(ctx, "auto-promote queued message failed, rolling back", slog.Error(promoteErr)) + return xerrors.Errorf("auto-promote queued message: %w", promoteErr) } else if result.promotedMessage != nil { status = database.ChatStatusPending } diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 0ad47577d2..6cc6956ab5 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -3524,9 +3524,9 @@ func TestProcessChat_IgnoresStaleControlNotification(t *testing.T) { clock: clock, workerID: workerID, chatHeartbeatInterval: time.Minute, + metrics: chatloop.NopMetrics(), configCache: newChatConfigCache(ctx, db, clock), heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry), - metrics: chatloop.NopMetrics(), } // Publish a stale "pending" notification on the control channel @@ -3680,6 +3680,7 @@ func TestHeartbeatTick_StolenChatIsInterrupted(t *testing.T) { clock: clock, workerID: workerID, chatHeartbeatInterval: time.Minute, + metrics: chatloop.NopMetrics(), heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry), } @@ -3760,6 +3761,7 @@ func TestHeartbeatTick_DBErrorDoesNotInterruptChats(t *testing.T) { clock: clock, workerID: uuid.New(), chatHeartbeatInterval: time.Minute, + metrics: chatloop.NopMetrics(), heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry), } @@ -4717,3 +4719,224 @@ func TestGetWorkspaceConn_DialErrorNotMisclassifiedAsTimeout(t *testing.T) { // The original dial error should propagate. require.ErrorContains(t, err, "authentication failed") } + +// TestAutoPromote_InsertFailureRollsBackTransaction verifies that when +// tryAutoPromoteQueuedMessage pops a queued message but the subsequent +// insert fails, the error propagates to the InTx callback, causing the +// transaction to roll back and preserving the queued message. +func TestAutoPromote_InsertFailureRollsBackTransaction(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + tx := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + ps := dbpubsub.NewInMemory() + clock := quartz.NewReal() + + chatID := uuid.New() + workerID := uuid.New() + ownerID := uuid.New() + modelConfigID := uuid.New() + + waitingChat := database.Chat{ + ID: chatID, + OwnerID: ownerID, + LastModelConfigID: modelConfigID, + Status: database.ChatStatusWaiting, + WorkerID: uuid.NullUUID{UUID: workerID, Valid: true}, + } + queuedMsg := database.ChatQueuedMessage{ + ID: 1, + ChatID: chatID, + Content: []byte(`[{"type":"text","text":"queued"}]`), + } + insertErr := xerrors.New("insert failed") + + server := &Server{ + db: db, + logger: logger, + pubsub: ps, + configCache: newChatConfigCache(ctx, db, clock), + } + + // The caller runs tryAutoPromoteQueuedMessage inside InTx. + // Wire the mock to execute the callback against the TX mock. + var txErr error + db.EXPECT().InTx(gomock.Any(), gomock.Any()).DoAndReturn( + func(fn func(database.Store) error, _ *database.TxOptions) error { + txErr = fn(tx) + return txErr + }, + ) + + // Inside the TX: lock chat, get queued messages, resolve model + // config, pop queued message, insert fails. + tx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(waitingChat, nil) + tx.EXPECT().GetChatQueuedMessages(gomock.Any(), chatID).Return([]database.ChatQueuedMessage{queuedMsg}, nil) + tx.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(database.ChatModelConfig{ID: modelConfigID}, nil) + tx.EXPECT().PopNextQueuedMessage(gomock.Any(), chatID).Return(queuedMsg, nil) + tx.EXPECT().InsertChatMessages(gomock.Any(), gomock.Any()).Return(nil, insertErr) + + // Invoke tryAutoPromoteQueuedMessage through the same InTx + // pattern the processChat defer uses. The test directly calls + // the production path to verify error propagation. + _ = db.InTx(func(txStore database.Store) error { + latestChat, err := txStore.GetChatByIDForUpdate(ctx, chatID) + if err != nil { + return err + } + + _, _, _, promoteErr := server.tryAutoPromoteQueuedMessage(ctx, txStore, latestChat) + if promoteErr != nil { + return promoteErr + } + + // This code path should not be reached when the insert + // fails, because promoteErr should be non-nil. + return nil + }, nil) + + // The InTx callback must return a non-nil error so the + // transaction rolls back, preserving the queued message. + require.Error(t, txErr, "InTx callback should return error when insert fails") +} + +// TestAutoPromote_WakesRunLoopAfterPromotion verifies that after the +func TestAutoPromote_InsertFailureSkipsStatusUpdate(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + tx := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + ps := dbpubsub.NewInMemory() + clock := quartz.NewReal() + + chatID := uuid.New() + workerID := uuid.New() + ownerID := uuid.New() + modelConfigID := uuid.New() + + waitingChat := database.Chat{ + ID: chatID, + OwnerID: ownerID, + LastModelConfigID: modelConfigID, + Status: database.ChatStatusWaiting, + WorkerID: uuid.NullUUID{UUID: workerID, Valid: true}, + } + queuedMsg := database.ChatQueuedMessage{ + ID: 1, + ChatID: chatID, + Content: []byte(`[{"type":"text","text":"queued"}]`), + } + + wakeCh := make(chan struct{}, 1) + server := &Server{ + db: db, + logger: logger, + pubsub: ps, + clock: clock, + workerID: workerID, + wakeCh: wakeCh, + chatHeartbeatInterval: time.Minute, + metrics: chatloop.NopMetrics(), + configCache: newChatConfigCache(ctx, db, clock), + heartbeatRegistry: make(map[uuid.UUID]*heartbeatEntry), + } + + // Block model resolution until the control subscriber fires. + modelBlocked := make(chan struct{}) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), gomock.Any()).DoAndReturn( + func(ctx context.Context, _ uuid.UUID) (database.ChatModelConfig, error) { + <-modelBlocked + return database.ChatModelConfig{}, xerrors.New("no model") + }, + ).AnyTimes() + db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return(nil, nil).AnyTimes() + db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return(nil, nil).AnyTimes() + db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return( + database.ChatUsageLimitConfig{}, sql.ErrNoRows, + ).AnyTimes() + db.EXPECT().GetChatMessagesForPromptByChatID(gomock.Any(), chatID).Return(nil, nil).AnyTimes() + + // The deferred cleanup transaction: InsertChatMessages fails, + // so UpdateChatStatus must NOT be called. + db.EXPECT().InTx(gomock.Any(), gomock.Any()).DoAndReturn( + func(fn func(database.Store) error, _ *database.TxOptions) error { + return fn(tx) + }, + ) + tx.EXPECT().GetChatByIDForUpdate(gomock.Any(), chatID).Return(waitingChat, nil) + tx.EXPECT().GetChatQueuedMessages(gomock.Any(), chatID).Return([]database.ChatQueuedMessage{queuedMsg}, nil) + tx.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(database.ChatModelConfig{ID: modelConfigID}, nil) + tx.EXPECT().PopNextQueuedMessage(gomock.Any(), chatID).Return(queuedMsg, nil) + tx.EXPECT().InsertChatMessages(gomock.Any(), gomock.Any()).Return( + nil, xerrors.New("insert failed"), + ) + tx.EXPECT().UpdateChatStatus(gomock.Any(), gomock.Any()).Times(0) + + // Subscribe BEFORE launching the goroutine. + runningCh := make(chan struct{}, 1) + unsubRunning, err := ps.SubscribeWithErr( + coderdpubsub.ChatStreamNotifyChannel(chatID), + func(_ context.Context, msg []byte, err error) { + if err != nil { + return + } + var notify coderdpubsub.ChatStreamNotifyMessage + if json.Unmarshal(msg, ¬ify) != nil { + return + } + if notify.Status == string(database.ChatStatusRunning) { + select { + case runningCh <- struct{}{}: + default: + } + } + }, + ) + require.NoError(t, err) + defer unsubRunning() + + chat := database.Chat{ID: chatID, OwnerID: ownerID, LastModelConfigID: modelConfigID} + processDone := make(chan struct{}) + go func() { + defer close(processDone) + server.processChat(ctx, chat) + }() + + select { + case <-runningCh: + case <-ctx.Done(): + t.Fatal("timed out waiting for running status") + } + + // Publish an interrupt so processChat exits runChat. + interruptMsg, err := json.Marshal(coderdpubsub.ChatStreamNotifyMessage{ + Status: string(database.ChatStatusWaiting), + }) + require.NoError(t, err) + err = ps.Publish(coderdpubsub.ChatStreamNotifyChannel(chatID), interruptMsg) + require.NoError(t, err) + + // Unblock model resolution so runChat can exit. + close(modelBlocked) + + select { + case <-processDone: + case <-ctx.Done(): + t.Fatal("processChat did not complete") + } + + // The wake channel should NOT have a signal because the + // transaction failed before reaching UpdateChatStatus. + select { + case <-wakeCh: + t.Fatal("wake channel should not have a signal after insert failure") + default: + // No signal, as expected. + } +} diff --git a/site/src/pages/AgentsPage/components/ChatConversation/chatStore.test.tsx b/site/src/pages/AgentsPage/components/ChatConversation/chatStore.test.tsx index 73d2c3224a..9f9de11388 100644 --- a/site/src/pages/AgentsPage/components/ChatConversation/chatStore.test.tsx +++ b/site/src/pages/AgentsPage/components/ChatConversation/chatStore.test.tsx @@ -731,7 +731,14 @@ describe("useChatStore", () => { }); await waitFor(() => { - expect(result.current.streamState).toBeNull(); + // Stream state is preserved after status=pending (the + // durable message event handles cleanup via + // needsStreamReset). Only new message_parts should be + // blocked by the shouldApplyMessagePart gate. + expect(result.current.streamState).not.toBeNull(); + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "first" }, + ]); }); act(() => { @@ -749,7 +756,12 @@ describe("useChatStore", () => { }); await waitFor(() => { - expect(result.current.streamState).toBeNull(); + // The late message_part should not be applied because + // shouldApplyMessagePart gates on pending/waiting. + // Stream state still shows the original "first". + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "first" }, + ]); }); }); @@ -3036,6 +3048,181 @@ describe("useChatStore", () => { expect(result.current.chatStatus).toBe("running"); }); }); + + it("preserves stream state when status transitions to waiting", async () => { + immediateAnimationFrame(); + + const chatID = "chat-preserve-stream"; + const existingMessage = makeMessage(chatID, 1, "user", "hello"); + const mockSocket = createMockSocket(); + mockWatchChatReturn(mockSocket); + + const queryClient = createTestQueryClient(); + const wrapper = ({ children }: PropsWithChildren) => ( + {children} + ); + const setChatErrorReason = vi.fn(); + const clearChatErrorReason = vi.fn(); + + const { result } = renderHook( + () => { + const { store } = useChatStore({ + chatID, + chatMessages: [existingMessage], + chatRecord: makeChat(chatID), + chatMessagesData: { + messages: [existingMessage], + queued_messages: [], + has_more: false, + }, + chatQueuedMessages: [], + setChatErrorReason, + clearChatErrorReason, + }); + return { + streamState: useChatSelector(store, selectStreamState), + }; + }, + { wrapper }, + ); + + await waitFor(() => { + expect(watchChat).toHaveBeenCalledWith(chatID, 1); + }); + + // Build up stream state with a message_part. + act(() => { + mockSocket.emitData({ + type: "message_part", + chat_id: chatID, + message_part: { + role: "assistant", + part: { type: "text", text: "thinking..." }, + }, + }); + }); + + await waitFor(() => { + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "thinking..." }, + ]); + }); + + // Deliver a status=waiting event (interrupt). Stream state + // should be preserved so the user continues to see the + // partial response until the durable message arrives. + act(() => { + mockSocket.emitData({ + type: "status", + chat_id: chatID, + status: { status: "waiting" }, + }); + }); + + await waitFor(() => { + expect(result.current.streamState).not.toBeNull(); + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "thinking..." }, + ]); + }); + }); + + it("clears stream state when durable message follows waiting status", async () => { + immediateAnimationFrame(); + + const chatID = "chat-durable-clears"; + const existingMessage = makeMessage(chatID, 1, "user", "hello"); + const mockSocket = createMockSocket(); + mockWatchChatReturn(mockSocket); + + const queryClient = createTestQueryClient(); + const wrapper = ({ children }: PropsWithChildren) => ( + {children} + ); + const setChatErrorReason = vi.fn(); + const clearChatErrorReason = vi.fn(); + + const { result } = renderHook( + () => { + const { store } = useChatStore({ + chatID, + chatMessages: [existingMessage], + chatRecord: makeChat(chatID), + chatMessagesData: { + messages: [existingMessage], + queued_messages: [], + has_more: false, + }, + chatQueuedMessages: [], + setChatErrorReason, + clearChatErrorReason, + }); + return { + streamState: useChatSelector(store, selectStreamState), + orderedIDs: useChatSelector(store, selectOrderedMessageIDs), + }; + }, + { wrapper }, + ); + + await waitFor(() => { + expect(watchChat).toHaveBeenCalledWith(chatID, 1); + }); + + // Build up stream state. + act(() => { + mockSocket.emitData({ + type: "message_part", + chat_id: chatID, + message_part: { + role: "assistant", + part: { type: "text", text: "partial response" }, + }, + }); + }); + + await waitFor(() => { + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "partial response" }, + ]); + }); + + // Deliver status=waiting (interrupt). Stream state should be + // preserved so the user continues to see the partial response + // until the durable message arrives. + act(() => { + mockSocket.emitData({ + type: "status", + chat_id: chatID, + status: { status: "waiting" }, + }); + }); + + // Stream state must still be present after the status change. + await waitFor(() => { + expect(result.current.streamState).not.toBeNull(); + expect(result.current.streamState?.blocks).toEqual([ + { type: "response", text: "partial response" }, + ]); + }); + + // Now deliver the durable assistant message. This should + // clear stream state via the needsStreamReset path. + act(() => { + mockSocket.emitData({ + type: "message", + chat_id: chatID, + message: makeMessage(chatID, 2, "assistant", "partial response"), + }); + }); + + // Stream state should now be null and the durable message + // should be in the message store. + await waitFor(() => { + expect(result.current.streamState).toBeNull(); + expect(result.current.orderedIDs).toContain(2); + }); + }); }); describe("thinking indicator event ordering", () => { diff --git a/site/src/pages/AgentsPage/components/ChatConversation/useChatStore.ts b/site/src/pages/AgentsPage/components/ChatConversation/useChatStore.ts index 5d2cda47eb..0e6e18a917 100644 --- a/site/src/pages/AgentsPage/components/ChatConversation/useChatStore.ts +++ b/site/src/pages/AgentsPage/components/ChatConversation/useChatStore.ts @@ -390,9 +390,9 @@ export const useChatStore = ( }; // Discard buffered parts without applying them. Used when - // stream state is about to be cleared (pending, waiting, - // retry) — flushing would re-populate the state that the - // event is about to clear. + // the stream is no longer active (pending, waiting, retry) + // so stale buffered parts are not applied after the + // status transition. const discardBufferedParts = () => { partsBuf.length = 0; if (partsFlushTimer !== null) { @@ -508,7 +508,6 @@ export const useChatStore = ( store.setChatStatus(nextStatus); if (nextStatus === "pending" || nextStatus === "waiting") { discardBufferedParts(); - store.clearStreamState(); store.clearRetryState(); } if (nextStatus === "running") {