fix: stream messages properly (#416)

This commit is contained in:
Brace Sproul 2025-07-15 13:35:37 -07:00 • committed by GitHub
parent 3b7dee85f7
commit 01f0698e02
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 39 additions and 114 deletions

View file

@ -233,6 +233,7 @@ export async function classifyMessage(
command: { command: {
resume: plannerResume, resume: plannerResume,
}, },
streamMode: ["values", "messages-tuple", "custom"],
}, },
); );
newPlannerId = newPlannerRun.run_id; newPlannerId = newPlannerRun.run_id;

View file

@ -92,7 +92,7 @@ ${ISSUE_CONTENT_CLOSE_TAG}`,
}, },
ifNotExists: "create", ifNotExists: "create",
streamResumable: true, streamResumable: true,
streamMode: ["values", "messages", "custom"], streamMode: ["values", "messages-tuple", "custom"],
}); });
return { return {

View file

@ -49,7 +49,7 @@ export async function startPlanner(
ifNotExists: "create", ifNotExists: "create",
multitaskStrategy: "enqueue", multitaskStrategy: "enqueue",
streamResumable: true, streamResumable: true,
streamMode: ["values", "messages", "custom"], streamMode: ["values", "messages-tuple", "custom"],
}, },
); );

View file

@ -61,6 +61,7 @@ export async function generatePlan(
if (isAIMessage(lastMessage) && lastMessage.tool_calls?.[0]) { if (isAIMessage(lastMessage) && lastMessage.tool_calls?.[0]) {
const lastMessageToolCall = lastMessage.tool_calls?.[0]; const lastMessageToolCall = lastMessage.tool_calls?.[0];
optionalToolMessage = new ToolMessage({ optionalToolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: lastMessageToolCall.id ?? "", tool_call_id: lastMessageToolCall.id ?? "",
name: lastMessageToolCall.name, name: lastMessageToolCall.name,
content: "Tool call not executed. Max actions reached.", content: "Tool call not executed. Max actions reached.",

View file

@ -106,7 +106,7 @@ async function startProgrammerRun(input: {
ifNotExists: "create", ifNotExists: "create",
streamResumable: true, streamResumable: true,
streamSubgraphs: true, streamSubgraphs: true,
streamMode: ["values", "messages", "custom", "events"], streamMode: ["values", "messages-tuple", "custom"],
}, },
); );

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
isAIMessage, isAIMessage,
isToolMessage, isToolMessage,
@ -30,6 +31,7 @@ import { getSandboxWithErrorHandling } from "../../../utils/sandbox.js";
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js"; import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
import { Command } from "@langchain/langgraph"; import { Command } from "@langchain/langgraph";
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js"; import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js";
import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants";
const logger = createLogger(LogLevel.INFO, "TakeAction"); const logger = createLogger(LogLevel.INFO, "TakeAction");
@ -79,6 +81,7 @@ export async function takeActions(
if (!tool) { if (!tool) {
logger.error(`Unknown tool: ${toolCall.name}`); logger.error(`Unknown tool: ${toolCall.name}`);
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: `${DO_NOT_RENDER_ID_PREFIX}${uuidv4()}`,
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`, content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name, name: toolCall.name,
@ -145,6 +148,7 @@ export async function takeActions(
: truncateOutput(result); : truncateOutput(result);
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: truncatedOutput, content: truncatedOutput,
name: toolCall.name, name: toolCall.name,

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
BaseMessage, BaseMessage,
isToolMessage, isToolMessage,
@ -136,6 +137,7 @@ export async function diagnoseError(
}); });
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Successfully diagnosed error. Please use the diagnosis to continue with the next action.`, content: `Successfully diagnosed error. Please use the diagnosis to continue with the next action.`,
name: toolCall.name, name: toolCall.name,

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
GraphConfig, GraphConfig,
GraphState, GraphState,
@ -140,6 +141,7 @@ export async function openPullRequest(
}, },
}), }),
new ToolMessage({ new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: pr content: pr
? `Created pull request: ${pr.html_url}` ? `Created pull request: ${pr.html_url}`

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createLogger, LogLevel } from "../../../utils/logger.js";
import { import {
GraphConfig, GraphConfig,
@ -108,6 +109,7 @@ Once you've determined the status of the current task, call either the \`mark_ta
const isCompleted = toolCall.name === markCompletedTool.name; const isCompleted = toolCall.name === markCompletedTool.name;
const currentTask = getCurrentPlanItem(activePlanItems); const currentTask = getCurrentPlanItem(activePlanItems);
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Saved task status as ${isCompleted ? "completed" : "not completed"} for task ${currentTask?.plan || "unknown"}`, content: `Saved task status as ${isCompleted ? "completed" : "not completed"} for task ${currentTask?.plan || "unknown"}`,
name: toolCall.name, name: toolCall.name,

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { isAIMessage, ToolMessage } from "@langchain/core/messages"; import { isAIMessage, ToolMessage } from "@langchain/core/messages";
import { import {
GraphConfig, GraphConfig,
@ -71,6 +72,7 @@ export async function requestHelp(
); );
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Human response: ${interruptRes.args}`, content: `Human response: ${interruptRes.args}`,
status: "success", status: "success",

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { isAIMessage, ToolMessage } from "@langchain/core/messages"; import { isAIMessage, ToolMessage } from "@langchain/core/messages";
import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createLogger, LogLevel } from "../../../utils/logger.js";
import { import {
@ -78,6 +79,7 @@ export async function takeAction(
if (!tool) { if (!tool) {
logger.error(`Unknown tool: ${toolCall.name}`); logger.error(`Unknown tool: ${toolCall.name}`);
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`, content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name, name: toolCall.name,
@ -127,6 +129,7 @@ export async function takeAction(
} }
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: truncateOutput(result), content: truncateOutput(result),
name: toolCall.name, name: toolCall.name,

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
GraphState, GraphState,
GraphConfig, GraphConfig,
@ -174,6 +175,7 @@ export async function updatePlan(
"agent", "agent",
); );
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: updatePlanToolCallId, tool_call_id: updatePlanToolCallId,
content: content:
"Successfully updated the plan. The complete updated plan items are as follow:\n\n" + "Successfully updated the plan. The complete updated plan items are as follow:\n\n" +

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
ReviewerGraphState, ReviewerGraphState,
ReviewerGraphUpdate, ReviewerGraphUpdate,
@ -96,6 +97,7 @@ export async function finalReview(
if (toolCall.name === completedTool.name) { if (toolCall.name === completedTool.name) {
// Marked as completed. No further actions necessary. // Marked as completed. No further actions necessary.
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: "Marked task as completed.", content: "Marked task as completed.",
}); });
@ -144,6 +146,7 @@ export async function finalReview(
); );
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: "Marked task as incomplete.", content: "Marked task as incomplete.",
}); });

