From 821ec9570dcf6bcd6bbb4a067351b8da867a50dc Mon Sep 17 00:00:00 2001 From: can1357 Date: Wed, 8 Apr 2026 06:23:27 +0200 Subject: [PATCH] feat(coding-agent): enabled bidirectional RPC tool execution with host-owned custom tools - Added RPC host-owned custom tools framework enabling bidirectional tool execution between RPC client and host. - Added `set_host_tools` RPC command and `refreshRpcHostTools()` method for dynamic tool registry management. - Added RpcHostToolBridge and RpcHostToolAdapter classes with abort signal handling for cancellable tool execution. - Added `defineRpcClientTool` helper and customTools option to RpcClient for embedding host tool definitions. - Added comprehensive test suite covering RPC host tool execution, cancellation, and client-host communication. --- packages/coding-agent/CHANGELOG.md | 6 + packages/coding-agent/src/modes/index.ts | 22 +- .../coding-agent/src/modes/rpc/host-tools.ts | 193 ++++++++++++++ .../coding-agent/src/modes/rpc/rpc-client.ts | 189 +++++++++++++- .../coding-agent/src/modes/rpc/rpc-mode.ts | 51 +++- .../coding-agent/src/modes/rpc/rpc-types.ts | 47 +++- .../coding-agent/src/session/agent-session.ts | 44 ++++ .../coding-agent/test/rpc-host-tools.test.ts | 236 ++++++++++++++++++ 8 files changed, 771 insertions(+), 17 deletions(-) create mode 100644 packages/coding-agent/src/modes/rpc/host-tools.ts create mode 100644 packages/coding-agent/test/rpc-host-tools.test.ts diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 06350ca71..0080c14a6 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Breaking Changes - Simplified chunk edit operations: removed `append_child`, `prepend_child`, `append_sibling`, `prepend_sibling`, and `replace_body` ops in favor of unified `replace`, `before`, `after`, `prepend`, and `append` with region targeting (`@container`, `@prologue`, `@body`, `@epilogue`) @@ -9,6 +10,11 @@ ### Added +- Host-owned custom tools support: RPC clients can now register custom tools via `setCustomTools()` and the RPC server will invoke them over the transport with `host_tool_call` requests +- RPC host tool framework: `RpcHostToolBridge` for managing host tool execution, `RpcHostToolDefinition` for tool metadata, and bidirectional `host_tool_call`, `host_tool_cancel`, `host_tool_update`, and `host_tool_result` frames +- RPC client tool API: `defineRpcClientTool()` helper, `RpcClientCustomTool` interface, and `RpcClientToolContext` for implementing host-side tool execution with update streaming and abort support +- `set_host_tools` RPC command to replace the active set of host-owned tools before the next model call +- `refreshRpcHostTools()` method on `AgentSession` to integrate host tools into the active tool registry with conflict detection and auto-activation of non-hidden tools - Instruction breakpoints support: `set_instruction_breakpoint` and `remove_instruction_breakpoint` debug actions for setting breakpoints at specific instruction addresses - Data breakpoints support: `data_breakpoint_info`, `set_data_breakpoint`, and `remove_data_breakpoint` debug actions for monitoring variable/memory access - Memory introspection: `read_memory` and `write_memory` debug actions for inspecting and modifying debugger memory diff --git a/packages/coding-agent/src/modes/index.ts b/packages/coding-agent/src/modes/index.ts index 143261ff1..9b2726d68 100644 --- a/packages/coding-agent/src/modes/index.ts +++ b/packages/coding-agent/src/modes/index.ts @@ -7,9 +7,27 @@ import { postmortem } from "@oh-my-pi/pi-utils"; export { runAcpMode } from "./acp"; export { InteractiveMode, type InteractiveModeOptions } from "./interactive-mode"; export { type PrintModeOptions, runPrintMode } from "./print-mode"; -export { type ModelInfo, RpcClient, type RpcClientOptions, type RpcEventListener } from "./rpc/rpc-client"; +export { + defineRpcClientTool, + type ModelInfo, + RpcClient, + type RpcClientCustomTool, + type RpcClientOptions, + type RpcClientToolContext, + type RpcClientToolResult, + type RpcEventListener, +} from "./rpc/rpc-client"; export { runRpcMode } from "./rpc/rpc-mode"; -export type { RpcCommand, RpcResponse, RpcSessionState } from "./rpc/rpc-types"; +export type { + RpcCommand, + RpcHostToolCallRequest, + RpcHostToolCancelRequest, + RpcHostToolDefinition, + RpcHostToolResult, + RpcHostToolUpdate, + RpcResponse, + RpcSessionState, +} from "./rpc/rpc-types"; postmortem.register("terminal-restore", () => { emergencyTerminalRestore(); diff --git a/packages/coding-agent/src/modes/rpc/host-tools.ts b/packages/coding-agent/src/modes/rpc/host-tools.ts new file mode 100644 index 000000000..7a45500d2 --- /dev/null +++ b/packages/coding-agent/src/modes/rpc/host-tools.ts @@ -0,0 +1,193 @@ +import type { AgentTool, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; +import { Snowflake } from "@oh-my-pi/pi-utils"; +import type { Static, TSchema } from "@sinclair/typebox"; +import { applyToolProxy } from "../../extensibility/tool-proxy"; +import type { Theme } from "../../modes/theme/theme"; +import type { + RpcHostToolCallRequest, + RpcHostToolCancelRequest, + RpcHostToolDefinition, + RpcHostToolResult, + RpcHostToolUpdate, +} from "./rpc-types"; + +type RpcHostToolOutput = (frame: RpcHostToolCallRequest | RpcHostToolCancelRequest) => void; + +type PendingHostToolCall = { + resolve: (result: AgentToolResult) => void; + reject: (error: Error) => void; + onUpdate?: AgentToolUpdateCallback; +}; + +function _createErrorToolResult(message: string): AgentToolResult { + return { + content: [{ type: "text", text: message }], + details: {}, + }; +} + +function isAgentToolResult(value: unknown): value is AgentToolResult { + if (!value || typeof value !== "object") return false; + const content = (value as { content?: unknown }).content; + return Array.isArray(content); +} + +export function isRpcHostToolResult(value: unknown): value is RpcHostToolResult { + if (!value || typeof value !== "object") return false; + const frame = value as { type?: unknown; id?: unknown; result?: unknown }; + return frame.type === "host_tool_result" && typeof frame.id === "string" && isAgentToolResult(frame.result); +} + +export function isRpcHostToolUpdate(value: unknown): value is RpcHostToolUpdate { + if (!value || typeof value !== "object") return false; + const frame = value as { type?: unknown; id?: unknown; partialResult?: unknown }; + return frame.type === "host_tool_update" && typeof frame.id === "string" && isAgentToolResult(frame.partialResult); +} + +class RpcHostToolAdapter + implements AgentTool +{ + declare name: string; + declare label: string; + declare description: string; + declare parameters: TParams; + readonly strict = true; + concurrency: "shared" | "exclusive" = "shared"; + #bridge: RpcHostToolBridge; + #definition: RpcHostToolDefinition; + + constructor(definition: RpcHostToolDefinition, bridge: RpcHostToolBridge) { + this.#definition = definition; + this.#bridge = bridge; + applyToolProxy(definition, this); + } + + execute( + toolCallId: string, + params: Static, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + ): Promise> { + return this.#bridge.requestExecution( + this.#definition, + toolCallId, + params as Record, + signal, + onUpdate, + ); + } +} + +export class RpcHostToolBridge { + #output: RpcHostToolOutput; + #definitions = new Map(); + #pendingCalls = new Map(); + + constructor(output: RpcHostToolOutput) { + this.#output = output; + } + + getToolNames(): string[] { + return Array.from(this.#definitions.keys()); + } + + setTools(tools: RpcHostToolDefinition[]): AgentTool[] { + this.#definitions = new Map(tools.map(tool => [tool.name, tool])); + return tools.map(tool => new RpcHostToolAdapter(tool, this)); + } + + handleResult(frame: RpcHostToolResult): boolean { + const pending = this.#pendingCalls.get(frame.id); + if (!pending) return false; + this.#pendingCalls.delete(frame.id); + if (frame.isError) { + const text = frame.result.content + .filter( + (item): item is { type: "text"; text: string } => item.type === "text" && typeof item.text === "string", + ) + .map(item => item.text) + .join("\n") + .trim(); + pending.reject(new Error(text || "Host tool execution failed")); + return true; + } + pending.resolve(frame.result); + return true; + } + + handleUpdate(frame: RpcHostToolUpdate): boolean { + const pending = this.#pendingCalls.get(frame.id); + if (!pending) return false; + pending.onUpdate?.(frame.partialResult); + return true; + } + + requestExecution( + definition: RpcHostToolDefinition, + toolCallId: string, + args: Record, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + ): Promise> { + if (signal?.aborted) { + return Promise.reject(new Error(`Host tool "${definition.name}" was aborted`)); + } + + const id = Snowflake.next() as string; + const { promise, resolve, reject } = Promise.withResolvers>(); + let settled = false; + + const cleanup = () => { + signal?.removeEventListener("abort", onAbort); + this.#pendingCalls.delete(id); + }; + + const onAbort = () => { + if (settled) return; + settled = true; + cleanup(); + this.#output({ + type: "host_tool_cancel", + id: Snowflake.next() as string, + targetId: id, + }); + reject(new Error(`Host tool "${definition.name}" was aborted`)); + }; + + signal?.addEventListener("abort", onAbort, { once: true }); + this.#pendingCalls.set(id, { + resolve: result => { + if (settled) return; + settled = true; + cleanup(); + resolve(result); + }, + reject: error => { + if (settled) return; + settled = true; + cleanup(); + reject(error); + }, + onUpdate, + }); + + this.#output({ + type: "host_tool_call", + id, + toolCallId, + toolName: definition.name, + arguments: args, + }); + + return promise; + } + + rejectAllPending(message: string): void { + const error = new Error(message); + const pendingCalls = Array.from(this.#pendingCalls.values()); + this.#pendingCalls.clear(); + for (const pending of pendingCalls) { + pending.reject(error); + } + } +} diff --git a/packages/coding-agent/src/modes/rpc/rpc-client.ts b/packages/coding-agent/src/modes/rpc/rpc-client.ts index 2e8056958..cccd43100 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-client.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-client.ts @@ -3,13 +3,22 @@ * * Spawns the agent in RPC mode and provides a typed API for all operations. */ -import type { AgentEvent, AgentMessage, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; +import type { AgentEvent, AgentMessage, AgentToolResult, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; import type { ImageContent, Model } from "@oh-my-pi/pi-ai"; import { isRecord, ptree, readJsonl } from "@oh-my-pi/pi-utils"; import type { BashResult } from "../../exec/bash-executor"; import type { SessionStats } from "../../session/agent-session"; import type { CompactionResult } from "../../session/compaction"; -import type { RpcCommand, RpcResponse, RpcSessionState } from "./rpc-types"; +import type { + RpcCommand, + RpcHostToolCallRequest, + RpcHostToolCancelRequest, + RpcHostToolDefinition, + RpcHostToolResult, + RpcHostToolUpdate, + RpcResponse, + RpcSessionState, +} from "./rpc-types"; /** Distributive Omit that works with union types */ type DistributiveOmit = T extends unknown ? Omit : never; @@ -32,12 +41,40 @@ export interface RpcClientOptions { sessionDir?: string; /** Additional CLI arguments */ args?: string[]; + /** Custom tools owned by the embedding host and exposed over the RPC transport */ + customTools?: RpcClientCustomTool[]; } export type ModelInfo = Pick; export type RpcEventListener = (event: AgentEvent) => void; +export interface RpcClientToolContext { + toolCallId: string; + signal: AbortSignal; + sendUpdate(partialResult: RpcClientToolResult): void; +} + +export type RpcClientToolResult = AgentToolResult | string; + +export interface RpcClientCustomTool< + TParams extends Record = Record, + TDetails = unknown, +> extends Omit { + parameters: Record; + execute( + params: TParams, + context: RpcClientToolContext, + ): Promise> | RpcClientToolResult; +} + +export function defineRpcClientTool< + TParams extends Record = Record, + TDetails = unknown, +>(tool: RpcClientCustomTool): RpcClientCustomTool { + return tool; +} + const agentEventTypes = new Set([ "agent_start", "agent_end", @@ -70,6 +107,31 @@ function isAgentEvent(value: unknown): value is AgentEvent { return agentEventTypes.has(type as AgentEvent["type"]); } +function isRpcHostToolCallRequest(value: unknown): value is RpcHostToolCallRequest { + if (!isRecord(value)) return false; + return ( + value.type === "host_tool_call" && + typeof value.id === "string" && + typeof value.toolCallId === "string" && + typeof value.toolName === "string" && + isRecord(value.arguments) + ); +} + +function isRpcHostToolCancelRequest(value: unknown): value is RpcHostToolCancelRequest { + if (!isRecord(value)) return false; + return value.type === "host_tool_cancel" && typeof value.id === "string" && typeof value.targetId === "string"; +} + +function normalizeToolResult(result: RpcClientToolResult): AgentToolResult { + if (typeof result === "string") { + return { + content: [{ type: "text", text: result }], + }; + } + return result; +} + // ============================================================================ // RPC Client // ============================================================================ @@ -79,10 +141,14 @@ export class RpcClient { #eventListeners: RpcEventListener[] = []; #pendingRequests: Map void; reject: (error: Error) => void }> = new Map(); + #customTools: RpcClientCustomTool[] = []; + #pendingHostToolCalls = new Map(); #requestId = 0; #abortController = new AbortController(); - constructor(private options: RpcClientOptions = {}) {} + constructor(private options: RpcClientOptions = {}) { + this.#customTools = [...(options.customTools ?? [])]; + } /** * Start the RPC agent process. @@ -162,6 +228,7 @@ export class RpcClient { try { await readyPromise; + await this.setCustomTools(this.#customTools); } finally { clearTimeout(readyTimeout); } @@ -177,6 +244,10 @@ export class RpcClient { this.#abortController.abort(); this.#process = null; this.#pendingRequests.clear(); + for (const pendingCall of this.#pendingHostToolCalls.values()) { + pendingCall.controller.abort(); + } + this.#pendingHostToolCalls.clear(); } /** @@ -434,6 +505,26 @@ export class RpcClient { return this.#getData<{ messages: AgentMessage[] }>(response).messages; } + /** + * Replace the host-owned custom tools exposed to the RPC session. + * Changes take effect before the next model call. + */ + async setCustomTools(tools: RpcClientCustomTool[]): Promise { + this.#customTools = [...tools]; + if (!this.#process) { + return this.#customTools.map(tool => tool.name); + } + const definitions: RpcHostToolDefinition[] = this.#customTools.map(tool => ({ + name: tool.name, + label: tool.label, + description: tool.description, + parameters: tool.parameters, + hidden: tool.hidden, + })); + const response = await this.#send({ type: "set_host_tools", tools: definitions }); + return this.#getData<{ toolNames: string[] }>(response).toolNames; + } + // ========================================================================= // Helpers // ========================================================================= @@ -514,6 +605,16 @@ export class RpcClient { } } + if (isRpcHostToolCallRequest(data)) { + void this.#handleHostToolCall(data); + return; + } + + if (isRpcHostToolCancelRequest(data)) { + this.#pendingHostToolCalls.get(data.targetId)?.controller.abort(); + return; + } + if (!isAgentEvent(data)) return; // Otherwise it's an event @@ -555,21 +656,83 @@ export class RpcClient { }, }); - // Write to stdin after registering the handler - const stdin = this.#process!.stdin as import("bun").FileSink; - stdin.write(`${JSON.stringify(fullCommand)}\n`); - // flush() returns number | Promise - handle both cases + this.#writeFrame(fullCommand, err => { + this.#pendingRequests.delete(id); + if (settled) return; + settled = true; + clearTimeout(timeoutId); + reject(err); + }); + return promise; + } + + async #handleHostToolCall(request: RpcHostToolCallRequest): Promise { + const tool = this.#customTools.find(candidate => candidate.name === request.toolName); + if (!tool) { + this.#writeFrame({ + type: "host_tool_result", + id: request.id, + result: { + content: [{ type: "text", text: `Host tool "${request.toolName}" is not registered` }], + details: {}, + }, + isError: true, + } satisfies RpcHostToolResult); + return; + } + + const controller = new AbortController(); + this.#pendingHostToolCalls.set(request.id, { controller }); + + const sendUpdate = (partialResult: RpcClientToolResult): void => { + if (controller.signal.aborted) return; + this.#writeFrame({ + type: "host_tool_update", + id: request.id, + partialResult: normalizeToolResult(partialResult), + } satisfies RpcHostToolUpdate); + }; + + try { + const result = await tool.execute(request.arguments, { + toolCallId: request.toolCallId, + signal: controller.signal, + sendUpdate, + }); + if (controller.signal.aborted) return; + this.#writeFrame({ + type: "host_tool_result", + id: request.id, + result: normalizeToolResult(result), + } satisfies RpcHostToolResult); + } catch (error) { + if (controller.signal.aborted) return; + this.#writeFrame({ + type: "host_tool_result", + id: request.id, + result: { + content: [{ type: "text", text: error instanceof Error ? error.message : String(error) }], + details: {}, + }, + isError: true, + } satisfies RpcHostToolResult); + } finally { + this.#pendingHostToolCalls.delete(request.id); + } + } + + #writeFrame(frame: RpcCommand | RpcHostToolResult | RpcHostToolUpdate, onError?: (error: Error) => void): void { + if (!this.#process?.stdin) { + throw new Error("Client not started"); + } + const stdin = this.#process.stdin as import("bun").FileSink; + stdin.write(`${JSON.stringify(frame)}\n`); const flushResult = stdin.flush(); if (flushResult instanceof Promise) { flushResult.catch((err: Error) => { - this.#pendingRequests.delete(id); - if (settled) return; - settled = true; - clearTimeout(timeoutId); - reject(err); + onError?.(err); }); } - return promise; } #getData(response: RpcResponse): T { diff --git a/packages/coding-agent/src/modes/rpc/rpc-mode.ts b/packages/coding-agent/src/modes/rpc/rpc-mode.ts index 8bf2a8c07..bda4637e5 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-mode.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-mode.ts @@ -18,10 +18,14 @@ import type { } from "../../extensibility/extensions"; import { type Theme, theme } from "../../modes/theme/theme"; import type { AgentSession } from "../../session/agent-session"; +import { isRpcHostToolResult, isRpcHostToolUpdate, RpcHostToolBridge } from "./host-tools"; import type { RpcCommand, RpcExtensionUIRequest, RpcExtensionUIResponse, + RpcHostToolCallRequest, + RpcHostToolCancelRequest, + RpcHostToolDefinition, RpcResponse, RpcSessionState, } from "./rpc-types"; @@ -34,7 +38,33 @@ export type PendingExtensionRequest = { reject: (error: Error) => void; }; -type RpcOutput = (obj: RpcResponse | RpcExtensionUIRequest | object) => void; +type RpcOutput = ( + obj: RpcResponse | RpcExtensionUIRequest | RpcHostToolCallRequest | RpcHostToolCancelRequest | object, +) => void; + +function normalizeHostToolDefinitions(tools: RpcHostToolDefinition[]): RpcHostToolDefinition[] { + return tools.map((tool, index) => { + const name = typeof tool.name === "string" ? tool.name.trim() : ""; + if (!name) { + throw new Error(`Host tool at index ${index} must provide a non-empty name`); + } + const description = typeof tool.description === "string" ? tool.description.trim() : ""; + if (!description) { + throw new Error(`Host tool "${name}" must provide a non-empty description`); + } + if (!tool.parameters || typeof tool.parameters !== "object" || Array.isArray(tool.parameters)) { + throw new Error(`Host tool "${name}" must provide a JSON Schema object`); + } + const label = typeof tool.label === "string" && tool.label.trim() ? tool.label.trim() : name; + return { + name, + label, + description, + parameters: tool.parameters, + hidden: tool.hidden === true, + }; + }); +} function shouldEmitRpcTitles(): boolean { const raw = $env.PI_RPC_EMIT_TITLE; @@ -135,6 +165,7 @@ export async function runRpcMode(session: AgentSession): Promise { }; const pendingExtensionRequests = new Map(); + const hostToolBridge = new RpcHostToolBridge(output); // Shutdown request flag (wrapped in object to allow mutation with const) const shutdownState = { requested: false }; @@ -569,6 +600,13 @@ export async function runRpcMode(session: AgentSession): Promise { return success(id, "set_todos", { todoPhases: session.getTodoPhases() }); } + case "set_host_tools": { + const tools = normalizeHostToolDefinitions(command.tools); + const rpcTools = hostToolBridge.setTools(tools); + await session.refreshRpcHostTools(rpcTools); + return success(id, "set_host_tools", { toolNames: tools.map(tool => tool.name) }); + } + // ================================================================= // Model // ================================================================= @@ -759,6 +797,16 @@ export async function runRpcMode(session: AgentSession): Promise { continue; } + if (isRpcHostToolResult(parsed)) { + hostToolBridge.handleResult(parsed); + continue; + } + + if (isRpcHostToolUpdate(parsed)) { + hostToolBridge.handleUpdate(parsed); + continue; + } + // Handle regular commands const command = parsed as RpcCommand; const response = await handleCommand(command); @@ -772,5 +820,6 @@ export async function runRpcMode(session: AgentSession): Promise { } // stdin closed — RPC client is gone, exit cleanly + hostToolBridge.rejectAllPending("RPC client disconnected before host tool execution completed"); process.exit(0); } diff --git a/packages/coding-agent/src/modes/rpc/rpc-types.ts b/packages/coding-agent/src/modes/rpc/rpc-types.ts index 18c467655..3751ad13b 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-types.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-types.ts @@ -4,7 +4,7 @@ * Commands are sent as JSON lines on stdin. * Responses and events are emitted as JSON lines on stdout. */ -import type { AgentMessage, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; +import type { AgentMessage, AgentToolResult, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; import type { Effort, ImageContent, Model } from "@oh-my-pi/pi-ai"; import type { BashResult } from "../../exec/bash-executor"; import type { SessionStats } from "../../session/agent-session"; @@ -27,6 +27,7 @@ export type RpcCommand = // State | { id?: string; type: "get_state" } | { id?: string; type: "set_todos"; phases: TodoPhase[] } + | { id?: string; type: "set_host_tools"; tools: RpcHostToolDefinition[] } // Model | { id?: string; type: "set_model"; provider: string; modelId: string } @@ -107,6 +108,7 @@ export type RpcResponse = // State | { id?: string; type: "response"; command: "get_state"; success: true; data: RpcSessionState } | { id?: string; type: "response"; command: "set_todos"; success: true; data: { todoPhases: TodoPhase[] } } + | { id?: string; type: "response"; command: "set_host_tools"; success: true; data: { toolNames: string[] } } // Model | { @@ -235,6 +237,49 @@ export type RpcExtensionUIRequest = | { type: "extension_ui_request"; id: string; method: "setTitle"; title: string } | { type: "extension_ui_request"; id: string; method: "set_editor_text"; text: string }; +// ============================================================================ +// Host Tool Frames (bidirectional) +// ============================================================================ + +export interface RpcHostToolDefinition { + name: string; + label?: string; + description: string; + parameters: Record; + hidden?: boolean; +} + +/** Emitted by the RPC server when it needs the host to execute a registered tool. */ +export interface RpcHostToolCallRequest { + type: "host_tool_call"; + id: string; + toolCallId: string; + toolName: string; + arguments: Record; +} + +/** Emitted by the RPC server when a pending host tool call should be aborted. */ +export interface RpcHostToolCancelRequest { + type: "host_tool_cancel"; + id: string; + targetId: string; +} + +/** Sent by the host to stream partial tool updates back to the RPC server. */ +export interface RpcHostToolUpdate { + type: "host_tool_update"; + id: string; + partialResult: AgentToolResult; +} + +/** Sent by the host to complete a pending tool call. */ +export interface RpcHostToolResult { + type: "host_tool_result"; + id: string; + result: AgentToolResult; + isError?: boolean; +} + // ============================================================================ // Extension UI Commands (stdin) // ============================================================================ diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 96231fca5..4ab319aae 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -482,6 +482,7 @@ export class AgentSession { #discoverableMCPTools = new Map(); #discoverableMCPSearchIndex: DiscoverableMCPSearchIndex | null = null; #selectedMCPToolNames = new Set(); + #rpcHostToolNames = new Set(); #defaultSelectedMCPServerNames = new Set(); #defaultSelectedMCPToolNames = new Set(); #sessionDefaultSelectedMCPToolNames = new Map(); @@ -2026,6 +2027,49 @@ export class AgentSession { await this.#applyActiveToolsByName(nextActive, { previousSelectedMCPToolNames }); } + /** + * Replace RPC host-owned tools and refresh the active tool set before the next model call. + */ + async refreshRpcHostTools(rpcTools: AgentTool[]): Promise { + const nextToolNames = rpcTools.map(tool => tool.name); + const uniqueToolNames = new Set(nextToolNames); + if (uniqueToolNames.size !== nextToolNames.length) { + throw new Error("RPC host tool names must be unique"); + } + + for (const name of uniqueToolNames) { + if (this.#toolRegistry.has(name) && !this.#rpcHostToolNames.has(name)) { + throw new Error(`RPC host tool "${name}" conflicts with an existing tool`); + } + } + + const previousRpcHostToolNames = new Set(this.#rpcHostToolNames); + const previousActiveToolNames = this.getActiveToolNames(); + for (const name of previousRpcHostToolNames) { + this.#toolRegistry.delete(name); + } + this.#rpcHostToolNames.clear(); + + for (const tool of rpcTools) { + const finalTool = ( + this.#extensionRunner ? new ExtensionToolWrapper(tool, this.#extensionRunner) : tool + ) as AgentTool; + this.#toolRegistry.set(finalTool.name, finalTool); + this.#rpcHostToolNames.add(finalTool.name); + } + + const activeNonRpcToolNames = previousActiveToolNames.filter(name => !previousRpcHostToolNames.has(name)); + const preservedRpcToolNames = previousActiveToolNames.filter( + name => previousRpcHostToolNames.has(name) && this.#rpcHostToolNames.has(name), + ); + const autoActivatedRpcToolNames = rpcTools + .filter(tool => !tool.hidden && !previousRpcHostToolNames.has(tool.name)) + .map(tool => tool.name); + await this.#applyActiveToolsByName( + Array.from(new Set([...activeNonRpcToolNames, ...preservedRpcToolNames, ...autoActivatedRpcToolNames])), + ); + } + /** Whether auto-compaction is currently running */ get isCompacting(): boolean { return this.#autoCompactionAbortController !== undefined || this.#compactionAbortController !== undefined; diff --git a/packages/coding-agent/test/rpc-host-tools.test.ts b/packages/coding-agent/test/rpc-host-tools.test.ts new file mode 100644 index 000000000..5d4e2ed03 --- /dev/null +++ b/packages/coding-agent/test/rpc-host-tools.test.ts @@ -0,0 +1,236 @@ +import { afterEach, describe, expect, it } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import type { AgentEvent } from "@oh-my-pi/pi-agent-core"; +import { defineRpcClientTool, RpcClient } from "@oh-my-pi/pi-coding-agent/modes"; +import { RpcHostToolBridge } from "@oh-my-pi/pi-coding-agent/modes/rpc/host-tools"; +import type { + RpcHostToolCallRequest, + RpcHostToolCancelRequest, + RpcHostToolUpdate, +} from "@oh-my-pi/pi-coding-agent/modes/rpc/rpc-types"; + +const tempPaths: string[] = []; + +afterEach(async () => { + await Promise.all( + tempPaths.splice(0).map(async filePath => { + try { + await fs.rm(filePath, { force: true }); + } catch {} + }), + ); +}); + +describe("RpcHostToolBridge", () => { + it("forwards host tool updates and results to the pending execution", async () => { + const frames: Array = []; + const bridge = new RpcHostToolBridge(frame => { + frames.push(frame); + }); + const [tool] = bridge.setTools([ + { + name: "host_sum", + label: "Host Sum", + description: "Adds numbers in the host process", + parameters: { + type: "object", + properties: { + left: { type: "number" }, + right: { type: "number" }, + }, + required: ["left", "right"], + additionalProperties: false, + }, + }, + ]); + + const updates: RpcHostToolUpdate["partialResult"][] = []; + const execution = tool.execute("toolu_1", { left: 2, right: 3 }, undefined, update => { + updates.push(update); + }); + + expect(frames).toHaveLength(1); + const request = frames[0]; + if (!request || request.type !== "host_tool_call") { + throw new Error("Expected host_tool_call frame"); + } + + bridge.handleUpdate({ + type: "host_tool_update", + id: request.id, + partialResult: { + content: [{ type: "text", text: "working" }], + }, + }); + expect(updates).toHaveLength(1); + expect(updates[0]?.content[0]).toEqual({ type: "text", text: "working" }); + + bridge.handleResult({ + type: "host_tool_result", + id: request.id, + result: { + content: [{ type: "text", text: "5" }], + }, + }); + + await expect(execution).resolves.toEqual({ + content: [{ type: "text", text: "5" }], + }); + }); + + it("emits a cancel frame when the host tool execution is aborted", async () => { + const frames: Array = []; + const bridge = new RpcHostToolBridge(frame => { + frames.push(frame); + }); + const [tool] = bridge.setTools([ + { + name: "host_wait", + description: "Waits in the host process", + parameters: { + type: "object", + properties: {}, + additionalProperties: false, + }, + }, + ]); + + const controller = new AbortController(); + const execution = tool.execute("toolu_2", {}, controller.signal); + const request = frames[0]; + if (!request || request.type !== "host_tool_call") { + throw new Error("Expected host_tool_call frame"); + } + + controller.abort(); + + expect(frames[1]).toMatchObject({ + type: "host_tool_cancel", + targetId: request.id, + }); + await expect(execution).rejects.toThrow('Host tool "host_wait" was aborted'); + }); +}); + +describe("RpcClient custom tools", () => { + it("registers host custom tools and serves tool calls over the RPC transport", async () => { + const scriptPath = path.join(os.tmpdir(), `omp-rpc-host-tools-${Date.now()}.js`); + tempPaths.push(scriptPath); + await Bun.write( + scriptPath, + ` +const encoder = new TextEncoder(); +let buffer = ""; + +function write(frame) { + process.stdout.write(JSON.stringify(frame) + "\\n"); +} + +write({ type: "ready" }); + +process.stdin.on("data", chunk => { + buffer += chunk.toString("utf8"); + let index = buffer.indexOf("\\n"); + while (index !== -1) { + const line = buffer.slice(0, index).trim(); + buffer = buffer.slice(index + 1); + if (line) handle(JSON.parse(line)); + index = buffer.indexOf("\\n"); + } +}); + +function handle(frame) { + if (frame.type === "set_host_tools") { + write({ + id: frame.id, + type: "response", + command: "set_host_tools", + success: true, + data: { toolNames: frame.tools.map(tool => tool.name) }, + }); + return; + } + if (frame.type === "prompt") { + write({ id: frame.id, type: "response", command: "prompt", success: true }); + write({ type: "agent_start" }); + write({ + type: "host_tool_call", + id: "host-call-1", + toolCallId: "toolu_host_1", + toolName: "echo_host", + arguments: { message: "hello" }, + }); + return; + } + if (frame.type === "host_tool_update") { + write({ + type: "tool_execution_update", + toolCallId: "toolu_host_1", + toolName: "echo_host", + args: { message: "hello" }, + partialResult: frame.partialResult, + }); + return; + } + if (frame.type === "host_tool_result") { + write({ + type: "tool_execution_end", + toolCallId: "toolu_host_1", + toolName: "echo_host", + result: frame.result, + isError: frame.isError === true, + }); + write({ type: "agent_end", messages: [] }); + } +} +`, + ); + + const client = new RpcClient({ + cliPath: scriptPath, + customTools: [ + defineRpcClientTool<{ message: string }>({ + name: "echo_host", + description: "Echo a value from the embedding host", + parameters: { + type: "object", + properties: { + message: { type: "string" }, + }, + required: ["message"], + additionalProperties: false, + }, + async execute(args, context) { + context.sendUpdate(`working:${args.message}`); + return `host:${args.message}`; + }, + }), + ], + }); + + try { + await client.start(); + const events = await client.promptAndWait("Trigger host tool"); + const toolEnd = events.find( + (event): event is Extract => + event.type === "tool_execution_end", + ); + expect(toolEnd?.toolName).toBe("echo_host"); + expect(toolEnd?.result).toEqual({ + content: [{ type: "text", text: "host:hello" }], + }); + + const toolUpdate = events.find( + (event): event is Extract => + event.type === "tool_execution_update", + ); + expect(toolUpdate?.partialResult).toEqual({ + content: [{ type: "text", text: "working:hello" }], + }); + } finally { + client.stop(); + } + }); +});