feat: Add internal messages state (#141)

* feat: Add internal messages state

* update frontend

* disable parallel tool calls
This commit is contained in:
Brace Sproul 2025-06-12 13:00:51 -07:00 • committed by GitHub
parent abf2e926f7
commit 0db1a02d13
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 173 additions and 92 deletions

View file

@ -32,8 +32,8 @@ import { plannerGraph } from "./subgraphs/index.js";
async function routeGeneratedAction(
state: GraphState,
): Promise<"open-pr" | "take-action" | "request-help" | Send> {
const { messages } = state;
const lastMessage = messages[messages.length - 1];
const { internalMessages } = state;
const lastMessage = internalMessages[internalMessages.length - 1];
// If the message is an AI message, and it has tool calls, we should take action.
if (isAIMessage(lastMessage) && lastMessage.tool_calls?.length) {

View file

@ -101,7 +101,7 @@ export async function diagnoseError(
state: GraphState,
config: GraphConfig,
): Promise<GraphUpdate> {
const lastFailedAction = state.messages.findLast(
const lastFailedAction = state.internalMessages.findLast(
(m) => isToolMessage(m) && m.status === "error",
);
if (!lastFailedAction?.content) {
@ -113,6 +113,7 @@ export async function diagnoseError(
const model = await loadModel(config, Task.SUMMARIZER);
const modelWithTools = model.bindTools([diagnoseErrorTool], {
tool_choice: diagnoseErrorTool.name,
parallel_tool_calls: false,
});
const response = await modelWithTools.invoke([
@ -126,7 +127,7 @@ export async function diagnoseError(
},
{
role: "user",
content: formatUserPrompt(state.messages),
content: formatUserPrompt(state.internalMessages),
},
]);
@ -153,5 +154,6 @@ export async function diagnoseError(
return {
messages: [response, toolMessage],
internalMessages: [response, toolMessage],
};
}

View file

@ -39,12 +39,12 @@ export async function generateConclusion(
): Promise<GraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER);
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
const userMessage = `The user's initial request is as follows:
${userRequest || "No user message found"}
The conversation history is as follows:
${state.messages.map(getMessageString).join("\n")}
${state.internalMessages.map(getMessageString).join("\n")}
Given all of this, please respond with the concise conclusion. Do not include any additional text besides the conclusion.`;
@ -71,6 +71,7 @@ Given all of this, please respond with the concise conclusion. Do not include an
return {
messages: [response],
internalMessages: [response],
plan: updatedTaskPlan,
};
}

View file

@ -147,14 +147,17 @@ export async function generateAction(
requestHumanHelpTool,
updatePlanTool,
];
const modelWithTools = model.bindTools(tools, { tool_choice: "auto" });
const modelWithTools = model.bindTools(tools, {
tool_choice: "auto",
parallel_tool_calls: false,
});
const response = await modelWithTools.invoke([
{
role: "system",
content: formatPrompt(state),
},
...state.messages,
...state.internalMessages,
]);
const hasToolCalls = !!response.tool_calls?.length;
@ -178,6 +181,7 @@ export async function generateAction(
return {
messages: [response],
internalMessages: [response],
...(newSandboxSessionId && { sandboxSessionId: newSandboxSessionId }),
};
}

View file

@ -41,7 +41,7 @@ export async function interruptPlan(state: GraphState): Promise<Command> {
throw new Error("No sandbox session ID found.");
}
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
if (interruptRes.type === "accept") {
const newSandboxSessionId = (await startSandbox(state.sandboxSessionId)).id;

View file

@ -106,9 +106,10 @@ export async function openPullRequest(
const model = await loadModel(config, Task.SUMMARIZER);
const modelWithTool = model.bindTools([openPrTool], {
tool_choice: openPrTool.name,
parallel_tool_calls: false,
});
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
const response = await modelWithTool.invoke([
{
role: "user",
@ -141,20 +142,23 @@ export async function openPullRequest(
sandboxDeleted = await deleteSandbox(sandboxSessionId);
}
const newMessages = [
response,
new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: pr
? `Created pull request: ${pr.html_url}`
: "Failed to create pull request.",
name: toolCall.name,
additional_kwargs: {
pull_request: pr,
},
}),
];
return {
messages: [
response,
new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: pr
? `Created pull request: ${pr.html_url}`
: "Failed to create pull request.",
name: toolCall.name,
additional_kwargs: {
pull_request: pr,
},
}),
],
messages: newMessages,
internalMessages: newMessages,
// If the sandbox was successfully deleted, we can remove it from the state.
...(sandboxDeleted && { sandboxSessionId: undefined }),
};

View file

@ -73,14 +73,15 @@ export async function progressPlanStep(
const model = await loadModel(config, Task.PROGRESS_PLAN_CHECKER);
const modelWithTools = model.bindTools([setTaskStatusTool], {
tool_choice: setTaskStatusTool.name,
parallel_tool_calls: false,
});
const userRequest = getUserRequest(state.messages, {
const userRequest = getUserRequest(state.internalMessages, {
returnFullMessage: true,
});
const conversationHistoryStr = `Here is the full conversation history after the user's request:
${removeFirstHumanMessage(state.messages).map(getMessageString).join("\n")}
${removeFirstHumanMessage(state.internalMessages).map(getMessageString).join("\n")}
Take all of this information, and determine whether or not you have completed this task in the plan.
Once you've determined the status of the current task, call the \`set_task_status\` tool.`;
@ -118,6 +119,8 @@ Once you've determined the status of the current task, call the \`set_task_statu
name: toolCall.name,
});
const newMessages = [response, toolMessage];
if (!isCompleted) {
logger.info(
"Current task has not been completed. Progressing to the next action.",
@ -125,7 +128,10 @@ Once you've determined the status of the current task, call the \`set_task_statu
reasoning: toolCall.args.reasoning,
},
);
const commandUpdate: GraphUpdate = { messages: [response, toolMessage] };
const commandUpdate: GraphUpdate = {
messages: newMessages,
internalMessages: newMessages,
};
return new Command({
goto: "generate-action",
update: commandUpdate,
@ -146,7 +152,8 @@ Once you've determined the status of the current task, call the \`set_task_statu
"Found no remaining tasks in the plan during the check plan step. Continuing to the conclusion generation step.",
);
const commandUpdate: GraphUpdate = {
messages: [response, toolMessage],
messages: newMessages,
internalMessages: newMessages,
// Even though there are no remaining tasks, still mark as completed so the UI reflects that the task is completed.
plan: updatedPlanTasks,
};
@ -164,7 +171,8 @@ Once you've determined the status of the current task, call the \`set_task_statu
});
const commandUpdate: GraphUpdate = {
messages: [response, toolMessage],
messages: newMessages,
internalMessages: newMessages,
plan: updatedPlanTasks,
};

View file

@ -13,7 +13,7 @@ ${helpRequest}
};
export async function requestHelp(state: GraphState): Promise<Command> {
const lastMessage = state.messages[state.messages.length - 1];
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
if (!isAIMessage(lastMessage) || !lastMessage.tool_calls?.length) {
throw new Error("Last message is not an AI message with tool calls.");
}
@ -53,14 +53,14 @@ export async function requestHelp(state: GraphState): Promise<Command> {
throw new Error("Interrupt response expected to be a string.");
}
await startSandbox(sandboxSessionId);
const toolMessage = new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Human response: ${interruptRes.args}`,
status: "success",
});
const commandUpdate: GraphUpdate = {
messages: [
new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Human response: ${interruptRes.args}`,
status: "success",
}),
],
messages: [toolMessage],
internalMessages: [toolMessage],
};
return new Command({
goto: "generate-action",

View file

@ -134,10 +134,11 @@ async function identifyTasksToModifyFunc(
{
// The model should always call the tool when identifying plan changes.
tool_choice: identifyPlanChangesTool.name,
parallel_tool_calls: false,
},
);
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
const response = await modelWithIdentifyChangesTool.invoke([
{
role: "user",
@ -197,9 +198,10 @@ async function updatePlanTasksFunc(
const modelWithUpdatePlanTasksTool = model.bindTools([updatePlanTasksTool], {
// The model should always call the tool when identifying plan changes.
tool_choice: updatePlanTasksTool.name,
parallel_tool_calls: false,
});
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
const response = await modelWithUpdatePlanTasksTool.invoke([
{
role: "user",

View file

@ -107,7 +107,7 @@ async function generateTaskSummaryFunc(
},
{
role: "user",
content: formatUserMessage(state.messages, activePlanItems),
content: formatUserMessage(state.internalMessages, activePlanItems),
},
]);
@ -140,7 +140,7 @@ export async function summarizeTaskSteps(
taskSummary.summary,
);
const removedMessages = removeLastTaskMessages(state.messages);
const removedMessages = removeLastTaskMessages(state.internalMessages);
logger.info(`Removing ${removedMessages.length} message(s) from state.`);
const condensedTaskMessage = new AIMessage({
@ -155,7 +155,8 @@ export async function summarizeTaskSteps(
const allTasksCompleted = activePlanItems.every((p) => p.completed);
if (allTasksCompleted) {
const commandUpdate: GraphUpdate = {
messages: newMessagesStateUpdate,
messages: [condensedTaskMessage],
internalMessages: newMessagesStateUpdate,
plan: updatedTaskPlan,
};
return new Command({
@ -165,7 +166,8 @@ export async function summarizeTaskSteps(
}
const commandUpdate: GraphUpdate = {
messages: newMessagesStateUpdate,
messages: [condensedTaskMessage],
internalMessages: newMessagesStateUpdate,
plan: updatedTaskPlan,
};
return new Command({

View file

@ -47,7 +47,7 @@ export async function takeAction(
state: GraphState,
config: GraphConfig,
): Promise<Command> {
const lastMessage = state.messages[state.messages.length - 1];
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
if (!isAIMessage(lastMessage) || !lastMessage.tool_calls?.length) {
throw new Error("Last message is not an AI message with tool calls.");
@ -68,17 +68,17 @@ export async function takeAction(
if (!tool) {
logger.error(`Unknown tool: ${toolCall.name}`);
const toolMessage = new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name,
status: "error",
});
return new Command({
goto: "progress-plan-step",
update: {
messages: [
new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name,
status: "error",
}),
],
messages: [toolMessage],
internalMessages: [toolMessage],
},
});
}
@ -150,7 +150,7 @@ export async function takeAction(
}
const shouldRouteDiagnoseNode = shouldDiagnoseError(
[...state.messages, toolMessage].filter(
[...state.internalMessages, toolMessage].filter(
(m): m is ToolMessage =>
isToolMessage(m) && !m.additional_kwargs?.is_diagnosis,
),
@ -162,6 +162,7 @@ export async function takeAction(
goto: shouldRouteDiagnoseNode ? "diagnose-error" : "progress-plan-step",
update: {
messages: [toolMessage],
internalMessages: [toolMessage],
...(branchName && { branchName }),
codebaseTree,
},

View file

@ -88,7 +88,7 @@ export async function updatePlan(
if (!state.planChangeRequest) {
throw new Error("No plan change request found.");
}
const lastMessage = state.messages[state.messages.length - 1];
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
if (
!lastMessage ||
!isAIMessage(lastMessage) ||
@ -108,6 +108,7 @@ export async function updatePlan(
const model = await loadModel(config, Task.PLANNER);
const modelWithTools = model.bindTools([updatePlanTool], {
tool_choice: updatePlanTool.name,
parallel_tool_calls: false,
});
const activeTask = getActiveTask(state.plan);
@ -124,7 +125,7 @@ export async function updatePlan(
state.planChangeRequest,
activePlanItems,
);
const userMessage = formatUserMessage(state.messages);
const userMessage = formatUserMessage(state.internalMessages);
const response = await modelWithTools.invoke([
{
@ -160,21 +161,21 @@ export async function updatePlan(
newPlanItems,
"agent",
);
const toolMessage = new ToolMessage({
tool_call_id: updatePlanToolCallId,
content:
"Successfully updated the plan. The complete updated plan items are as follow:\n\n" +
newPlanItems
.map(
(p) =>
`<plan-item completed="${p.completed}" index="${p.index}">${p.plan}</plan-item>`,
)
.join("\n"),
});
return {
messages: [
new ToolMessage({
tool_call_id: updatePlanToolCallId,
content:
"Successfully updated the plan. The complete updated plan items are as follow:\n\n" +
newPlanItems
.map(
(p) =>
`<plan-item completed="${p.completed}" index="${p.index}">${p.plan}</plan-item>`,
)
.join("\n"),
}),
],
messages: [toolMessage],
internalMessages: [toolMessage],
plan: newTaskPlan,
planChangeRequest: null,
};

View file

@ -37,7 +37,7 @@ The user's request is the first user message in the conversation below. Ensure y
function formatSystemPrompt(state: PlannerGraphState): string {
// It's a followup if there's more than one human message.
const isFollowup = state.messages.filter(isHumanMessage).length > 1;
const isFollowup = state.internalMessages.filter(isHumanMessage).length > 1;
return systemPrompt
.replace(
@ -58,11 +58,15 @@ export async function generateAction(
): Promise<PlannerGraphUpdate> {
const model = await loadModel(config, Task.ACTION_GENERATOR);
const tools = [shellTool];
const modelWithTools = model.bindTools(tools, { tool_choice: "auto" });
const modelWithTools = model.bindTools(tools, {
tool_choice: "auto",
parallel_tool_calls: false,
});
const userRequest = getUserRequest(state.messages, {
const userRequest = getUserRequest(state.internalMessages, {
returnFullMessage: true,
});
const response = await modelWithTools
.withConfig({ tags: ["nostream"] })
.invoke([
@ -85,6 +89,7 @@ export async function generateAction(
});
return {
messages: [response],
plannerMessages: [response],
};
}

View file

@ -36,8 +36,8 @@ The user's request is as follows. Ensure you generate your plan in accordance wi
function formatSystemPrompt(state: PlannerGraphState): string {
// It's a followup if there's more than one human message.
const isFollowup = state.messages.filter(isHumanMessage).length > 1;
const userRequest = getUserRequest(state.messages);
const isFollowup = state.internalMessages.filter(isHumanMessage).length > 1;
const userRequest = getUserRequest(state.internalMessages);
return systemPrompt
.replace(
@ -54,6 +54,7 @@ export async function generatePlan(
const model = await loadModel(config, Task.PLANNER);
const modelWithTools = model.bindTools([sessionPlanTool], {
tool_choice: sessionPlanTool.name,
parallel_tool_calls: false,
});
let optionalToolMessage: ToolMessage | undefined;
@ -89,6 +90,7 @@ export async function generatePlan(
}
return {
messages: [response],
proposedPlan: response.tool_calls[0].args.plan,
...(newSessionId && { sandboxSessionId: newSessionId }),
// Do this so that the planner state is up to date with the tool call.

View file

@ -48,9 +48,10 @@ export async function summarizer(
const model = await loadModel(config, Task.SUMMARIZER);
const modelWithTools = model.bindTools([condenseContextTool], {
tool_choice: condenseContextTool.name,
parallel_tool_calls: false,
});
const userRequest = getUserRequest(state.messages);
const userRequest = getUserRequest(state.internalMessages);
const conversationHistoryStr = `Here is the full conversation history:
${state.plannerMessages.map(getMessageString).join("\n")}`;
@ -72,6 +73,7 @@ ${state.plannerMessages.map(getMessageString).join("\n")}`;
}
return {
messages: [response],
planContextSummary: toolCall.args.context,
};
}

View file

@ -34,18 +34,22 @@ export async function takeAction(
if (!tool) {
logger.error(`Unknown tool: ${toolCall.name}`);
const toolMessage = new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name,
status: "error",
});
return {
plannerMessages: [
new ToolMessage({
tool_call_id: toolCall.id ?? "",
content: `Unknown tool: ${toolCall.name}`,
name: toolCall.name,
status: "error",
}),
],
messages: [toolMessage],
plannerMessages: [toolMessage],
};
}
logger.info("Executing planner tool action", {
...toolCall,
});
let result = "";
let toolCallStatus: "success" | "error" = "success";
try {
@ -86,7 +90,13 @@ export async function takeAction(
status: toolCallStatus,
});
logger.info("Completed planner tool action", {
tool_call_id: toolCall.id,
status: toolCallStatus,
});
return {
messages: [toolMessage],
plannerMessages: [toolMessage],
};
}

View file

@ -1,14 +1,24 @@
import "@langchain/langgraph/zod";
import { z } from "zod";
import { addMessages, Messages } from "@langchain/langgraph";
import { Messages, messagesStateReducer } from "@langchain/langgraph";
import { BaseMessage } from "@langchain/core/messages";
import { GraphAnnotation } from "@open-swe/shared/open-swe/types";
import { withLangGraph } from "@langchain/langgraph/zod";
export const PlannerGraphStateObj = GraphAnnotation.extend({
plannerMessages: z
.custom<BaseMessage[]>()
.default(() => [])
.langgraph.reducer<Messages>((state, update) => addMessages(state, update)),
plannerMessages: withLangGraph<BaseMessage[], Messages>(
z.custom<BaseMessage[]>(),
{
reducer: {
schema: z.custom<Messages>(),
fn: messagesStateReducer,
},
jsonSchemaExtra: {
langgraph_type: "messages",
},
default: () => [],
},
),
});
export type PlannerGraphState = z.infer<typeof PlannerGraphStateObj>;

View file

@ -246,12 +246,14 @@ export function Thread() {
const context =
Object.keys(artifactContext).length > 0 ? artifactContext : undefined;
const newMessages = [
...toolMessages,
newHumanMessage,
] as unknown as BaseMessage[];
stream.submit(
{
messages: [
...toolMessages,
newHumanMessage,
] as unknown as BaseMessage[],
messages: newMessages,
internalMessages: newMessages,
context,
targetRepository: selectedRepository,
},

View file

@ -53,9 +53,12 @@ export function HumanMessage({
const handleSubmitEdit = () => {
setIsEditing(false);
const newMessage: Message = { type: "human", content: value };
const newMessage = {
type: "human",
content: value,
} as unknown as BaseMessage;
thread.submit(
{ messages: [newMessage] as unknown as BaseMessage[] },
{ messages: [newMessage], internalMessages: [newMessage] },
{
checkpoint: parentCheckpoint,
streamMode: ["values"],

View file

@ -2,6 +2,8 @@ import "@langchain/langgraph/zod";
import { z } from "zod";
import {
LangGraphRunnableConfig,
Messages,
messagesStateReducer,
MessagesZodState,
} from "@langchain/langgraph/web";
import { MODEL_OPTIONS, MODEL_OPTIONS_NO_THINKING } from "./models.js";
@ -12,6 +14,8 @@ import {
type RemoveUIMessage,
} from "@langchain/langgraph-sdk/react-ui";
import { GITHUB_TOKEN_COOKIE } from "../constants.js";
import { withLangGraph } from "@langchain/langgraph/zod";
import { BaseMessage } from "@langchain/core/messages";
export type PlanItem = {
/**
@ -116,6 +120,24 @@ export type TargetRepository = {
};
export const GraphAnnotation = MessagesZodState.extend({
/**
* The internal messages. These are the messages which are
* passed to the LLM, truncated, removed etc. The main `messages`
* key is never modified to persist the content show on the client.
*/
internalMessages: withLangGraph<BaseMessage[], Messages>(
z.custom<BaseMessage[]>(),
{
reducer: {
schema: z.custom<Messages>(),
fn: messagesStateReducer,
},
jsonSchemaExtra: {
langgraph_type: "messages",
},
default: () => [],
},
),
proposedPlan: z
.array(z.string())
.default(() => [])