mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user