From 2cd871e88f77b8c213f7d6ba0d6c9256ab3215e1 Mon Sep 17 00:00:00 2001 From: Kyle Carberry Date: Fri, 6 Mar 2026 12:47:28 -0800 Subject: [PATCH] fix: resolve bugs in chat frontend ChatContext and streamState (#22721) Split from #22693 per review feedback. Fixes race conditions, TOCTOU bugs, and state management issues in the chat frontend streaming layer. --- .../AgentsPage/AgentDetail/ChatContext.ts | 146 ++++++++++++------ .../AgentDetail/streamState.test.ts | 65 +++++++- .../AgentsPage/AgentDetail/streamState.ts | 14 +- 3 files changed, 174 insertions(+), 51 deletions(-) diff --git a/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts b/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts index e687a77f88..2b1d20d595 100644 --- a/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts +++ b/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts @@ -217,6 +217,7 @@ export const createChatStore = (): ChatStore => { const nextMessagesByID = buildMessageMap(safeMessages); const nextOrderedMessageIDs = buildOrderedMessageIDs(safeMessages); + // Fast-path: skip setState entirely when nothing changed. if ( mapsEqualByRef(state.messagesByID, nextMessagesByID) && arraysEqual(state.orderedMessageIDs, nextOrderedMessageIDs) @@ -224,35 +225,62 @@ export const createChatStore = (): ChatStore => { return; } - setState((current) => ({ - ...current, - messagesByID: nextMessagesByID, - orderedMessageIDs: nextOrderedMessageIDs, - })); + setState((current) => { + // Re-check equality against `current` inside the updater + // to avoid overwriting a concurrent state change. + if ( + mapsEqualByRef(current.messagesByID, nextMessagesByID) && + arraysEqual(current.orderedMessageIDs, nextOrderedMessageIDs) + ) { + return current; + } + return { + ...current, + messagesByID: nextMessagesByID, + orderedMessageIDs: nextOrderedMessageIDs, + }; + }); }; const upsertDurableMessage = (message: TypesGen.ChatMessage) => { + // Use `state` for the early-return guard so we can return + // the result synchronously. The actual mutation below uses + // `current` inside the updater to avoid overwriting a + // concurrent state change (TOCTOU). const existing = state.messagesByID.get(message.id); const isDuplicate = state.messagesByID.has(message.id); if (existing && chatMessagesEqualByValue(existing, message)) { return { isDuplicate, changed: false }; } - const nextMessagesByID = new Map(state.messagesByID); - nextMessagesByID.set(message.id, message); + let actuallyChanged = false; + setState((current) => { + // Re-check inside the updater: another call may have + // already applied this exact message. + const curExisting = current.messagesByID.get(message.id); + if (curExisting && chatMessagesEqualByValue(curExisting, message)) { + return current; + } - const needsReorder = - !isDuplicate || nextMessagesByID.size !== state.messagesByID.size; - const nextOrderedMessageIDs = needsReorder - ? buildOrderedMessageIDs(Array.from(nextMessagesByID.values())) - : state.orderedMessageIDs; + actuallyChanged = true; - setState((current) => ({ - ...current, - messagesByID: nextMessagesByID, - orderedMessageIDs: nextOrderedMessageIDs, - })); - return { isDuplicate, changed: true }; + const nextMessagesByID = new Map(current.messagesByID); + nextMessagesByID.set(message.id, message); + + const curIsDuplicate = current.messagesByID.has(message.id); + const needsReorder = + !curIsDuplicate || nextMessagesByID.size !== current.messagesByID.size; + const nextOrderedMessageIDs = needsReorder + ? buildOrderedMessageIDs(Array.from(nextMessagesByID.values())) + : current.orderedMessageIDs; + + return { + ...current, + messagesByID: nextMessagesByID, + orderedMessageIDs: nextOrderedMessageIDs, + }; + }); + return { isDuplicate, changed: actuallyChanged }; }; const applyMessageParts = (parts: readonly Record[]) => { @@ -260,17 +288,19 @@ export const createChatStore = (): ChatStore => { return; } - let nextStreamState: StreamState | null = state.streamState; - for (const part of parts) { - nextStreamState = applyMessagePartToStreamState(nextStreamState, part); - } - if (nextStreamState === state.streamState) { - return; - } - setState((current) => ({ - ...current, - streamState: nextStreamState, - })); + setState((current) => { + let nextStreamState: StreamState | null = current.streamState; + for (const part of parts) { + nextStreamState = applyMessagePartToStreamState(nextStreamState, part); + } + if (nextStreamState === current.streamState) { + return current; + } + return { + ...current, + streamState: nextStreamState, + }; + }); }; return { @@ -287,15 +317,17 @@ export const createChatStore = (): ChatStore => { applyMessageParts, setQueuedMessages: (queuedMessages) => { const nextQueuedMessages = queuedMessages ?? []; - if ( - chatQueuedMessagesEqualByID(state.queuedMessages, nextQueuedMessages) - ) { - return; - } - setState((current) => ({ - ...current, - queuedMessages: nextQueuedMessages, - })); + setState((current) => { + if ( + chatQueuedMessagesEqualByID( + current.queuedMessages, + nextQueuedMessages, + ) + ) { + return current; + } + return { ...current, queuedMessages: nextQueuedMessages }; + }); }, setChatStatus: (status) => { if (state.chatStatus === status) { @@ -355,12 +387,14 @@ export const createChatStore = (): ChatStore => { if (state.subagentStatusOverrides.get(chatID) === status) { return; } - const nextOverrides = new Map(state.subagentStatusOverrides); - nextOverrides.set(chatID, status); - setState((current) => ({ - ...current, - subagentStatusOverrides: nextOverrides, - })); + setState((current) => { + if (current.subagentStatusOverrides.get(chatID) === status) { + return current; + } + const nextOverrides = new Map(current.subagentStatusOverrides); + nextOverrides.set(chatID, status); + return { ...current, subagentStatusOverrides: nextOverrides }; + }); }, resetTransientState: () => { if ( @@ -564,6 +598,9 @@ export const useChatStore = ( const handleMessage = ( payload: OneWayMessageEvent, ) => { + if (disposed) { + return; + } if (payload.parseError || !payload.parsedMessage) { store.setStreamError("Failed to parse chat stream update."); return; @@ -684,6 +721,10 @@ export const useChatStore = ( continue; } case "error": { + const eventChatID = asString(streamEvent.chat_id); + if (eventChatID && eventChatID !== chatID) { + continue; + } const error = asRecord(streamEvent.error); const reason = asString(error?.message).trim() || "Chat processing failed."; @@ -698,8 +739,13 @@ export const useChatStore = ( continue; } case "retry": { + const eventChatID = asString(streamEvent.chat_id); + if (eventChatID && eventChatID !== chatID) { + continue; + } const retry = streamEvent.retry; if (retry) { + store.clearStreamState(); store.setRetryState({ attempt: retry.attempt, error: retry.error, @@ -720,6 +766,9 @@ export const useChatStore = ( if (disposed) { return; } + if (reconnectTimer !== null) { + clearTimeout(reconnectTimer); + } const delay = Math.min( RECONNECT_BASE_MS * 2 ** reconnectAttempt, RECONNECT_MAX_MS, @@ -732,6 +781,9 @@ export const useChatStore = ( if (disposed) { return; } + if (activeSocket) { + activeSocket.close(); + } // Use the latest known message ID so the server only // sends events the client hasn't seen yet. @@ -746,9 +798,13 @@ export const useChatStore = ( }; const handleDisconnect = () => { - if (disposed) { + // Guard against duplicate calls: browsers fire both + // "error" and "close" on a failed WebSocket, so we + // only process the first event per socket instance. + if (activeSocket !== socket || disposed) { return; } + activeSocket = null; // Show the error only on the first disconnect (not // while we are already retrying). if (reconnectAttempt === 0) { diff --git a/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts b/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts index e1c6daf2e6..08d449f7fb 100644 --- a/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts +++ b/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts @@ -134,7 +134,70 @@ describe("applyMessagePartToStreamState", () => { expect(result).not.toBeNull(); const ids = Object.keys(result!.toolCalls); expect(ids).toHaveLength(1); - expect(ids[0]).toBe("tool-call-1"); + expect(ids[0]).toMatch(/^tool-call-1-\d+$/); + }); + + it("does not collide IDs for multiple tool calls with the same name", () => { + let state: StreamState | null = null; + state = applyMessagePartToStreamState(state, { + type: "tool-call", + tool_name: "bash", + args: { command: "ls" }, + }); + state = applyMessagePartToStreamState(state, { + type: "tool-call", + tool_name: "bash", + args: { command: "pwd" }, + }); + const ids = Object.keys(state!.toolCalls); + expect(ids).toHaveLength(2); + // The two calls must have distinct IDs. + expect(ids[0]).not.toBe(ids[1]); + // Each call must retain its own args. + const calls = Object.values(state!.toolCalls); + const args = calls.map((c) => c.args); + expect(args).toContainEqual({ command: "ls" }); + expect(args).toContainEqual({ command: "pwd" }); + }); + + it("does not collide IDs for multiple tool results with the same name", () => { + let state: StreamState | null = null; + // Two tool calls with the same name, each receiving a separate result. + state = applyMessagePartToStreamState(state, { + type: "tool-call", + tool_name: "bash", + args: { command: "ls" }, + }); + state = applyMessagePartToStreamState(state, { + type: "tool-call", + tool_name: "bash", + args: { command: "pwd" }, + }); + const callIds = Object.keys(state!.toolCalls); + expect(callIds).toHaveLength(2); + + // First result arrives without an explicit tool_call_id. + state = applyMessagePartToStreamState(state, { + type: "tool-result", + tool_name: "bash", + result: { output: "file.txt" }, + }); + // Second result arrives without an explicit tool_call_id. + state = applyMessagePartToStreamState(state, { + type: "tool-result", + tool_name: "bash", + result: { output: "/home" }, + }); + + const resultIds = Object.keys(state!.toolResults); + expect(resultIds).toHaveLength(2); + // The two results must have distinct IDs. + expect(resultIds[0]).not.toBe(resultIds[1]); + // Each result must retain its own output. + const results = Object.values(state!.toolResults); + const outputs = results.map((r) => r.result); + expect(outputs).toContainEqual({ output: "file.txt" }); + expect(outputs).toContainEqual({ output: "/home" }); }); it("creates tool result entry from tool-result part", () => { diff --git a/site/src/pages/AgentsPage/AgentDetail/streamState.ts b/site/src/pages/AgentsPage/AgentDetail/streamState.ts index b4291d6ff9..289075fd9e 100644 --- a/site/src/pages/AgentsPage/AgentDetail/streamState.ts +++ b/site/src/pages/AgentsPage/AgentDetail/streamState.ts @@ -9,6 +9,8 @@ import { import { mergeStreamPayload } from "./streamingJson"; import type { MergedTool, RenderBlock, StreamState } from "./types"; +let nextFallbackID = 0; + export const createEmptyStreamState = (): StreamState => ({ blocks: [], toolCalls: {}, @@ -85,8 +87,8 @@ export const applyMessagePartToStreamState = ( ); const toolCallID = asString(part.tool_call_id) || - existingByName?.id || - `tool-call-${Object.keys(nextState.toolCalls).length + 1}`; + (existingByName && !existingByName.args ? existingByName.id : null) || + `tool-call-${Object.keys(nextState.toolCalls).length + 1}-${++nextFallbackID}`; const existing = nextState.toolCalls[toolCallID]; const nextArgs = mergeStreamPayload( existing?.args, @@ -120,9 +122,11 @@ export const applyMessagePartToStreamState = ( ); const toolCallID = asString(part.tool_call_id) || - existingByName?.id || - existingCallByName?.id || - `tool-result-${Object.keys(nextState.toolResults).length + 1}`; + (existingByName && !existingByName.result ? existingByName.id : null) || + (existingCallByName && !nextState.toolResults[existingCallByName.id] + ? existingCallByName.id + : null) || + `tool-result-${Object.keys(nextState.toolResults).length + 1}-${++nextFallbackID}`; const existing = nextState.toolResults[toolCallID]; const nextResult = mergeStreamPayload( existing?.result,