diff --git a/packages/agent/CHANGELOG.md b/packages/agent/CHANGELOG.md index d90fbcf3a..0566aba54 100644 --- a/packages/agent/CHANGELOG.md +++ b/packages/agent/CHANGELOG.md @@ -4,7 +4,7 @@ ### Changed -- `beforeToolCall` now runs at arg-prep time, per call in batch order before the call is chained into the schedule — ahead of concurrency resolution, `tool_execution_start`, and telemetry span start — instead of inside the already-scheduled execution slot. It receives the resolved `tool` in its context and may return `args` to replace the call's arguments; a replacement is revalidated against the tool schema, re-resolves argument-dependent interruptibility, and becomes the single source of truth for scheduling, execution events, the persisted assistant message, and `tool.execute`. Argument validation moved into the same prepare phase, so functional `concurrency` resolvers now see validated (and possibly revised) arguments rather than raw pre-validation ones. The steering watch is installed before the prepare loop so a hung hook cannot stall interrupt detection. +- `beforeToolCall` now runs during arg-prep in a pre-dispatch prepare phase — on the streamed path before the assistant message's `message_start`/`message_end` are emitted, and always ahead of concurrency resolution, `tool_execution_start`, telemetry span start, and `tool.execute` — instead of inside the already-scheduled execution slot. It receives the resolved `tool` in its context and may return `args` to replace the call's arguments; a replacement is revalidated against the tool schema, written back to the assistant message's tool-call block, and re-resolves argument-dependent interruptibility, making it the single source of truth for history, persistence, provider replay, scheduling, execution events, and `tool.execute`. Argument validation moved into the same prepare phase, so functional `concurrency` resolvers now see validated (and possibly revised) arguments rather than raw pre-validation ones. The hook now receives the run's request abort signal rather than the per-tool signal. ## [17.1.6] - 2026-07-27 diff --git a/packages/agent/src/agent-loop.ts b/packages/agent/src/agent-loop.ts index 763dbf9bd..dd6a86622 100644 --- a/packages/agent/src/agent-loop.ts +++ b/packages/agent/src/agent-loop.ts @@ -70,6 +70,7 @@ import type { AgentMessage, AgentPreModelCallResult, AgentTool, + AgentToolCall, AgentToolResult, AgentTurnEndContext, AsideMessage, @@ -1753,6 +1754,17 @@ async function streamAssistantResponse( if (config.transformAssistantMessage) { await config.transformAssistantMessage(finalMessage, requestSignal); } + // Prepare tool dispatch (validation + the `beforeToolCall` hook) + // BEFORE the message is snapshotted for consumers: a hook args + // revision is written back into this message's toolCall blocks, + // so history, the UI, persistence, provider replay, scheduling, + // and execution all carry the revised arguments. + if (finalMessage.content.some(c => c.type === "toolCall")) { + preparedDispatchByMessage.set( + finalMessage, + await prepareToolCallDispatch(finalMessage, context, config, requestSignal), + ); + } if (addedPartial) { context.messages[context.messages.length - 1] = finalMessage; } else { @@ -2052,6 +2064,136 @@ function emitAbortedAssistantMessage( return abortedMessage; } +/** Per-call outcome of the pre-dispatch prepare phase (validation + `beforeToolCall`). */ +interface PreparedToolCall { + tool: AgentTool | undefined; + /** Validated (possibly hook-revised) execution args; raw args when validation failed. */ + args: Record; + validationErrorMessage?: string; + blocked?: boolean; + blockReason?: string; + prepareError?: unknown; +} + +/** + * Prepare results computed in the stream-done branch (before `message_start`/ + * `message_end`) so a `beforeToolCall` args revision is baked into the message + * every consumer snapshots. `executeToolCalls` consumes them; a message that + * bypassed the streamed path (e.g. Harmony-recovered) is prepared at dispatch + * time instead. + */ +const preparedDispatchByMessage = new WeakMap>(); + +function resolveToolForCall( + tools: AgentTool[] | undefined, + toolCall: AgentToolCall, + resolveFallbackTool: AgentLoopConfig["resolveFallbackTool"], +): AgentTool | undefined { + // Tools emitted via OpenAI's custom-tool path (e.g. `apply_patch` on GPT-5) + // come back under their wire-level name, which may differ from the + // harness-internal `name`. Match on either, preferring `name` for + // determinism if both somehow collide. + return ( + tools?.find(t => t.name === toolCall.name) ?? + tools?.find(t => t.customWireName !== undefined && t.customWireName === toolCall.name) ?? + // Not in the advertised set: let the host route side-transport tools + // (e.g. xd:// device mounts) called by their top-level name. + resolveFallbackTool?.(toolCall.name) + ); +} + +/** + * Pre-dispatch phase for every pending tool call on `assistantMessage`, run in + * call order: intent extraction, argument validation, and the `beforeToolCall` + * hook. A hook `args` revision is revalidated against the tool schema and + * written back to `toolCall.arguments`; run before `message_start`/`message_end` + * (the streamed path) that makes the revision the single source of truth — + * history, execution events, persistence, provider replay, concurrency + * scheduling, and `tool.execute` all agree. Failures are recorded per call and + * surfaced by `executeToolCalls` at the record's scheduled slot. + */ +async function prepareToolCallDispatch( + assistantMessage: AssistantMessage, + context: AgentContext, + config: AgentLoopConfig, + signal: AbortSignal | undefined, +): Promise> { + const { resolveFallbackTool, intentTracing, beforeToolCall } = config; + const prepared = new Map(); + for (const toolCall of assistantMessage.content) { + if (toolCall.type !== "toolCall") continue; + if ((toolCall as CursorExecResolvedCarrier)[kCursorExecResolved] === true) continue; + const tool = resolveToolForCall(context.tools, toolCall, resolveFallbackTool); + const entry: PreparedToolCall = { tool, args: toolCall.arguments as Record }; + prepared.set(toolCall.id, entry); + let argsForExecution = toolCall.arguments as Record; + if (intentTracing) { + const { intent, strippedArgs } = extractIntent(toolCall.arguments); + argsForExecution = strippedArgs; + if (intent) { + toolCall.intent = intent; + } else if (typeof tool?.intent === "function") { + try { + const derived = tool.intent(strippedArgs as never)?.trim(); + if (derived) { + toolCall.intent = derived; + } + } catch { + // intent function must never break tool execution + } + } + } + const validate = (args: Record): Record | undefined => { + try { + if (!tool) throw new Error(`Tool ${toolCall.name} not found`); + return validateToolArguments(tool, { ...toolCall, arguments: args }); + } catch (validationError) { + if (tool?.lenientArgValidation) { + const fallback = { ...args }; + delete fallback.__parseError; + delete fallback.__rawJson; + return fallback; + } + entry.args = "__parseError" in args ? { __parseError: args.__parseError } : args; + entry.validationErrorMessage = + validationError instanceof Error ? validationError.message : String(validationError); + return undefined; + } + }; + const effectiveArgs = validate(argsForExecution); + if (effectiveArgs === undefined) continue; + entry.args = effectiveArgs; + if (!beforeToolCall || !tool) continue; + let beforeResult: BeforeToolCallResult | undefined; + try { + beforeResult = await beforeToolCall( + { assistantMessage, toolCall, tool, args: effectiveArgs, context }, + signal, + ); + } catch (e) { + // Contract: a throwing hook surfaces as a tool-error result without + // aborting the batch — rethrown inside the execution span in runTool. + entry.prepareError = e; + continue; + } + if (beforeResult?.block) { + entry.blocked = true; + entry.blockReason = beforeResult.reason; + continue; + } + if (beforeResult?.args !== undefined) { + // Revalidate: a hook revision is untrusted input to the tool schema. + const revised = validate(beforeResult.args); + if (revised === undefined) continue; + // Bake the revision into the message itself. On the streamed path this + // precedes every consumer snapshot, so there is exactly one version of + // the call anywhere downstream. + toolCall.arguments = beforeResult.args; + entry.args = revised; + } + } + return prepared; +} /** * Execute tool calls from an assistant message. */ @@ -2072,8 +2214,6 @@ async function executeToolCalls( getToolContext, transformToolCallArguments, resolveFallbackTool, - intentTracing, - beforeToolCall, afterToolCall, } = config; type ToolCallContent = Extract; @@ -2107,22 +2247,25 @@ async function executeToolCalls( : AbortSignal.any([steeringAbortController.signal, ircAbortController.signal]); const interruptState: { triggered: boolean; source?: SteeringInterruptSource | "irc" } = { triggered: false }; + // Streamed messages were prepared (validation + `beforeToolCall`) before + // `message_end`, so hook revisions are already part of the message; anything + // that bypassed the streamed path is prepared here instead. + const preparedDispatch = + preparedDispatchByMessage.get(assistantMessage) ?? + (await prepareToolCallDispatch(assistantMessage, currentContext, config, signal)); + const records = toolCalls.map(toolCall => { - // Tools emitted via OpenAI's custom-tool path (e.g. `apply_patch` on GPT-5) - // come back under their wire-level name, which may differ from the - // harness-internal `name`. Match on either, preferring `name` for - // determinism if both somehow collide. - const tool = - tools?.find(t => t.name === toolCall.name) ?? - tools?.find(t => t.customWireName !== undefined && t.customWireName === toolCall.name) ?? - // Not in the advertised set: let the host route side-transport tools - // (e.g. xd:// device mounts) called by their top-level name. - resolveFallbackTool?.(toolCall.name); - const args = toolCall.arguments as Record; + const prepared = preparedDispatch.get(toolCall.id) ?? { + tool: resolveToolForCall(tools, toolCall, resolveFallbackTool), + args: toolCall.arguments as Record, + }; + const { tool, args } = prepared; const interruptibleMode = tool?.interruptible; let interruptible = false; if (typeof interruptibleMode === "function") { try { + // Resolved from the prepared (possibly hook-revised) args so an + // argument-dependent policy governs the call that actually runs. interruptible = interruptibleMode(args); } catch { // Resolver failures default to preserving the tool's outcome. @@ -2143,10 +2286,10 @@ async function executeToolCalls( skipped: false, toolResultMessage: undefined as ToolResultMessage | undefined, resultEmitted: false, - validationErrorMessage: undefined as string | undefined, - blocked: false, - blockReason: undefined as string | undefined, - prepareError: undefined as unknown, + validationErrorMessage: prepared.validationErrorMessage, + blocked: prepared.blocked === true, + blockReason: prepared.blockReason, + prepareError: prepared.prepareError, }; }); @@ -2454,100 +2597,6 @@ async function executeToolCalls( await checkSteering(); }; - // Prepare phase, run per record in call order before the record is chained - // into the schedule: intent extraction, argument validation, and the - // `beforeToolCall` hook. A hook revision recorded here governs execution — - // concurrency resolution, `tool_execution_start`, telemetry, and - // `tool.execute` all observe it; the assistant message keeps the model's - // proposal (see the note inside the revision branch). - const prepareToolCall = async (record: (typeof records)[number]): Promise => { - // Interrupted/aborted records are settled by runTool and the tail sweep; - // running hooks for a call that will never execute would be misleading. - if (interruptState.triggered || record.signal.aborted) return; - const { toolCall, tool } = record; - let argsForExecution = toolCall.arguments as Record; - if (intentTracing) { - const { intent, strippedArgs } = extractIntent(toolCall.arguments); - argsForExecution = strippedArgs; - if (intent) { - toolCall.intent = intent; - } else if (typeof tool?.intent === "function") { - try { - const derived = tool.intent(strippedArgs as never)?.trim(); - if (derived) { - toolCall.intent = derived; - } - } catch { - // intent function must never break tool execution - } - } - } - const validate = (args: Record): Record | undefined => { - try { - if (!tool) throw new Error(`Tool ${toolCall.name} not found`); - return validateToolArguments(tool, { ...toolCall, arguments: args }); - } catch (validationError) { - if (tool?.lenientArgValidation) { - const fallback = { ...args }; - delete fallback.__parseError; - delete fallback.__rawJson; - return fallback; - } - record.args = "__parseError" in args ? { __parseError: args.__parseError } : args; - record.validationErrorMessage = - validationError instanceof Error ? validationError.message : String(validationError); - return undefined; - } - }; - const effectiveArgs = validate(argsForExecution); - if (effectiveArgs === undefined) return; - record.args = effectiveArgs; - if (!beforeToolCall || !tool) return; - let beforeResult: BeforeToolCallResult | undefined; - try { - beforeResult = await beforeToolCall( - { assistantMessage, toolCall, tool, args: effectiveArgs, context: currentContext }, - record.signal, - ); - } catch (e) { - // Contract: a throwing hook surfaces as a tool-error result without - // aborting the batch — rethrown inside the execution span in runTool. - record.prepareError = e; - return; - } - if (beforeResult?.block) { - record.blocked = true; - record.blockReason = beforeResult.reason; - return; - } - if (beforeResult?.args !== undefined) { - // Revalidate: a hook revision is untrusted input to the tool schema. - const revised = validate(beforeResult.args); - if (revised === undefined) return; - // Like `transformToolCallArguments`, the revision governs execution — - // scheduling, execution events, telemetry, and `tool.execute` — while - // the assistant message keeps the model's proposal. Consumers already - // received `message_end` snapshots of that message, so mutating it here - // would fork provider replay from persisted history. - record.args = revised; - // Interruptibility may be argument-dependent; re-resolve it so the - // revised call runs under the right abort signal. - const interruptibleMode = tool.interruptible; - let interruptible = false; - if (typeof interruptibleMode === "function") { - try { - interruptible = interruptibleMode(revised); - } catch { - interruptible = false; - } - } else { - interruptible = interruptibleMode === true; - } - record.interruptible = interruptible; - record.signal = interruptible ? interruptibleSignal : nonInterruptibleSignal; - } - }; - let lastExclusive: Promise = Promise.resolve(); let sharedTasks: Promise[] = []; const tasks: Promise[] = []; @@ -2557,9 +2606,7 @@ async function executeToolCalls( // detection hard-aborts interruptible waits, soft-signals cooperative tools // (auto-background bash), and skips not-yet-started tools, so the boundary // dequeue below injects the message promptly. Gated on immediate-interrupt - // mode; checkSteering is idempotent (no-op once triggered). Installed before - // the prepare/schedule loop so a hung `beforeToolCall` handler cannot stall - // steering detection. + // mode; checkSteering is idempotent (no-op once triggered). const watchSteeringWhileRunning = shouldInterruptImmediately && (hasSteeringMessages !== undefined || hasIrcInterrupts !== undefined); const eventDrivenSteeringWatch = @@ -2610,35 +2657,34 @@ async function executeToolCalls( STEERING_INTERRUPT_POLL_MS, ) : undefined; - try { - for (let index = 0; index < records.length; index++) { - const record = records[index]; - await prepareToolCall(record); - const concurrencyMode = record.tool?.concurrency; - let concurrency: "shared" | "exclusive"; - if (typeof concurrencyMode === "function") { - // Resolved from the validated (possibly hook-revised) args — raw - // args only when validation failed, and those records error out - // before executing. A throwing resolver must not take down the - // whole batch, so fall back to the safe (serial) mode. - try { - concurrency = concurrencyMode(record.args); - } catch { - concurrency = "exclusive"; - } - } else { - concurrency = concurrencyMode ?? "shared"; - } - const start = concurrency === "exclusive" ? Promise.all([lastExclusive, ...sharedTasks]) : lastExclusive; - const task = start.then(() => runTool(record, index)); - tasks.push(task); - if (concurrency === "exclusive") { - lastExclusive = task; - sharedTasks = []; - } else { - sharedTasks.push(task); + for (let index = 0; index < records.length; index++) { + const record = records[index]; + const concurrencyMode = record.tool?.concurrency; + let concurrency: "shared" | "exclusive"; + if (typeof concurrencyMode === "function") { + // Resolved from the prepared (possibly hook-revised) args — raw args + // only when validation failed, and those records error out before + // executing. A throwing resolver must not take down the whole batch, + // so fall back to the safe (serial) mode. + try { + concurrency = concurrencyMode(record.args); + } catch { + concurrency = "exclusive"; } + } else { + concurrency = concurrencyMode ?? "shared"; } + const start = concurrency === "exclusive" ? Promise.all([lastExclusive, ...sharedTasks]) : lastExclusive; + const task = start.then(() => runTool(record, index)); + tasks.push(task); + if (concurrency === "exclusive") { + lastExclusive = task; + sharedTasks = []; + } else { + sharedTasks.push(task); + } + } + try { await Promise.allSettled(tasks); } finally { steeringWatchAbortController.abort(); diff --git a/packages/agent/src/types.ts b/packages/agent/src/types.ts index 6c6a00107..51c3bbe1d 100644 --- a/packages/agent/src/types.ts +++ b/packages/agent/src/types.ts @@ -456,23 +456,22 @@ export interface AgentLoopConfig extends SimpleStreamOptions { getCwd?: () => string | undefined; /** - * Called once per tool call after argument validation, before the call is - * scheduled — ahead of concurrency resolution, `tool_execution_start`, and - * `tool.execute`. Hooks in a batch run in call order while earlier tools may - * already be executing. + * Called once per tool call after argument validation, in call order, before + * the call is scheduled — ahead of concurrency resolution, + * `tool_execution_start`, and `tool.execute`. On the streamed path it runs + * before the assistant message's `message_start`/`message_end` are emitted. * * Return `{ block: true }` to prevent execution. The loop emits an error tool * result instead (using `reason` as the error text, or a default if omitted). * * Return `{ args }` to replace the arguments the call runs with. The - * replacement is revalidated against the tool schema and governs execution: - * concurrency scheduling, execution events, telemetry spans, and - * `tool.execute` all see the revised arguments, while the assistant message - * keeps the model's original proposal (matching `transformToolCallArguments` - * semantics). Mutating `context.args` in place also survives into execution, - * but a returned `args` object wins. + * replacement is revalidated against the tool schema and written back to the + * tool-call block, making it the single source of truth: history, execution + * events, persistence, provider replay, concurrency scheduling, and + * `tool.execute` all see the revised arguments. Mutating `context.args` in + * place also survives into execution, but a returned `args` object wins. * - * The hook receives the tool abort signal (`signal`) and is responsible for + * The hook receives the run's request abort signal and is responsible for * honoring it. Throwing surfaces as a tool-error result and does not abort the * rest of the batch. */ @@ -558,8 +557,9 @@ export type AgentToolCall = Extract { expect(toolStart?.type === "tool_execution_start" && toolStart.args).toEqual({ value: "revised" }); const messages = await stream.result(); const assistant = messages.find(m => m.role === "assistant"); - const toolCallBlock = assistant?.role === "assistant" ? assistant.content.find(c => c.type === "toolCall") : undefined; + const toolCallBlock = + assistant?.role === "assistant" ? assistant.content.find(c => c.type === "toolCall") : undefined; expect(toolCallBlock?.type === "toolCall" && toolCallBlock.arguments).toEqual({ value: "revised" }); }); diff --git a/packages/coding-agent/test/agent-session-message-pipeline.test.ts b/packages/coding-agent/test/agent-session-message-pipeline.test.ts index af5fcfb8c..b7662a9a7 100644 --- a/packages/coding-agent/test/agent-session-message-pipeline.test.ts +++ b/packages/coding-agent/test/agent-session-message-pipeline.test.ts @@ -855,10 +855,16 @@ describe("AgentSession message pipeline", () => { queueMicrotask(() => { if (requests === 1) { const message = createAssistantMessage(""); - message.content = [ - { type: "toolCall", id: "call-revise-1", name: "bash", arguments: { command: "echo original" } }, - ]; + const toolCall = { + type: "toolCall", + id: "call-revise-1", + name: "bash", + arguments: { command: "echo original" }, + } as const; + message.content = [toolCall]; message.stopReason = "toolUse"; + stream.push({ type: "toolcall_start", contentIndex: 0, partial: message }); + stream.push({ type: "toolcall_end", contentIndex: 0, toolCall: toolCall as never, partial: message }); stream.push({ type: "done", reason: "toolUse", message }); } else { const message = createAssistantMessage("done"); diff --git a/packages/coding-agent/test/dbg-revision.test.ts b/packages/coding-agent/test/dbg-revision.test.ts deleted file mode 100644 index 61d2ba6fa..000000000 --- a/packages/coding-agent/test/dbg-revision.test.ts +++ /dev/null @@ -1,79 +0,0 @@ -import { expect, it } from "bun:test"; -import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; -import { type Api, type Model, type ModelSpec, registerCustomApi } from "@oh-my-pi/pi-ai"; -import { buildModel } from "@oh-my-pi/pi-catalog/build"; -import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; -import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; -import { createAgentSession, type ExtensionFactory } from "@oh-my-pi/pi-coding-agent/sdk"; -import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; -import { TempDir } from "@oh-my-pi/pi-utils"; -import { createAssistantMessage } from "./helpers/agent-session-setup"; - -it("debug revision", async () => { - using tempDir = TempDir.createSync("@pi-dbg-rev-"); - const api = "test-dbg-revision"; - let requests = 0; - const seenContexts: unknown[] = []; - registerCustomApi(api, (_model, context) => { - seenContexts.push(JSON.parse(JSON.stringify(context.messages))); - requests++; - const stream = new AssistantMessageEventStream(); - queueMicrotask(() => { - if (requests === 1) { - const message = createAssistantMessage(""); - const toolCall = { type: "toolCall", id: "call-revise-1", name: "bash", arguments: { command: "echo original" } } as const; - message.content = [toolCall]; - message.stopReason = "toolUse"; - stream.push({ type: "toolcall_start", contentIndex: 0, partial: message }); - stream.push({ type: "toolcall_end", contentIndex: 0, toolCall: toolCall as never, partial: message }); - stream.push({ type: "done", reason: "toolUse", message }); - } else { - const message = createAssistantMessage("done"); - stream.push({ type: "done", reason: "stop", message }); - } - }); - return stream; - }); - const model = buildModel({ - id: "local-dbg", name: "dbg", api, provider: "ollama", baseUrl: "http://127.0.0.1:11434", - reasoning: false, input: ["text"], cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, - contextWindow: 4096, maxTokens: 1024, - } as ModelSpec) as Model; - let handlerCalls = 0; - const reviseBash: ExtensionFactory = pi => { - pi.on("tool_call", async event => { - handlerCalls++; - console.log("HANDLER tool_call", event.toolName, JSON.stringify(event.input)); - if (event.toolName !== "bash") return undefined; - return { input: { command: "echo revised" } }; - }); - }; - const authStorage = await AuthStorage.create(tempDir.join("auth.db")); - const modelRegistry = new ModelRegistry(authStorage, tempDir.join("models.yml")); - const { session } = await createAgentSession({ - cwd: tempDir.path(), agentDir: tempDir.path(), sessionManager: SessionManager.inMemory(tempDir.path()), - authStorage, modelRegistry, - settings: Settings.isolated({ "compaction.enabled": false, "bash.autoBackground.enabled": false, "bashInterceptor.enabled": false }), - model, disableExtensionDiscovery: true, extensions: [reviseBash], skills: [], contextFiles: [], - promptTemplates: [], slashCommands: [], enableMCP: false, enableLsp: false, skipPythonPreflight: true, - toolNames: ["bash"], - }); - try { - console.log("hasHandlers tool_call:", session.extensionRunner?.hasHandlers("tool_call")); - console.log("agent.beforeToolCall set:", typeof session.agent.beforeToolCall); - session.subscribe(ev => { - if (ev.type.startsWith("tool_")) console.log("EVENT", ev.type, JSON.stringify((ev as any).args ?? "")); - }); - await session.sendUserMessage("run it"); - console.log("handlerCalls:", handlerCalls); - console.log("REQ2 CONTEXT", JSON.stringify(seenContexts[1]).slice(0, 400)); - for (const m of session.agent.state.messages) { - console.log("MSG", m.role, JSON.stringify((m as any).content).slice(0, 200)); - } - } finally { - await session.dispose(); - authStorage.close(); - } - expect(true).toBe(true); -}, 30000);