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.
This commit is contained in:
Kyle Carberry
2026-03-06 15:47:28 -05:00
committed by GitHub
parent b9b3c67c73
commit 2cd871e88f
3 changed files with 174 additions and 51 deletions
@@ -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<string, unknown>[]) => {
@@ -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<TypesGen.ServerSentEvent>,
) => {
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) {
@@ -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", () => {
@@ -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,