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. // Need to return a `Send` here so that we can update the state to include the plan change request.
return new Send("update-plan", { return new Send("update-plan", {
...state,
planChangeRequest: toolCall.args?.update_plan_reasoning, planChangeRequest: toolCall.args?.update_plan_reasoning,
}); });
} }

View file

@ -17,6 +17,7 @@ import {
updateTaskPlanItems, updateTaskPlanItems,
} from "@open-swe/shared/open-swe/tasks"; } from "@open-swe/shared/open-swe/tasks";
import { import {
AIMessage,
BaseMessage, BaseMessage,
isAIMessage, isAIMessage,
ToolMessage, ToolMessage,
@ -28,6 +29,7 @@ import { createUpdatePlanToolFields } from "@open-swe/shared/open-swe/tools";
import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js"; import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js";
import { trackCachePerformance } from "../../../utils/caching.js"; import { trackCachePerformance } from "../../../utils/caching.js";
import { getModelManager } from "../../../utils/llms/model-manager.js"; import { getModelManager } from "../../../utils/llms/model-manager.js";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
const logger = createLogger(LogLevel.INFO, "UpdatePlanNode"); const logger = createLogger(LogLevel.INFO, "UpdatePlanNode");
@ -79,6 +81,8 @@ const updatePlanTool = {
schema: updatePlanToolSchema, schema: updatePlanToolSchema,
}; };
const updatePlanReasoningTool = createUpdatePlanToolFields();
const formatSystemPrompt = ( const formatSystemPrompt = (
userRequest: string, userRequest: string,
reasoning: string, reasoning: string,
@ -98,14 +102,34 @@ const formatUserMessage = (messages: BaseMessage[]): string => {
${messages.map(getMessageString).join("\n")}`; ${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( export async function updatePlan(
state: GraphState, state: GraphState,
config: GraphConfig, config: GraphConfig,
): Promise<GraphUpdate> { ): Promise<GraphUpdate> {
const lastMessage = state.internalMessages[state.internalMessages.length - 1]; 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"); throw new Error("Last message was not an AI message");
} }
@ -190,6 +214,16 @@ export async function updatePlan(
newPlanItems, newPlanItems,
"agent", "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({ const toolMessage = new ToolMessage({
id: uuidv4(), id: uuidv4(),
tool_call_id: updatePlanToolCallId, tool_call_id: updatePlanToolCallId,
@ -204,8 +238,8 @@ export async function updatePlan(
}); });
return { return {
messages: [toolMessage], messages: [removeUncalledTools(lastMessage), toolMessage],
internalMessages: [toolMessage], internalMessages: [removeUncalledTools(lastMessage), toolMessage],
taskPlan: newTaskPlan, taskPlan: newTaskPlan,
tokenData: trackCachePerformance(response, modelName), tokenData: trackCachePerformance(response, modelName),
}; };