feat(agent): restructured tool call dispatch to validate arguments earlier

- Added `prepareToolCallDispatch` and `PreparedToolCall` to handle argument validation and `beforeToolCall` before message snapshotting.
- Implemented `preparedDispatchByMessage` WeakMap to store pre-dispatch results for streamed messages.
- Updated `executeToolCalls` to consume pre-computed dispatch preparation results.
- Updated documentation and changelog to specify `beforeToolCall` timing on the streamed path.
This commit is contained in:
can1357
2026-07-27 23:01:50 +02:00
parent da6d11de0e
commit daeb683528
6 changed files with 212 additions and 238 deletions
+1 -1
View File
@@ -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
+187 -141
View File
@@ -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<any> | undefined;
/** Validated (possibly hook-revised) execution args; raw args when validation failed. */
args: Record<string, unknown>;
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<AssistantMessage, Map<string, PreparedToolCall>>();
function resolveToolForCall(
tools: AgentTool<any>[] | undefined,
toolCall: AgentToolCall,
resolveFallbackTool: AgentLoopConfig["resolveFallbackTool"],
): AgentTool<any> | 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<Map<string, PreparedToolCall>> {
const { resolveFallbackTool, intentTracing, beforeToolCall } = config;
const prepared = new Map<string, PreparedToolCall>();
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<string, unknown> };
prepared.set(toolCall.id, entry);
let argsForExecution = toolCall.arguments as Record<string, unknown>;
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<string, unknown>): Record<string, unknown> | 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<AssistantMessage["content"][number], { type: "toolCall" }>;
@@ -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<string, unknown>;
const prepared = preparedDispatch.get(toolCall.id) ?? {
tool: resolveToolForCall(tools, toolCall, resolveFallbackTool),
args: toolCall.arguments as Record<string, unknown>,
};
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<void> => {
// 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<string, unknown>;
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<string, unknown>): Record<string, unknown> | 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<void> = Promise.resolve();
let sharedTasks: Promise<void>[] = [];
const tasks: Promise<void>[] = [];
@@ -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();
+13 -13
View File
@@ -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<AssistantMessage["content"][number], { type:
* result instead, using `reason` as the error text (or a default if omitted).
*
* Set `args` to replace the tool-call arguments. The replacement is revalidated
* against the tool schema (a failure surfaces as a validation-error tool result)
* and is seen by scheduling, execution events, and `tool.execute` alike. It is
* against the tool schema (a failure surfaces as a validation-error tool result),
* written back to the tool-call block on the assistant message, and seen by
* history, scheduling, execution events, and `tool.execute` alike. It is
* ignored when `block` is true.
*/
export interface BeforeToolCallResult {
+2 -1
View File
@@ -3796,7 +3796,8 @@ describe("agentLoopContinue with AgentMessage", () => {
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" });
});
@@ -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");
@@ -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<Api>) as Model<Api>;
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);