fix: Bugs in planner node (#610)

* fix: Bugs in planner node

* cr

* cr
This commit is contained in:
Brace Sproul 2025-07-30 12:17:22 -07:00 • committed by GitHub
parent 6440542b23
commit 3105a318d2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 39 additions and 4 deletions

View file

@ -69,6 +69,7 @@ function routeGeneratedAction(
) {
// Need to return a `Send` here so that we can update the state to include the plan change request.
return new Send("update-plan", {
...state,
planChangeRequest: toolCall.args?.update_plan_reasoning,
});
}

View file

@ -17,6 +17,7 @@ import {
updateTaskPlanItems,
} from "@open-swe/shared/open-swe/tasks";
import {
AIMessage,
BaseMessage,
isAIMessage,
ToolMessage,
@ -28,6 +29,7 @@ import { createUpdatePlanToolFields } from "@open-swe/shared/open-swe/tools";
import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js";
import { trackCachePerformance } from "../../../utils/caching.js";
import { getModelManager } from "../../../utils/llms/model-manager.js";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
const logger = createLogger(LogLevel.INFO, "UpdatePlanNode");
@ -79,6 +81,8 @@ const updatePlanTool = {
schema: updatePlanToolSchema,
};
const updatePlanReasoningTool = createUpdatePlanToolFields();
const formatSystemPrompt = (
userRequest: string,
reasoning: string,
@ -98,14 +102,34 @@ const formatUserMessage = (messages: BaseMessage[]): string => {
${messages.map(getMessageString).join("\n")}`;
};
function removeUncalledTools(lastMessage: AIMessage): AIMessage {
if (!lastMessage.tool_calls?.length || lastMessage.tool_calls?.length === 1) {
// check for no tool calls. will never happen, but need for type safety
// only one tool call, this is the update plan tool call. no-op
return lastMessage;
}
const updatePlanReasoningToolCall = lastMessage.tool_calls?.find(
(tc) => tc.name === updatePlanReasoningTool.name,
);
if (!updatePlanReasoningToolCall) {
throw new Error("Update plan reasoning tool call not found.");
}
// Return the last message, only changing the tool calls to only include the update plan reasoning tool call.
return new AIMessage({
...lastMessage,
tool_calls: [updatePlanReasoningToolCall],
});
}
export async function updatePlan(
state: GraphState,
config: GraphConfig,
): Promise<GraphUpdate> {
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
const updatePlanReasoningTool = createUpdatePlanToolFields();
if (!lastMessage || !isAIMessage(lastMessage)) {
if (!lastMessage || !isAIMessage(lastMessage) || !lastMessage.id) {
throw new Error("Last message was not an AI message");
}
@ -190,6 +214,16 @@ export async function updatePlan(
newPlanItems,
"agent",
);
// Update the github issue to reflect the changes in the plan
await addTaskPlanToIssue(
{
githubIssueId: state.githubIssueId,
targetRepository: state.targetRepository,
},
config,
newTaskPlan,
);
const toolMessage = new ToolMessage({
id: uuidv4(),
tool_call_id: updatePlanToolCallId,
@ -204,8 +238,8 @@ export async function updatePlan(
});
return {
messages: [toolMessage],
internalMessages: [toolMessage],
messages: [removeUncalledTools(lastMessage), toolMessage],
internalMessages: [removeUncalledTools(lastMessage), toolMessage],
taskPlan: newTaskPlan,
tokenData: trackCachePerformance(response, modelName),
};