From bb0026cb328f0eeed18a53192d3ab3f190fb74db Mon Sep 17 00:00:00 2001 From: can1357 Date: Wed, 11 Mar 2026 00:52:46 +0100 Subject: [PATCH] feat(coding-agent): added eager todo configuration and per-turn tool choice overrides - Added 'todo.eager' configuration setting to automatically create a comprehensive todo list after the first user message. - Added 'buildNamedToolChoice' utility function to build provider-aware tool choice constraints for named tools. - Modified tool choice resolution to support per-turn tool choice overrides via consumeNextToolChoiceOverride() method. - Implemented eager todo enforcement mechanism that injects a synthetic prompt to encourage todo creation when conditions are met. - Extracted tool choice building logic into reusable utility module for better code organization. - Added comprehensive test coverage for eager todo enforcement functionality in AgentSession. --- packages/coding-agent/CHANGELOG.md | 4 + .../src/config/settings-schema.ts | 9 + .../src/prompts/system/eager-todo.md | 16 ++ packages/coding-agent/src/sdk.ts | 2 +- .../coding-agent/src/session/agent-session.ts | 87 +++++- packages/coding-agent/src/task/executor.ts | 20 +- .../coding-agent/src/utils/tool-choice.ts | 28 ++ .../test/agent-session-eager-todo.test.ts | 262 ++++++++++++++++++ 8 files changed, 404 insertions(+), 24 deletions(-) create mode 100644 packages/coding-agent/src/prompts/system/eager-todo.md create mode 100644 packages/coding-agent/src/utils/tool-choice.ts create mode 100644 packages/coding-agent/test/agent-session-eager-todo.test.ts diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 6ea297f15..73c0c926f 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,13 +1,17 @@ # Changelog ## [Unreleased] + ### Added +- Added `todo.eager` setting to automatically create a comprehensive todo list after the first user message +- Added `buildNamedToolChoice` utility function to build provider-aware tool choice constraints for named tools - Support for comma/space-separated path lists in `find`, `grep`, `ast_grep`, and `ast_edit` tools (e.g., `apps/,packages/,phases/` or `apps/ packages/ phases/`) - New `resolveMultiSearchPath` and `resolveMultiFindPattern` functions to handle multi-path search inputs with automatic common base path detection ### Changed +- Modified tool choice resolution to support per-turn tool choice overrides via `consumeNextToolChoiceOverride()` - Updated tool documentation to clarify that `path` parameter accepts files, directories, glob patterns, or comma/space-separated path lists - Refactored path resolution logic in `find`, `grep`, `ast_grep`, and `ast_edit` tools to use unified multi-path handling diff --git a/packages/coding-agent/src/config/settings-schema.ts b/packages/coding-agent/src/config/settings-schema.ts index 7b6c8312e..3228b6c5b 100644 --- a/packages/coding-agent/src/config/settings-schema.ts +++ b/packages/coding-agent/src/config/settings-schema.ts @@ -470,6 +470,15 @@ export const SETTINGS_SCHEMA = { submenu: true, }, }, + "todo.eager": { + type: "boolean", + default: false, + ui: { + tab: "agent", + label: "Eager todos", + description: "Automatically create a comprehensive todo list after the first message", + }, + }, // ───────────────────────────────────────────────────────────────────────── // Optional tools diff --git a/packages/coding-agent/src/prompts/system/eager-todo.md b/packages/coding-agent/src/prompts/system/eager-todo.md new file mode 100644 index 000000000..187ec9336 --- /dev/null +++ b/packages/coding-agent/src/prompts/system/eager-todo.md @@ -0,0 +1,16 @@ + +Create a comprehensive phased todo for the upcoming user request now. + +The todo **MUST** cover this request: + +{{userRequest}} + + +You **MUST** call `todo_write` in this turn. +You **MUST** initialize the todo list with a single `replace` op. +You **MUST** cover the entire request from investigation through implementation and verification — not just the next immediate step. +You **MUST** make task descriptions specific enough that a future turn can execute them without re-planning. +You **MUST** keep exactly one task `in_progress` and all later tasks `pending`. + +You **MUST NOT** output plain text in this turn. + \ No newline at end of file diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index ef5a6709f..541c9b269 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -1409,7 +1409,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} if (pendingActionStore.hasPending) { return { type: "function", name: "resolve" }; } - return undefined; + return session?.consumeNextToolChoiceOverride(); }, }); cursorEventEmitter = event => agent.emitExternalEvent(event); diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index f8a574f7e..1215c6b96 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -89,6 +89,7 @@ import { getCurrentThemeName, theme } from "../modes/theme/theme"; import { normalizeDiff, normalizeToLF, ParseError, previewPatch, stripBom } from "../patch"; import type { PlanModeState } from "../plan-mode/state"; import autoHandoffThresholdFocusPrompt from "../prompts/system/auto-handoff-threshold-focus.md" with { type: "text" }; +import eagerTodoPrompt from "../prompts/system/eager-todo.md" with { type: "text" }; import handoffDocumentPrompt from "../prompts/system/handoff-document.md" with { type: "text" }; import planModeActivePrompt from "../prompts/system/plan-mode-active.md" with { type: "text" }; import planModeReferencePrompt from "../prompts/system/plan-mode-reference.md" with { type: "text" }; @@ -106,6 +107,7 @@ import { getLatestTodoPhasesFromEntries, type TodoItem, type TodoPhase } from ". import { parseCommandArgs } from "../utils/command-args"; import { resolveFileDisplayMode } from "../utils/file-display-mode"; import { extractFileMentions, generateFileMentionMessages } from "../utils/file-mentions"; +import { buildNamedToolChoice } from "../utils/tool-choice"; import { type CompactionResult, calculateContextTokens, @@ -350,6 +352,9 @@ export class AgentSession { // Todo completion reminder state #todoReminderCount = 0; #todoPhases: TodoPhase[] = []; + #eagerTodoInjected = false; + #skipNextTodoCompletionReminder = false; + #nextToolChoiceOverride: ToolChoice | undefined = undefined; // Bash execution state #bashAbortController: AbortController | undefined = undefined; @@ -457,6 +462,12 @@ export class AgentSession { return this.#modelRegistry; } + consumeNextToolChoiceOverride(): ToolChoice | undefined { + const toolChoice = this.#nextToolChoiceOverride; + this.#nextToolChoiceOverride = undefined; + return toolChoice; + } + /** Provider-scoped mutable state store for transport/session caches. */ get providerSessionState(): Map { return this.#providerSessionState; @@ -791,7 +802,15 @@ export class AgentSession { const compactionTask = this.#checkCompaction(msg); this.#trackPostPromptTask(compactionTask); await compactionTask; - // Check for incomplete todos (unless there was an error or abort) + // Check for incomplete todos only after a final assistant stop, not intermediate tool-use turns. + const hasToolCalls = msg.content.some(content => content.type === "toolCall"); + if (hasToolCalls) { + return; + } + if (this.#skipNextTodoCompletionReminder) { + this.#skipNextTodoCompletionReminder = false; + return; + } if (msg.stopReason !== "error" && msg.stopReason !== "aborted") { if (this.#enforceRewindBeforeYield()) { return; @@ -1349,6 +1368,9 @@ export class AgentSession { if (!this.#extensionRunner) return; if (event.type === "agent_start") { this.#turnIndex = 0; + this.#eagerTodoInjected = false; + this.#skipNextTodoCompletionReminder = false; + this.#nextToolChoiceOverride = undefined; await this.#extensionRunner.emit({ type: "agent_start" }); } else if (event.type === "agent_end") { await this.#extensionRunner.emit({ type: "agent_end", messages: event.messages }); @@ -1945,6 +1967,10 @@ export class AgentSession { return; } + if (!options?.synthetic) { + await this.#enforceEagerTodo(expandedText); + } + const userContent: (TextContent | ImageContent)[] = [{ type: "text", text: expandedText }]; if (options?.images) { userContent.push(...options.images); @@ -3481,6 +3507,57 @@ export class AgentSession { this.agent.setTools(previousTools); } } + + async #enforceEagerTodo(userRequest: string): Promise { + if (this.#eagerTodoInjected) { + return; + } + const eagerTodosEnabled = this.settings.get("todo.eager"); + const todosEnabled = this.settings.get("todo.enabled"); + if (!eagerTodosEnabled || !todosEnabled) { + return; + } + + this.#eagerTodoInjected = true; + + if (this.#planModeState?.enabled) { + return; + } + if (this.getTodoPhases().length > 0) { + return; + } + + if (!this.#toolRegistry.has("todo_write")) { + logger.warn("Eager todo enforcement skipped because todo_write is unavailable", { + activeToolNames: this.agent.state.tools.map(tool => tool.name), + }); + return; + } + + const todoWriteToolChoice = buildNamedToolChoice("todo_write", this.model); + if (!todoWriteToolChoice) { + logger.warn("Eager todo enforcement skipped because the current model does not support forcing todo_write", { + modelApi: this.model?.api, + modelId: this.model?.id, + }); + return; + } + + const eagerTodoReminder = renderPromptTemplate(eagerTodoPrompt, { + userRequest, + }); + + this.#nextToolChoiceOverride = todoWriteToolChoice; + this.#skipNextTodoCompletionReminder = true; + try { + await this.prompt(eagerTodoReminder, { + synthetic: true, + expandPromptTemplates: false, + }); + } finally { + this.#nextToolChoiceOverride = undefined; + } + } /** * Check if agent stopped with incomplete todos and prompt to continue. */ @@ -5191,8 +5268,8 @@ export class AgentSession { } for (const msg of this.messages) { - if (msg.role === "user") { - lines.push("## User\n"); + if (msg.role === "user" || msg.role === "developer") { + lines.push(msg.role === "developer" ? "## Developer\n" : "## User\n"); if (typeof msg.content === "string") { lines.push(msg.content); } else { @@ -5316,8 +5393,8 @@ export class AgentSession { lines.push(""); for (const msg of this.messages) { - if (msg.role === "user") { - lines.push("## User"); + if (msg.role === "user" || msg.role === "developer") { + lines.push(msg.role === "developer" ? "## Developer" : "## User"); lines.push(""); if (typeof msg.content === "string") { lines.push(msg.content); diff --git a/packages/coding-agent/src/task/executor.ts b/packages/coding-agent/src/task/executor.ts index 45454c5e0..5f976e031 100644 --- a/packages/coding-agent/src/task/executor.ts +++ b/packages/coding-agent/src/task/executor.ts @@ -5,7 +5,6 @@ */ import path from "node:path"; import type { AgentEvent, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; -import type { Api, Model, ToolChoice } from "@oh-my-pi/pi-ai"; import { logger, untilAborted } from "@oh-my-pi/pi-utils"; import type { TSchema } from "@sinclair/typebox"; import Ajv, { type ValidateFunction } from "ajv"; @@ -28,6 +27,7 @@ import { type ContextFileEntry, truncateTail } from "../tools"; import { jtdToJsonSchema } from "../tools/jtd-to-json-schema"; import { ToolAbortError } from "../tools/tool-errors"; import type { EventBus } from "../utils/event-bus"; +import { buildNamedToolChoice } from "../utils/tool-choice"; import { subprocessToolRegistry } from "./subprocess-tool-registry"; import { type AgentDefinition, @@ -117,22 +117,6 @@ function getReportFindingKey(value: unknown): string | null { return `${filePath}:${lineStart}:${lineEnd}:${priority ?? ""}:${title}`; } -function buildSubmitResultToolChoice(model?: Model): ToolChoice | undefined { - if (!model) return undefined; - if ( - model.api === "openai-codex-responses" || - model.api === "openai-responses" || - model.api === "openai-completions" || - model.api === "azure-openai-responses" - ) { - return { type: "function", name: "submit_result" }; - } - if (model.api === "anthropic-messages" || model.api === "bedrock-converse-stream") { - return { type: "tool", name: "submit_result" }; - } - return undefined; -} - /** Options for subagent execution */ export interface ExecutorOptions { cwd: string; @@ -1091,7 +1075,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise): ToolChoice | undefined { + if (!model) return undefined; + + if (model.api === "anthropic-messages" || model.api === "bedrock-converse-stream") { + return { type: "tool", name: toolName }; + } + + if ( + model.api === "openai-codex-responses" || + model.api === "openai-responses" || + model.api === "openai-completions" || + model.api === "azure-openai-responses" + ) { + return { type: "function", name: toolName }; + } + + if (model.api === "google-generative-ai" || model.api === "google-gemini-cli" || model.api === "google-vertex") { + return "required"; + } + + return undefined; +} diff --git a/packages/coding-agent/test/agent-session-eager-todo.test.ts b/packages/coding-agent/test/agent-session-eager-todo.test.ts new file mode 100644 index 000000000..f94ba3a8f --- /dev/null +++ b/packages/coding-agent/test/agent-session-eager-todo.test.ts @@ -0,0 +1,262 @@ +import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import * as path from "node:path"; +import { Agent, type AgentMessage, type AgentTool } from "@oh-my-pi/pi-agent-core"; +import { type AssistantMessage, getBundledModel, type TextContent, type ToolCall } from "@oh-my-pi/pi-ai"; +import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; +import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { AgentSession } from "@oh-my-pi/pi-coding-agent/session/agent-session"; +import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; +import { convertToLlm } from "@oh-my-pi/pi-coding-agent/session/messages"; +import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; +import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; +import { TodoWriteTool } from "@oh-my-pi/pi-coding-agent/tools"; +import { TempDir } from "@oh-my-pi/pi-utils"; +import { Type } from "@sinclair/typebox"; + +class MockAssistantStream extends AssistantMessageEventStream {} + +type ObservedPromptCall = { + toolChoice: string | undefined; + toolNames: string[]; + lastMessageRole: AgentMessage["role"]; + lastMessageText: string; +}; + +function isTextContentBlock(value: unknown): value is TextContent { + if (!value || typeof value !== "object") return false; + return (value as TextContent).type === "text" && typeof (value as TextContent).text === "string"; +} + +function getToolChoiceName(choice: unknown): string | undefined { + if (!choice) return undefined; + if (typeof choice === "string") return choice; + if (typeof choice !== "object" || !("type" in choice)) return undefined; + const toolChoice = choice as { type?: string; name?: string; function?: { name?: string } }; + if (toolChoice.type === "tool") return toolChoice.name; + if (toolChoice.type === "function") return toolChoice.name ?? toolChoice.function?.name; + return undefined; +} + +function createAssistantMessage(text: string): AssistantMessage { + return { + role: "assistant", + content: [{ type: "text", text }], + api: "anthropic-messages", + provider: "anthropic", + model: "mock", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: Date.now(), + }; +} + +function createToolCallAssistantMessage(name: string, args: Record): AssistantMessage { + const toolCall: ToolCall = { + type: "toolCall", + id: `call_${name}`, + name, + arguments: args, + }; + return { + role: "assistant", + content: [toolCall], + api: "anthropic-messages", + provider: "anthropic", + model: "mock", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "toolUse", + timestamp: Date.now(), + }; +} + +function getMessageText(message: AgentMessage): string { + if (!("content" in message)) { + return ""; + } + if (typeof message.content === "string") { + return message.content; + } + if (!Array.isArray(message.content)) { + return ""; + } + return message.content + .filter(isTextContentBlock) + .map(content => content.text) + .join("\n"); +} + +describe("AgentSession eager todo enforcement", () => { + let tempDir: TempDir; + let session: AgentSession; + let streamCallCount = 0; + let scriptedResponses: AssistantMessage[] = []; + const observedCalls: ObservedPromptCall[] = []; + + beforeEach(async () => { + tempDir = TempDir.createSync("@pi-agent-session-eager-todo-"); + streamCallCount = 0; + scriptedResponses = []; + observedCalls.length = 0; + + const model = getBundledModel("anthropic", "claude-sonnet-4-5"); + if (!model) throw new Error("Expected claude-sonnet-4-5 model to exist"); + + const authStorage = await AuthStorage.create(path.join(tempDir.path(), "testauth.db")); + authStorage.setRuntimeApiKey("anthropic", "test-key"); + const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir.path(), "models.yml")); + const settings = Settings.isolated({ + "compaction.enabled": false, + "todo.enabled": true, + "todo.eager": true, + }); + const sessionManager = SessionManager.inMemory(tempDir.path()); + + const toolSession: ToolSession = { + cwd: tempDir.path(), + hasUI: false, + getSessionFile: () => sessionManager.getSessionFile() ?? null, + getSessionSpawns: () => "*", + settings, + }; + const todoWriteTool = new TodoWriteTool(toolSession); + const mockBashTool: AgentTool = { + name: "bash", + label: "Bash", + description: "Mock bash tool", + parameters: Type.Object({}), + execute: async () => ({ content: [{ type: "text" as const, text: "ok" }] }), + }; + + const agent = new Agent({ + getApiKey: () => "test-key", + initialState: { + model, + systemPrompt: "Test", + tools: [todoWriteTool, mockBashTool], + messages: [], + }, + convertToLlm, + getToolChoice: () => session?.consumeNextToolChoiceOverride(), + streamFn: (_model, context, options) => { + streamCallCount++; + const lastMessage = context.messages.at(-1); + if (!lastMessage) { + throw new Error("Expected prompt context to include a message"); + } + observedCalls.push({ + toolChoice: getToolChoiceName(options?.toolChoice), + toolNames: (context.tools ?? []).map(tool => tool.name), + lastMessageRole: lastMessage.role, + lastMessageText: getMessageText(lastMessage), + }); + const response = scriptedResponses.shift() ?? createAssistantMessage("done"); + const stream = new MockAssistantStream(); + queueMicrotask(() => { + stream.push({ type: "start", partial: response }); + const reason = response.stopReason === "toolUse" || response.stopReason === "length" ? response.stopReason : "stop"; + stream.push({ type: "done", reason, message: response }); + }); + return stream; + }, + }); + + const toolRegistry = new Map([ + [todoWriteTool.name, todoWriteTool as unknown as AgentTool], + [mockBashTool.name, mockBashTool], + ]); + + session = new AgentSession({ + agent, + sessionManager, + settings, + modelRegistry, + toolRegistry, + }); + }); + + afterEach(async () => { + if (session) { + await session.dispose(); + } + tempDir.removeSync(); + }); + + it("forces a synthetic todo-only turn before the first real user prompt", async () => { + await session.prompt("list all work trees"); + + const dumpText = session.formatSessionAsText(); + + expect(observedCalls).toHaveLength(2); + expect(observedCalls[0]).toEqual({ + toolChoice: "todo_write", + toolNames: ["todo_write", "bash"], + lastMessageRole: "developer", + lastMessageText: expect.stringContaining("list all work trees"), + }); + expect(observedCalls[1]).toEqual({ + toolChoice: undefined, + toolNames: ["todo_write", "bash"], + lastMessageRole: "user", + lastMessageText: "list all work trees", + }); + expect(dumpText).toContain("## Developer"); + expect(dumpText).toContain("Create a comprehensive phased todo for the upcoming user request now."); + }); + + it("initializes todos once, then continues to the real user turn without looping todo_write", async () => { + scriptedResponses = [ + createToolCallAssistantMessage("todo_write", { + ops: [ + { + op: "replace", + phases: [ + { + name: "List worktrees", + tasks: [ + { content: "List all git worktrees in the current repository", status: "in_progress" }, + ], + }, + ], + }, + ], + }), + createAssistantMessage("todo initialized"), + createAssistantMessage("real user turn handled"), + ]; + + await session.prompt("list all work trees"); + + expect(streamCallCount).toBe(3); + expect(observedCalls).toHaveLength(3); + expect(observedCalls[0]).toEqual({ + toolChoice: "todo_write", + toolNames: ["todo_write", "bash"], + lastMessageRole: "developer", + lastMessageText: expect.stringContaining("list all work trees"), + }); + expect(observedCalls[1]?.lastMessageRole).toBe("toolResult"); + expect(observedCalls[1]?.toolChoice).toBeUndefined(); + expect(observedCalls[2]).toEqual({ + toolChoice: undefined, + toolNames: ["todo_write", "bash"], + lastMessageRole: "user", + lastMessageText: "list all work trees", }); + expect(session.getTodoPhases()).toHaveLength(1); + expect(session.getTodoPhases()[0]?.tasks[0]?.content).toBe("List all git worktrees in the current repository"); + }); +});