View file

@ -1,3 +1,4 @@
import { v4 as uuidv4 } from "uuid";
import { import {
BaseMessage, BaseMessage,
isToolMessage, isToolMessage,
@ -115,6 +116,7 @@ export async function diagnoseError(
}); });
const toolMessage = new ToolMessage({ const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "", tool_call_id: toolCall.id ?? "",
content: `Successfully diagnosed error. Please use the diagnosis to continue with the next action.`, content: `Successfully diagnosed error. Please use the diagnosis to continue with the next action.`,
name: toolCall.name, name: toolCall.name,

View file

@ -166,7 +166,7 @@ webhooks.on("issues.labeled", async ({ payload }) => {
}, },
ifNotExists: "create", ifNotExists: "create",
streamResumable: true, streamResumable: true,
streamMode: ["values", "messages", "custom"], streamMode: ["values", "messages-tuple", "custom"],
}); });
logger.info("Created new run from GitHub issue.", { logger.info("Created new run from GitHub issue.", {

View file

@ -2,10 +2,7 @@ import { isAIMessageSDK, isHumanMessageSDK } from "@/lib/langchain-messages";
import { UseStream, useStream } from "@langchain/langgraph-sdk/react"; import { UseStream, useStream } from "@langchain/langgraph-sdk/react";
import { AssistantMessage } from "../thread/messages/ai"; import { AssistantMessage } from "../thread/messages/ai";
import { Dispatch, SetStateAction, useEffect, useRef, useState } from "react"; import { Dispatch, SetStateAction, useEffect, useRef, useState } from "react";
import { import { ManagerGraphState } from "@open-swe/shared/open-swe/manager/types";
ManagerGraphState,
ManagerGraphUpdate,
} from "@open-swe/shared/open-swe/manager/types";
import { useCancelStream } from "@/hooks/useCancelStream"; import { useCancelStream } from "@/hooks/useCancelStream";
import { import {
isCustomNodeEvent, isCustomNodeEvent,
@ -26,7 +23,6 @@ import { GraphState, PlanItem } from "@open-swe/shared/open-swe/types";
import { HumanResponse } from "@langchain/langgraph/prebuilt"; import { HumanResponse } from "@langchain/langgraph/prebuilt";
import { LoadingActionsCardContent } from "./thread-view-loading"; import { LoadingActionsCardContent } from "./thread-view-loading";
import { Interrupt } from "../thread/messages/interrupt"; import { Interrupt } from "../thread/messages/interrupt";
import { debounce } from "lodash";
interface AcceptedPlanEventData { interface AcceptedPlanEventData {
planTitle: string; planTitle: string;
@ -99,71 +95,6 @@ const getCustomNodeEventsFromMessages = (
.flat(); .flat();
}; };
function addMessagesToState(
existingMessages: Message[],
newMessages: Message[],
): Message[] {
const existingIds = new Set(existingMessages.map((message) => message.id));
// First deduplicate within newMessages array itself
const seenNewIds = new Set<string>();
const uniqueNewMessages = newMessages.filter((message) => {
// Skip messages without IDs or those already in existingMessages
if (message.id && existingIds.has(message.id)) return false;
// Handle duplicates within newMessages
if (message.id) {
if (seenNewIds.has(message.id)) return false;
seenNewIds.add(message.id);
}
return true;
});
return [...existingMessages, ...uniqueNewMessages];
}
function isNodeEndMessagesUpdate(
data: unknown,
): data is { output: { messages: Message[] } } {
return !!(
typeof data === "object" &&
data !== null &&
"output" in data &&
data.output &&
typeof data.output === "object" &&
"messages" in data.output &&
data.output.messages &&
Array.isArray(data.output.messages)
);
}
function isNodeEndCommandUpdate(data: unknown): data is {
output: { lg_name: string; goto: string; update: ManagerGraphUpdate };
} {
return !!(
typeof data === "object" &&
data !== null &&
"output" in data &&
data.output &&
typeof data.output === "object" &&
"lg_name" in data.output &&
"goto" in data.output &&
"update" in data.output &&
typeof data.output.lg_name === "string" &&
(typeof data.output.goto === "string" || Array.isArray(data.output.goto)) &&
typeof data.output.update === "object"
);
}
const REVIEWER_NODE_IDS = [
"initialize-state",
"generate-review-actions",
"take-review-actions",
"diagnose-reviewer-error",
"final-review",
];
export function ActionsRenderer<State extends PlannerGraphState | GraphState>({ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
graphId, graphId,
threadId, threadId,
@ -178,12 +109,6 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
); );
const joinedRunId = useRef<string | undefined>(undefined); const joinedRunId = useRef<string | undefined>(undefined);
const [streamLoading, setStreamLoading] = useState(false); const [streamLoading, setStreamLoading] = useState(false);
const [mergedMessages, setMergedMessages] = useState<Message[]>([]);
const debouncedSetMessages = useRef(
debounce((messages: Message[]) => {
setMergedMessages((prev) => addMessagesToState(prev, messages));
}, 100),
).current;
const stream = useStream<State>({ const stream = useStream<State>({
apiUrl: process.env.NEXT_PUBLIC_API_URL, apiUrl: process.env.NEXT_PUBLIC_API_URL,
@ -195,23 +120,6 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
setCustomNodeEvents((prev) => [...prev, event]); setCustomNodeEvents((prev) => [...prev, event]);
} }
}, },
onLangChainEvent: (data) => {
if (
data.event === "on_chain_end" &&
data.metadata?.langgraph_node &&
REVIEWER_NODE_IDS.includes(data.metadata.langgraph_node as string) &&
data.data
) {
if (isNodeEndCommandUpdate(data.data)) {
const outputMessages = data.data.output.update
.messages as unknown as Message[];
debouncedSetMessages(outputMessages);
} else if (isNodeEndMessagesUpdate(data.data)) {
const outputMessages = data.data.output.messages;
debouncedSetMessages(outputMessages);
}
}
},
fetchStateHistory: false, fetchStateHistory: false,
}); });
@ -240,7 +148,7 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
} }
useEffect(() => { useEffect(() => {
const allCustomEvents = getCustomNodeEventsFromMessages(mergedMessages); const allCustomEvents = getCustomNodeEventsFromMessages(stream.messages);
if (!allCustomEvents?.length) { if (!allCustomEvents?.length) {
return; return;
} }
@ -263,18 +171,18 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
return prev; return prev;
}); });
}, [mergedMessages]); }, [stream.messages]);
// Clear streamLoading as soon as we get any content (agent has started running) // Clear streamLoading as soon as we get any content (agent has started running)
useEffect(() => { useEffect(() => {
const hasContent = const hasContent =
(mergedMessages && mergedMessages.length > 0) || (stream.messages && stream.messages.length > 0) ||
customNodeEvents.length > 0; customNodeEvents.length > 0;
if (hasContent && streamLoading) { if (hasContent && streamLoading) {
setStreamLoading(false); setStreamLoading(false);
} }
}, [mergedMessages, customNodeEvents, streamLoading]); }, [stream.messages, customNodeEvents, streamLoading]);
// TODO: If the SDK changes go in, use this instead: // TODO: If the SDK changes go in, use this instead:
// stream.joinStream(runId, undefined, { streamMode: ["values", "messages", "custom"]}).catch(console.error); // stream.joinStream(runId, undefined, { streamMode: ["values", "messages", "custom"]}).catch(console.error);
@ -300,15 +208,15 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
}, [onStreamReady, runId]); // Depend on runId instead of cancelRun to avoid infinite loops }, [onStreamReady, runId]); // Depend on runId instead of cancelRun to avoid infinite loops
// Filter out human & do not render messages // Filter out human & do not render messages
const filteredMessages = mergedMessages?.filter( const filteredMessages = stream.messages?.filter(
(m) => (m) =>
!isHumanMessageSDK(m) && !isHumanMessageSDK(m) &&
!(m.id && m.id.startsWith(DO_NOT_RENDER_ID_PREFIX)), !(m.id && m.id.startsWith(DO_NOT_RENDER_ID_PREFIX)),
); );
const isLastMessageHidden = !!( const isLastMessageHidden = !!(
mergedMessages?.length > 0 && stream.messages?.length > 0 &&
mergedMessages[mergedMessages.length - 1].id && stream.messages[stream.messages.length - 1].id &&
mergedMessages[mergedMessages.length - 1].id?.startsWith( stream.messages[stream.messages.length - 1].id?.startsWith(
DO_NOT_RENDER_ID_PREFIX, DO_NOT_RENDER_ID_PREFIX,
) )
); );
@ -335,13 +243,6 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
} }
}, [stream.values, graphId]); }, [stream.values, graphId]);
useEffect(() => {
debouncedSetMessages(stream.messages);
return () => {
debouncedSetMessages.cancel();
};
}, [stream.messages, debouncedSetMessages]);
if (streamLoading) { if (streamLoading) {
return <LoadingActionsCardContent />; return <LoadingActionsCardContent />;
} }
@ -360,7 +261,7 @@ export function ActionsRenderer<State extends PlannerGraphState | GraphState>({
<AssistantMessage <AssistantMessage
key={m.id} key={m.id}
thread={stream as UseStream<Record<string, unknown>>} thread={stream as UseStream<Record<string, unknown>>}
threadMessages={mergedMessages} threadMessages={stream.messages}
message={m} message={m}
isLoading={false} isLoading={false}
handleRegenerate={() => {}} handleRegenerate={() => {}}

View file

@ -101,7 +101,7 @@ export function TerminalInput({
}, },
ifNotExists: "create", ifNotExists: "create",
streamResumable: true, streamResumable: true,
streamMode: ["values", "messages", "custom"], streamMode: ["values", "messages-tuple", "custom"],
}, },
); );