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( async function routeGeneratedAction(
state: GraphState, state: GraphState,
): Promise<"open-pr" | "take-action" | "request-help" | Send> { ): Promise<"open-pr" | "take-action" | "request-help" | Send> {
const { messages } = state; const { internalMessages } = state;
const lastMessage = messages[messages.length - 1]; const lastMessage = internalMessages[internalMessages.length - 1];
// If the message is an AI message, and it has tool calls, we should take action. // If the message is an AI message, and it has tool calls, we should take action.
if (isAIMessage(lastMessage) && lastMessage.tool_calls?.length) { if (isAIMessage(lastMessage) && lastMessage.tool_calls?.length) {

View file

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

View file

@ -39,12 +39,12 @@ export async function generateConclusion(
): Promise<GraphUpdate> { ): Promise<GraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER); 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: const userMessage = `The user's initial request is as follows:
${userRequest || "No user message found"} ${userRequest || "No user message found"}
The conversation history is as follows: 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.`; 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 { return {
messages: [response], messages: [response],
internalMessages: [response],
plan: updatedTaskPlan, plan: updatedTaskPlan,
}; };
} }

View file

@ -147,14 +147,17 @@ export async function generateAction(
requestHumanHelpTool, requestHumanHelpTool,
updatePlanTool, 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([ const response = await modelWithTools.invoke([
{ {
role: "system", role: "system",
content: formatPrompt(state), content: formatPrompt(state),
}, },
...state.messages, ...state.internalMessages,
]); ]);
const hasToolCalls = !!response.tool_calls?.length; const hasToolCalls = !!response.tool_calls?.length;
@ -178,6 +181,7 @@ export async function generateAction(
return { return {
messages: [response], messages: [response],
internalMessages: [response],
...(newSandboxSessionId && { sandboxSessionId: newSandboxSessionId }), ...(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."); throw new Error("No sandbox session ID found.");
} }
const userRequest = getUserRequest(state.messages); const userRequest = getUserRequest(state.internalMessages);
if (interruptRes.type === "accept") { if (interruptRes.type === "accept") {
const newSandboxSessionId = (await startSandbox(state.sandboxSessionId)).id; 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 model = await loadModel(config, Task.SUMMARIZER);
const modelWithTool = model.bindTools([openPrTool], { const modelWithTool = model.bindTools([openPrTool], {
tool_choice: openPrTool.name, tool_choice: openPrTool.name,
parallel_tool_calls: false,
}); });
const userRequest = getUserRequest(state.messages); const userRequest = getUserRequest(state.internalMessages);
const response = await modelWithTool.invoke([ const response = await modelWithTool.invoke([
{ {
role: "user", role: "user",
@ -141,20 +142,23 @@ export async function openPullRequest(
sandboxDeleted = await deleteSandbox(sandboxSessionId); 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 { return {
messages: [ messages: newMessages,
response, internalMessages: newMessages,
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,
},
}),
],
// If the sandbox was successfully deleted, we can remove it from the state. // If the sandbox was successfully deleted, we can remove it from the state.
...(sandboxDeleted && { sandboxSessionId: undefined }), ...(sandboxDeleted && { sandboxSessionId: undefined }),
}; };

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -88,7 +88,7 @@ export async function updatePlan(
if (!state.planChangeRequest) { if (!state.planChangeRequest) {
throw new Error("No plan change request found."); 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 ( if (
!lastMessage || !lastMessage ||
!isAIMessage(lastMessage) || !isAIMessage(lastMessage) ||
@ -108,6 +108,7 @@ export async function updatePlan(
const model = await loadModel(config, Task.PLANNER); const model = await loadModel(config, Task.PLANNER);
const modelWithTools = model.bindTools([updatePlanTool], { const modelWithTools = model.bindTools([updatePlanTool], {
tool_choice: updatePlanTool.name, tool_choice: updatePlanTool.name,
parallel_tool_calls: false,
}); });
const activeTask = getActiveTask(state.plan); const activeTask = getActiveTask(state.plan);
@ -124,7 +125,7 @@ export async function updatePlan(
state.planChangeRequest, state.planChangeRequest,
activePlanItems, activePlanItems,
); );
const userMessage = formatUserMessage(state.messages); const userMessage = formatUserMessage(state.internalMessages);
const response = await modelWithTools.invoke([ const response = await modelWithTools.invoke([
{ {
@ -160,21 +161,21 @@ export async function updatePlan(
newPlanItems, newPlanItems,
"agent", "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 { return {
messages: [ messages: [toolMessage],
new ToolMessage({ internalMessages: [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"),
}),
],
plan: newTaskPlan, plan: newTaskPlan,
planChangeRequest: null, 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 { function formatSystemPrompt(state: PlannerGraphState): string {
// It's a followup if there's more than one human message. // 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 return systemPrompt
.replace( .replace(
@ -58,11 +58,15 @@ export async function generateAction(
): Promise<PlannerGraphUpdate> { ): Promise<PlannerGraphUpdate> {
const model = await loadModel(config, Task.ACTION_GENERATOR); const model = await loadModel(config, Task.ACTION_GENERATOR);
const tools = [shellTool]; 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, returnFullMessage: true,
}); });
const response = await modelWithTools const response = await modelWithTools
.withConfig({ tags: ["nostream"] }) .withConfig({ tags: ["nostream"] })
.invoke([ .invoke([
@ -85,6 +89,7 @@ export async function generateAction(
}); });
return { return {
messages: [response],
plannerMessages: [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 { function formatSystemPrompt(state: PlannerGraphState): string {
// It's a followup if there's more than one human message. // 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;
const userRequest = getUserRequest(state.messages); const userRequest = getUserRequest(state.internalMessages);
return systemPrompt return systemPrompt
.replace( .replace(
@ -54,6 +54,7 @@ export async function generatePlan(
const model = await loadModel(config, Task.PLANNER); const model = await loadModel(config, Task.PLANNER);
const modelWithTools = model.bindTools([sessionPlanTool], { const modelWithTools = model.bindTools([sessionPlanTool], {
tool_choice: sessionPlanTool.name, tool_choice: sessionPlanTool.name,
parallel_tool_calls: false,
}); });
let optionalToolMessage: ToolMessage | undefined; let optionalToolMessage: ToolMessage | undefined;
@ -89,6 +90,7 @@ export async function generatePlan(
} }
return { return {
messages: [response],
proposedPlan: response.tool_calls[0].args.plan, proposedPlan: response.tool_calls[0].args.plan,
...(newSessionId && { sandboxSessionId: newSessionId }), ...(newSessionId && { sandboxSessionId: newSessionId }),
// Do this so that the planner state is up to date with the tool call. // 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 model = await loadModel(config, Task.SUMMARIZER);
const modelWithTools = model.bindTools([condenseContextTool], { const modelWithTools = model.bindTools([condenseContextTool], {
tool_choice: condenseContextTool.name, 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: const conversationHistoryStr = `Here is the full conversation history:
${state.plannerMessages.map(getMessageString).join("\n")}`; ${state.plannerMessages.map(getMessageString).join("\n")}`;
@ -72,6 +73,7 @@ ${state.plannerMessages.map(getMessageString).join("\n")}`;
} }
return { return {
messages: [response],
planContextSummary: toolCall.args.context, planContextSummary: toolCall.args.context,
}; };
} }

View file

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

View file

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

View file

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

View file

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

View file

@ -2,6 +2,8 @@ import "@langchain/langgraph/zod";
import { z } from "zod"; import { z } from "zod";
import { import {
LangGraphRunnableConfig, LangGraphRunnableConfig,
Messages,
messagesStateReducer,
MessagesZodState, MessagesZodState,
} from "@langchain/langgraph/web"; } from "@langchain/langgraph/web";
import { MODEL_OPTIONS, MODEL_OPTIONS_NO_THINKING } from "./models.js"; import { MODEL_OPTIONS, MODEL_OPTIONS_NO_THINKING } from "./models.js";
@ -12,6 +14,8 @@ import {
type RemoveUIMessage, type RemoveUIMessage,
} from "@langchain/langgraph-sdk/react-ui"; } from "@langchain/langgraph-sdk/react-ui";
import { GITHUB_TOKEN_COOKIE } from "../constants.js"; import { GITHUB_TOKEN_COOKIE } from "../constants.js";
import { withLangGraph } from "@langchain/langgraph/zod";
import { BaseMessage } from "@langchain/core/messages";
export type PlanItem = { export type PlanItem = {
/** /**
@ -116,6 +120,24 @@ export type TargetRepository = {
}; };
export const GraphAnnotation = MessagesZodState.extend({ 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 proposedPlan: z
.array(z.string()) .array(z.string())
.default(() => []) .default(() => [])