mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 20:32:12 +00:00
feat: Add internal messages state (#141)
* feat: Add internal messages state * update frontend * disable parallel tool calls
This commit is contained in:
parent
abf2e926f7
commit
0db1a02d13
20 changed files with 173 additions and 92 deletions
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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],
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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 }),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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;
|
||||||
|
|
|
||||||
|
|
@ -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 }),
|
||||||
};
|
};
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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({
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
};
|
};
|
||||||
|
|
|
||||||
|
|
@ -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],
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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],
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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>;
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -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"],
|
||||||
|
|
|
||||||
|
|
@ -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(() => [])
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue