From da3797605fe49bf1b55e07fc5790474b3832277e Mon Sep 17 00:00:00 2001 From: youwang <61931019+MelodyVAR@users.noreply.github.com> Date: Sun, 27 Sep 2026 17:44:42 +0000 Subject: [PATCH 1/2] fix(mcp): refresh complete catalogs and discover active tools --- packages/coding-agent/docs/mcp-catalog.md | 29 + .../src/core/extensions/loader.ts | 10 +- .../src/core/extensions/runner.ts | 6 + .../coding-agent/src/core/extensions/types.ts | 11 +- .../coding-agent/src/step/mcp-catalog.test.ts | 882 ++++++++++++++++++ packages/coding-agent/src/step/mcp-client.ts | 115 ++- packages/coding-agent/src/step/mcp.ts | 133 ++- .../coding-agent/src/step/tool-profile.ts | 7 +- .../test/extensions-tool-catalog.test.ts | 217 +++++ 9 files changed, 1337 insertions(+), 73 deletions(-) create mode 100644 packages/coding-agent/docs/mcp-catalog.md create mode 100644 packages/coding-agent/src/step/mcp-catalog.test.ts create mode 100644 packages/coding-agent/test/extensions-tool-catalog.test.ts diff --git a/packages/coding-agent/docs/mcp-catalog.md b/packages/coding-agent/docs/mcp-catalog.md new file mode 100644 index 00000000..fd45bcee --- /dev/null +++ b/packages/coding-agent/docs/mcp-catalog.md @@ -0,0 +1,29 @@ +# MCP catalogs and tool discovery + +Step publishes each MCP server's complete tool catalog after connecting. It follows every `tools/list` cursor, including cursors on empty pages and empty-string cursors. A repeated cursor fails the listing instead of looping. The server's `startup_timeout_sec` covers connection and initial pagination together; shutdown cancels pending discovery. No partial catalog is registered. + +The interactive TUI starts before discovery finishes. Each server publishes independently after yielding to the event loop. Print, JSON, and RPC session binding waits for initial discovery to settle. Registering one server's catalog refreshes the tool registry once, regardless of its tool count. + +For every catalog, `enabled_tools` is an allow list when present, and `disabled_tools` takes precedence. These filters apply before registration, including during refreshes. Registered names keep the existing `__` convention; calls use the server's original tool name. + +Step listens for `notifications/tools/list_changed` from the start of the connection. Bursts are coalesced, with one listing in progress per server and a pending follow-up when notifications arrive during that listing. A successful refresh replaces the server's registered tools in one batch, adding new names, updating definitions, and removing names that disappeared. An empty successful catalog removes all of that server's tools. Pagination or output-schema preparation failures leave the previous catalog and tool count intact and produce a warning; a later notification can retry. + +A closed connection removes its tools and changes its status to failed. Session shutdown removes its MCP registrations and closes the clients. Results that settle after cancellation or connection closure cannot publish a catalog. Calls already in progress retain their captured tool definition when a live catalog is updated or a tool is removed. This behavior does not add an automatic reconnect policy. + +`find_tools` reads the currently active session tools on every invocation, so late MCP and extension registrations become discoverable. Inactive and removed tools disappear from search. Each match includes its callable name, description, and JSON parameter schema. The optional `ExtensionContext.getToolCatalog()` accessor supplies this metadata through the runner's existing session actions. Direct SDK callers whose context does not provide that accessor retain the builtin Step catalog fallback. + +Extensions can replace part of their own catalog with the existing batch API: + +```ts +pi.registerTools(nextDefinitions, { remove: previousNames }); +``` + +Removal affects only the calling extension's registrations. Names supplied in both lists receive the new definition. The registry refreshes once after the batch, so callers see the finished replacement. Removing a winning extension definition can reveal another extension or builtin definition under the existing precedence rules. Empty batches and removal of names the caller does not own do not refresh the registry. Stale extension APIs reject the operation. + +MCP input schemas pass unchanged through registration to the existing agent-core/provider argument-validation boundary. The regression tests exercise integers and bounds, nested required properties, additional properties, nullable enums, array constraints, `anyOf`, `oneOf`, `allOf`, and local `definitions`/`$defs` references. Provider coercion and optional-null handling still apply; preserving JSON Schema does not make validation a strict, non-coercing JSON Schema validator. Format and dialect support remain those of the installed validator. In particular, an unregistered custom format is retained in the schema but is not thereby enforced. + +Both session tools and the one-shot remote MCP client use the same catalog and call adapters. With MCP SDK 1.27.1, `client.listTools()` replaces its cached output validators and task metadata for each page, and `client.callTool()` reads those mutable validators after receiving a response. Step stages pages with the public `client.request()` and `ListToolsResultSchema` instead. Each callable definition captures its own output validator using the SDK's public `AjvJsonSchemaValidator`, with isolated schema-ID scope. A failed refresh or a later schema using the same `$id` cannot change an older call's validation. Successful schema-bearing calls must return matching `structuredContent`; error results may omit it. Output-schema support remains that of the SDK's default validator, without a blanket guarantee for arbitrary dialects or external references. + +Tools declaring required task execution continue to fail before sending a normal tool call, with an explicit task-execution error. This adapter does not implement the SDK's experimental task execution protocol. The SDK version remains 1.27.1. + +Focused regressions are in `src/step/mcp-catalog.test.ts`, `src/step/mcp-startup.test.ts`, and `test/extensions-tool-catalog.test.ts`. They use a real HTTP/SSE MCP peer, the installed SDK, real session registry actions, provider argument validation, and the agent loop. diff --git a/packages/coding-agent/src/core/extensions/loader.ts b/packages/coding-agent/src/core/extensions/loader.ts index 40630c7c..b0a62784 100644 --- a/packages/coding-agent/src/core/extensions/loader.ts +++ b/packages/coding-agent/src/core/extensions/loader.ts @@ -311,16 +311,20 @@ function createExtensionAPI( runtime.refreshTools(); }, - registerTools(tools: readonly ToolDefinition[]): void { + registerTools(tools: readonly ToolDefinition[], options?: { remove?: readonly string[] }): void { assertActive(); - if (tools.length === 0) return; + let changed = false; + for (const name of options?.remove ?? []) { + if (extension.tools.delete(name)) changed = true; + } for (const tool of tools) { extension.tools.set(tool.name, { definition: tool, sourceInfo: extension.sourceInfo, }); + changed = true; } - runtime.refreshTools(); + if (changed) runtime.refreshTools(); }, registerCommand(name: string, options: Omit): void { diff --git a/packages/coding-agent/src/core/extensions/runner.ts b/packages/coding-agent/src/core/extensions/runner.ts index b91f96e0..819747e4 100644 --- a/packages/coding-agent/src/core/extensions/runner.ts +++ b/packages/coding-agent/src/core/extensions/runner.ts @@ -801,6 +801,12 @@ export class ExtensionRunner { runner.assertActive(); return runner.getSystemPromptFn(); }, + getToolCatalog: () => { + runner.assertActive(); + runner.runtime.assertActive(); + const active = new Set(runner.runtime.getActiveTools()); + return runner.runtime.getAllTools().filter((tool) => active.has(tool.name)); + }, get autoRetryEnabled() { runner.assertActive(); return runner.getAutoRetryEnabledFn?.(); diff --git a/packages/coding-agent/src/core/extensions/types.ts b/packages/coding-agent/src/core/extensions/types.ts index ea9c6946..3191f448 100644 --- a/packages/coding-agent/src/core/extensions/types.ts +++ b/packages/coding-agent/src/core/extensions/types.ts @@ -388,6 +388,8 @@ export interface ExtensionContext { compact(options?: CompactOptions): void; /** Get the current effective system prompt. */ getSystemPrompt(): string; + /** Read the currently active tool definitions at call time, when the host provides a live catalog. */ + getToolCatalog?(): readonly ToolInfo[]; /** Whether Pi's native provider retry loop is enabled, when exposed by the host. */ readonly autoRetryEnabled?: boolean; /** Toggle Pi's native provider retry loop, when exposed by the host. */ @@ -1356,11 +1358,12 @@ export interface ExtensionAPI { ): void; /** - * Register several tools and refresh the tool registry once. Registering a - * large catalog one tool at a time rebuilds the registry and the system - * prompt per tool, which is quadratic work on the startup path. + * Register several tools and optionally remove this extension's registrations + * in one registry refresh. Removals never delete another extension's tools; + * removing an override reveals the next definition under normal precedence. + * Names present in both lists are replaced by the supplied definitions. */ - registerTools(tools: readonly ToolDefinition[]): void; + registerTools(tools: readonly ToolDefinition[], options?: { remove?: readonly string[] }): void; // ========================================================================= // Command, Shortcut, Flag Registration diff --git a/packages/coding-agent/src/step/mcp-catalog.test.ts b/packages/coding-agent/src/step/mcp-catalog.test.ts new file mode 100644 index 00000000..7d1919d6 --- /dev/null +++ b/packages/coding-agent/src/step/mcp-catalog.test.ts @@ -0,0 +1,882 @@ +import { mkdtemp, rm } from "node:fs/promises"; +import { createServer, type ServerResponse } from "node:http"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { setTimeout as delay, setImmediate as yieldToEventLoop } from "node:timers/promises"; +import { Client } from "@modelcontextprotocol/sdk/client/index.js"; +import type { CallToolResult, ListToolsResult, Tool as McpTool } from "@modelcontextprotocol/sdk/types.js"; +import { Agent } from "@step-harness/agent-core"; +import { + type AssistantMessage, + createAssistantMessageEventStream, + type ToolCall, + validateToolArguments, +} from "@step-harness/providers"; +import { afterEach, expect, test, vi } from "vitest"; +import { createTestExtensionsResult, createTestResourceLoader, stepModel } from "../../test/utilities.ts"; +import { createEventBus } from "../core/event-bus.ts"; +import { createExtensionRuntime, loadExtensionFromFactory } from "../core/extensions/loader.ts"; +import type { ExtensionMode, ToolDefinition } from "../core/extensions/types.ts"; +import { createAgentSession } from "../core/sdk.ts"; +import { SessionManager } from "../core/session-manager.ts"; +import { SettingsManager } from "../core/settings-manager.ts"; +import { wrapToolDefinition } from "../core/tools/tool-definition-wrapper.ts"; +import type { StepConfigDocument } from "./config-toml.ts"; +import { connectStepMcpServer, createStepMcpExtension, getStepMcpStatuses } from "./mcp.ts"; +import { invokeRemoteMcpTool } from "./mcp-client.ts"; +import { createStepToolProfile } from "./tool-profile.ts"; + +const config = vi.hoisted(() => ({ value: {} as StepConfigDocument })); +vi.mock("./config-toml.ts", () => ({ readGlobalStepConfig: () => config.value })); +vi.mock("./plugins.ts", () => ({ + defaultStepPluginsDir: () => "/unused-test-plugins", + listStepPluginDirectories: async () => [], + ensureBuiltinPluginsInstalled: async () => ({ installed: [], warnings: [] }), + provisionBuiltinPlugin: async () => undefined, +})); +vi.mock("./mcp-oauth.ts", () => ({ hasStoredMcpOAuthCredential: () => false })); + +const cleanups: Array<() => Promise> = []; +const releases: Array<() => void> = []; +afterEach(async () => { + for (const release of releases.splice(0)) release(); + for (const cleanup of cleanups.splice(0).reverse()) await cleanup(); + vi.restoreAllMocks(); +}); + +function gate() { + let release = () => {}; + const promise = new Promise((resolve) => { + release = resolve; + }); + releases.push(release); + return { promise, release }; +} + +function tool(name: string, description = name, inputSchema: McpTool["inputSchema"] = { type: "object" }): McpTool { + return { name, description, inputSchema }; +} + +type ListHandler = (cursor: string | undefined, request: number) => ListToolsResult | Promise; + +/** A real HTTP/SSE peer exercises the pinned SDK's requests and notification routing. */ +async function catalogServer(initial: ListToolsResult[], initializedResponse?: Promise) { + let pages = initial; + let initializedRequested = false; + let initializedResponded = false; + const streams = new Set(); + const requests: Array = []; + const calls: Array<{ name: string; arguments: Record }> = []; + let onCall = async (_args: { name: string; arguments: Record }): Promise => ({ + content: [{ type: "text", text: "called" }], + }); + let listing = 0; + let maxConcurrentLists = 0; + let onList: ListHandler = (cursor) => pages[cursor === undefined ? 0 : Number(cursor)]; + const server = createServer(async (req, res) => { + if (req.method === "GET") { + res.writeHead(200, { "Content-Type": "text/event-stream", "Cache-Control": "no-cache" }); + res.flushHeaders(); + streams.add(res); + res.on("close", () => streams.delete(res)); + return; + } + if (req.method !== "POST") { + res.writeHead(405).end(); + return; + } + const chunks: Buffer[] = []; + for await (const chunk of req) chunks.push(Buffer.from(chunk)); + const message = JSON.parse(Buffer.concat(chunks).toString()) as { + id?: number; + method: string; + params?: { cursor?: string; name: string; arguments: Record }; + }; + if (message.id === undefined) { + if (message.method === "notifications/initialized") { + initializedRequested = true; + await initializedResponse; + initializedResponded = true; + } + res.writeHead(202).end(); + return; + } + try { + let result: unknown; + if (message.method === "initialize") { + result = { + protocolVersion: "2025-03-26", + capabilities: { tools: { listChanged: true } }, + serverInfo: { name: "catalog-test", version: "1" }, + }; + } else if (message.method === "tools/list") { + requests.push(message.params?.cursor); + listing++; + maxConcurrentLists = Math.max(maxConcurrentLists, listing); + try { + result = await onList(message.params?.cursor, requests.length); + } finally { + listing--; + } + } else if (message.method === "tools/call") { + calls.push(message.params!); + result = await onCall(message.params!); + } else { + throw new Error(`Unexpected method ${message.method}`); + } + res.writeHead(200, { "Content-Type": "application/json" }); + res.end(JSON.stringify({ jsonrpc: "2.0", id: message.id, result })); + } catch (error) { + res.writeHead(200, { "Content-Type": "application/json" }); + res.end( + JSON.stringify({ + jsonrpc: "2.0", + id: message.id, + error: { code: -32603, message: error instanceof Error ? error.message : String(error) }, + }), + ); + } + }); + await new Promise((resolve) => server.listen(0, "127.0.0.1", resolve)); + const address = server.address(); + if (!address || typeof address === "string") throw new Error("Missing test port"); + cleanups.push(async () => { + server.closeAllConnections(); + await new Promise((resolve) => server.close(() => resolve())); + }); + return { + url: `http://127.0.0.1:${address.port}/mcp`, + requests, + calls, + maxConcurrentLists: () => maxConcurrentLists, + initializedRequested: () => initializedRequested, + initializedResponded: () => initializedResponded, + setPages(next: ListToolsResult[]) { + pages = next; + }, + setListHandler(handler: ListHandler) { + onList = handler; + }, + setCallHandler(handler: typeof onCall) { + onCall = handler; + }, + async notify(count = 1) { + await vi.waitFor(() => expect(streams.size).toBeGreaterThan(0)); + const notification = `data: ${JSON.stringify({ jsonrpc: "2.0", method: "notifications/tools/list_changed" })}\n\n`; + for (const stream of streams) stream.write(notification.repeat(count)); + }, + }; +} + +async function setup(mode: ExtensionMode = "print") { + const runtime = createExtensionRuntime(); + const extension = await loadExtensionFromFactory(createStepMcpExtension(), process.cwd(), createEventBus(), runtime); + const notify = vi.fn(); + const ctx = { cwd: process.cwd(), mode, isProjectTrusted: () => true, ui: { notify } }; + const batches: string[][] = []; + runtime.refreshTools = () => batches.push([...extension.tools.keys()].sort()); + const start = () => extension.handlers.get("session_start")![0]({ type: "session_start" }, ctx); + const stop = () => extension.handlers.get("session_shutdown")![0]({ type: "session_shutdown" }, ctx); + cleanups.push(stop); + return { start, stop, extension, runtime, batches, notify }; +} + +test.each(["print", "rpc"] as const)( + "%s readiness waits for all pages, including an empty continuation page", + async (mode) => { + const lastPage = gate(); + const server = await catalogServer([]); + server.setListHandler(async (cursor) => { + if (cursor === undefined) return { tools: [tool("first")], nextCursor: "1" }; + if (cursor === "1") return { tools: [], nextCursor: "2" }; + await lastPage.promise; + return { tools: [tool("last")] }; + }); + config.value = { mcp_servers: { pages: { url: server.url } } }; + const harness = await setup(mode); + let ready = false; + const starting = harness.start().then(() => { + ready = true; + }); + await vi.waitFor(() => expect(server.requests).toEqual([undefined, "1", "2"])); + expect(ready).toBe(false); + expect(harness.extension.tools.size).toBe(0); + expect(harness.batches).toEqual([]); + lastPage.release(); + await starting; + expect(harness.batches).toEqual([["pages__first", "pages__last"]]); + expect(getStepMcpStatuses()).toEqual([{ name: "pages", status: "connected", toolCount: 2 }]); + }, +); + +test("the one-shot remote client consumes empty continuation pages too", async () => { + const server = await catalogServer([ + { tools: [tool("first")], nextCursor: "1" }, + { tools: [], nextCursor: "2" }, + { tools: [tool("last")] }, + ]); + const result = await invokeRemoteMcpTool({ + serverName: "pages", + serverUrl: server.url, + toolName: "last", + arguments: {}, + }); + expect(result).toEqual({ content: "called" }); + expect(server.requests).toEqual([undefined, "1", "2"]); +}); + +test.each(["timeout", "caller cancellation"])( + "one-shot handshake honors %s while the initialized notification response is withheld", + async (kind) => { + const withheld = gate(); + const server = await catalogServer([{ tools: [tool("check")] }], withheld.promise); + const controller = new AbortController(); + let outcome: { error: unknown } | { result: unknown } | undefined; + const invocation = invokeRemoteMcpTool({ + serverName: "handshake", + serverUrl: server.url, + toolName: "check", + arguments: {}, + timeoutMs: kind === "timeout" ? 100 : 30_000, + signal: controller.signal, + }).then( + (result) => { + outcome = { result }; + }, + (error: unknown) => { + outcome = { error }; + }, + ); + // Release only in cleanup so a hung baseline settles without leaking a client. + cleanups.push(async () => { + withheld.release(); + await invocation; + }); + await vi.waitFor(() => expect(server.initializedRequested()).toBe(true)); + if (kind === "caller cancellation") controller.abort(); + await vi.waitFor(() => expect(outcome).toEqual({ error: expect.any(Error) }), { timeout: 700, interval: 10 }); + expect(outcome).toMatchObject({ + error: { message: expect.stringContaining("MCP tool handshake.check failed:") }, + }); + expect(server.initializedResponded()).toBe(false); + expect(server.requests).toEqual([]); + expect(server.calls).toEqual([]); + }, +); + +test("one-shot handshake rejects an already-aborted caller before starting the client", async () => { + const server = await catalogServer([{ tools: [tool("check")] }]); + const connecting = vi.spyOn(Client.prototype, "connect"); + const controller = new AbortController(); + controller.abort(new Error("cancelled before connect")); + await expect( + invokeRemoteMcpTool({ + serverName: "handshake", + serverUrl: server.url, + toolName: "check", + arguments: {}, + signal: controller.signal, + }), + ).rejects.toThrow("cancelled before connect"); + expect(connecting).not.toHaveBeenCalled(); + expect(server.initializedRequested()).toBe(false); + expect(server.requests).toEqual([]); + expect(server.calls).toEqual([]); +}); + +test.each(["session", "one-shot"])("%s rejects repeated cursors before requesting a page twice", async (kind) => { + const server = await catalogServer([]); + server.setListHandler((cursor, request) => { + if (request > 4) throw new Error("test stopped an unbounded pagination loop"); + return { tools: [tool("first")], nextCursor: cursor === "1" ? "2" : "1" }; + }); + if (kind === "session") { + const connecting = connectStepMcpServer({ name: "cycle", declaration: { url: server.url } }).then((connected) => { + cleanups.push(() => connected.client.close()); + return connected; + }); + await expect(connecting).rejects.toThrow(/cursor/i); + } else { + await expect( + invokeRemoteMcpTool({ serverName: "cycle", serverUrl: server.url, toolName: "missing", arguments: {} }), + ).rejects.toThrow(/cursor/i); + } + expect(server.requests).toEqual([undefined, "1", "2"]); +}); + +test("an empty string cursor is opaque and is still followed", async () => { + const server = await catalogServer([]); + server.setListHandler((cursor) => + cursor === undefined ? { tools: [], nextCursor: "" } : { tools: [tool("last")] }, + ); + const connecting = await connectStepMcpServer({ name: "opaque", declaration: { url: server.url } }); + cleanups.push(() => connecting.client.close()); + expect(connecting.tools.map((entry) => entry.name)).toEqual(["last"]); + expect(server.requests).toEqual([undefined, ""]); +}); + +test("one startup deadline covers the handshake and all catalog pages", async () => { + const server = await catalogServer([]); + server.setListHandler(async (cursor) => { + await delay(600); + return cursor === undefined ? { tools: [tool("first")], nextCursor: "1" } : { tools: [tool("last")] }; + }); + const connecting = connectStepMcpServer({ + name: "deadline", + declaration: { url: server.url, startup_timeout_sec: 1 }, + }).then((connected) => { + cleanups.push(() => connected.client.close()); + return connected; + }); + await expect(connecting).rejects.toThrow(/timeout|timed out|aborted/i); + expect(server.requests).toEqual([undefined, "1"]); +}); + +test("session abort cancels a later page and never publishes a partial catalog", async () => { + const pending = gate(); + const server = await catalogServer([]); + server.setListHandler(async (cursor) => { + if (cursor === undefined) return { tools: [tool("first")], nextCursor: "1" }; + await pending.promise; + return { tools: [tool("late")] }; + }); + config.value = { mcp_servers: { abort: { url: server.url } } }; + const harness = await setup("tui"); + await harness.start(); + await vi.waitFor(() => expect(server.requests).toEqual([undefined, "1"])); + await harness.stop(); + pending.release(); + await yieldToEventLoop(); + expect(harness.batches).toEqual([]); + expect(harness.extension.tools.size).toBe(0); + expect(harness.notify).not.toHaveBeenCalled(); +}); + +test("list changes add, update and remove filtered tools in one complete batch", async () => { + const server = await catalogServer([{ tools: [tool("keep", "old"), tool("remove"), tool("denied")] }]); + config.value = { + mcp_servers: { + live: { url: server.url, enabled_tools: ["keep", "remove", "add", "denied"], disabled_tools: ["denied"] }, + }, + }; + const harness = await setup(); + await harness.start(); + const schema = { + type: "object" as const, + properties: { limit: { type: "integer", minimum: 2 } }, + required: ["limit"], + }; + server.setPages([ + { tools: [tool("keep", "updated", schema), tool("denied"), tool("unlisted")], nextCursor: "1" }, + { tools: [tool("add")] }, + ]); + await server.notify(); + await vi.waitFor(() => expect([...harness.extension.tools.keys()].sort()).toEqual(["live__add", "live__keep"])); + expect(harness.batches).toEqual([ + ["live__keep", "live__remove"], + ["live__add", "live__keep"], + ]); + expect(harness.extension.tools.get("live__keep")!.definition).toMatchObject({ + description: "updated", + parameters: schema, + }); + expect(getStepMcpStatuses()).toEqual([{ name: "live", status: "connected", toolCount: 2 }]); +}); + +test("a successful empty refresh removes the server's entire registered catalog", async () => { + const server = await catalogServer([{ tools: [tool("old")] }]); + config.value = { mcp_servers: { live: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + server.setPages([{ tools: [] }]); + await server.notify(); + await vi.waitFor(() => expect(harness.extension.tools.size).toBe(0)); + expect(harness.batches).toEqual([["live__old"], []]); + expect(getStepMcpStatuses()).toEqual([{ name: "live", status: "connected", toolCount: 0 }]); +}); + +test("a failed later refresh page keeps the last good catalog and a later notification can recover", async () => { + const server = await catalogServer([{ tools: [tool("good")] }]); + config.value = { mcp_servers: { live: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const original = harness.extension.tools.get("live__good"); + server.setListHandler((cursor) => { + if (cursor === undefined) return { tools: [tool("partial")], nextCursor: "1" }; + throw new Error("catalog unavailable"); + }); + await server.notify(); + await vi.waitFor(() => + expect(harness.notify).toHaveBeenCalledWith(expect.stringContaining("catalog unavailable"), "warning"), + ); + expect(harness.extension.tools.get("live__good")).toBe(original); + expect(harness.batches).toEqual([["live__good"]]); + expect(getStepMcpStatuses()).toEqual([{ name: "live", status: "connected", toolCount: 1 }]); + server.setListHandler(() => ({ tools: [tool("recovered")] })); + await server.notify(); + await vi.waitFor(() => expect(harness.extension.tools.has("live__recovered")).toBe(true)); + expect(harness.batches).toEqual([["live__good"], ["live__recovered"]]); +}); + +test("notification bursts coalesce to one running refresh and one follow-up", async () => { + const server = await catalogServer([{ tools: [tool("initial")] }]); + config.value = { mcp_servers: { live: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const pending = gate(); + server.setListHandler(async (_cursor, request) => { + if (request === 2) { + await pending.promise; + return { tools: [tool("middle")] }; + } + return { tools: [tool("latest")] }; + }); + await server.notify(30); + await vi.waitFor(() => expect(server.requests).toHaveLength(2)); + await server.notify(30); + // Give notification routing an event-loop turn while the request is held. + await delay(30); + expect(server.requests).toHaveLength(2); + pending.release(); + await vi.waitFor(() => expect(harness.extension.tools.has("live__latest")).toBe(true)); + expect(server.requests).toHaveLength(3); + expect(server.maxConcurrentLists()).toBe(1); + expect(harness.batches.at(-1)).toEqual(["live__latest"]); +}); + +test("a notification during initial listing is retained for a follow-up refresh", async () => { + const initial = gate(); + const server = await catalogServer([]); + server.setListHandler(async (_cursor, request) => { + if (request === 1) { + await initial.promise; + return { tools: [tool("initial")] }; + } + return { tools: [tool("latest")] }; + }); + config.value = { mcp_servers: { live: { url: server.url } } }; + const harness = await setup("tui"); + await harness.start(); + await vi.waitFor(() => expect(server.requests).toHaveLength(1)); + await server.notify(); + await delay(30); + initial.release(); + await vi.waitFor(() => expect(harness.extension.tools.has("live__latest")).toBe(true)); + expect(server.requests).toHaveLength(2); +}); + +test.each(["shutdown", "disconnect"])("%s removes its catalog and ignores a late refresh result", async (action) => { + const connecting = vi.spyOn(Client.prototype, "connect"); + const server = await catalogServer([{ tools: [tool("initial")] }]); + config.value = { mcp_servers: { live: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const client = connecting.mock.contexts[0] as Client; + const late = gate(); + // A response already in userland may settle after SDK cancellation/close. + server.setPages([{ tools: [tool("late")] }]); + const request = client.request.bind(client); + const listing = vi.spyOn(client, "request").mockImplementationOnce(async (...args) => { + const result = await request(...args); + await late.promise; + return result; + }); + await server.notify(); + await vi.waitFor(() => expect(listing).toHaveBeenCalledTimes(1)); + if (action === "shutdown") await harness.stop(); + else await client.close(); + expect(harness.extension.tools.size).toBe(0); + late.release(); + await yieldToEventLoop(); + await yieldToEventLoop(); + expect(harness.batches).toEqual([["live__initial"], []]); + expect(harness.notify).not.toHaveBeenCalled(); + if (action === "shutdown") expect(getStepMcpStatuses()).toEqual([]); + else expect(getStepMcpStatuses()).toEqual([{ name: "live", status: "failed", toolCount: 0 }]); +}); + +const complexSchema: McpTool["inputSchema"] = { + type: "object", + properties: { + count: { type: "integer", minimum: 1, maximum: 10 }, + config: { + type: "object", + properties: { + mode: { enum: ["fast", "slow", null] }, + paths: { type: "array", items: { type: "string", minLength: 2 }, minItems: 1, uniqueItems: true }, + }, + required: ["mode", "paths"], + additionalProperties: false, + }, + choice: { + oneOf: [ + { type: "string", const: "auto" }, + { type: "integer", minimum: 3 }, + ], + }, + target: { anyOf: [{ type: "string", pattern: "^ok:" }, { type: "null" }] }, + range: { allOf: [{ type: "integer", minimum: 2 }, { maximum: 5 }] }, + window: { $ref: "#/definitions/window" }, + }, + required: ["count", "config", "choice", "target", "range", "window"], + additionalProperties: false, + definitions: { + window: { + type: "object", + properties: { size: { type: "integer", minimum: 1 } }, + required: ["size"], + additionalProperties: false, + }, + }, +}; + +function validArguments() { + return { + count: 2, + config: { mode: "fast", paths: ["ab"] }, + choice: "auto", + target: "ok:yes", + range: 3, + window: { size: 2 }, + }; +} + +async function schemaTool(schema = complexSchema) { + const server = await catalogServer([{ tools: [tool("check", "Validate complex arguments", schema)] }]); + config.value = { mcp_servers: { schema: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + return { server, definition: harness.extension.tools.get("schema__check")!.definition }; +} + +function validate(definition: ToolDefinition, args: Record) { + return validateToolArguments(definition, { + type: "toolCall", + id: "schema-call", + name: definition.name, + arguments: args, + }); +} + +test("MCP input schemas survive registration with their original JSON semantics", async () => { + const { definition } = await schemaTool(); + expect(JSON.parse(JSON.stringify(definition.parameters))).toEqual(complexSchema); + expect(validate(definition, validArguments())).toEqual(validArguments()); +}); + +test("nullable enum members remain valid through provider argument validation", async () => { + const { definition } = await schemaTool(); + const args = { ...validArguments(), config: { mode: null, paths: ["ab"] } }; + expect(validate(definition, args)).toEqual(args); +}); + +test.each([ + ["integer", { count: 1.5 }], + ["minimum", { count: 0 }], + ["maximum", { count: 11 }], + ["nested required", { config: { paths: ["ab"] } }], + ["nested additionalProperties", { config: { mode: "fast", paths: ["ab"], extra: true } }], + ["enum", { config: { mode: "invalid", paths: ["ab"] } }], + ["item minLength", { config: { mode: "fast", paths: ["a"] } }], + ["uniqueItems", { config: { mode: "fast", paths: ["ab", "ab"] } }], + ["minItems", { config: { mode: "fast", paths: [] } }], + ["oneOf", { choice: "manual" }], + ["anyOf", { target: "bad:value" }], + ["allOf", { range: 6 }], + ["local definitions", { window: { size: 1.5 } }], + ["root additionalProperties", { extra: true }], +])("provider validation enforces the preserved %s constraint", async (_name, invalid) => { + const { definition } = await schemaTool(); + expect(() => validate(definition, { ...validArguments(), ...invalid })).toThrow(/Validation failed/); +}); + +test("local $defs references remain enforceable", async () => { + const schema: McpTool["inputSchema"] = { + type: "object", + properties: { count: { $ref: "#/$defs/count" } }, + required: ["count"], + $defs: { count: { type: "integer", minimum: 2 } }, + }; + const { definition } = await schemaTool(schema); + expect(validate(definition, { count: 3 })).toEqual({ count: 3 }); + expect(() => validate(definition, { count: 1.5 })).toThrow(/Validation failed/); +}); + +test("the agent loop rejects invalid MCP arguments before a remote call", async () => { + const { definition, server } = await schemaTool(); + const model = stepModel(); + let turn = 0; + const agent = new Agent({ + initialState: { model, tools: [wrapToolDefinition(definition)] }, + streamFn: () => { + const stream = createAssistantMessageEventStream(); + const toolCall = (id: string, args: Record): ToolCall => ({ + type: "toolCall", + id, + name: definition.name, + arguments: args, + }); + const content: AssistantMessage["content"] = + turn++ === 0 + ? [toolCall("invalid", { ...validArguments(), count: 1.5 }), toolCall("valid", validArguments())] + : [{ type: "text", text: "done" }]; + const message: AssistantMessage = { + role: "assistant", + content, + api: model.api, + model: model.id, + provider: model.provider, + stopReason: turn === 1 ? "toolUse" : "stop", + timestamp: Date.now(), + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + }; + stream.push({ type: "done", reason: message.stopReason as "stop" | "toolUse", message }); + stream.end(message); + return stream; + }, + }); + await agent.prompt("Check the arguments"); + expect(server.calls).toEqual([{ name: "check", arguments: validArguments() }]); + const invalid = agent.state.messages.find( + (message) => message.role === "toolResult" && message.toolCallId === "invalid", + ); + expect(invalid).toMatchObject({ + isError: true, + content: [{ type: "text", text: expect.stringContaining("Validation failed") }], + }); +}); + +function outputTool(name: string, valueType: "integer" | "string", description = name): McpTool { + return { + ...tool(name, description), + outputSchema: { + // Catalog generations may reuse a schema ID with a different definition. + $id: "urn:catalog-test:output", + type: "object", + properties: { value: { type: valueType } }, + required: ["value"], + }, + }; +} + +test.each(["session", "one-shot"])("%s validates an output schema from the first catalog page", async (kind) => { + const server = await catalogServer([ + { tools: [outputTool("first", "integer")], nextCursor: "1" }, + { tools: [tool("last")] }, + ]); + server.setCallHandler(async () => ({ content: [], structuredContent: { value: "invalid integer" } })); + if (kind === "session") { + config.value = { mcp_servers: { output: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const registered = harness.extension.tools.get("output__first")!.definition; + await expect(registered.execute("output", {}, undefined, undefined, undefined as never)).rejects.toThrow( + /output schema/, + ); + } else { + await expect( + invokeRemoteMcpTool({ serverName: "output", serverUrl: server.url, toolName: "first", arguments: {} }), + ).rejects.toThrow(/output schema/); + } + expect(server.requests).toEqual([undefined, "1"]); +}); + +test.each(["session", "one-shot"])( + "%s rejects required-task tools from the first catalog page before execution", + async (kind) => { + const server = await catalogServer([ + { tools: [{ ...tool("required"), execution: { taskSupport: "required" } }], nextCursor: "1" }, + { tools: [tool("last")] }, + ]); + if (kind === "session") { + config.value = { mcp_servers: { tasks: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const registered = harness.extension.tools.get("tasks__required")!.definition; + await expect(registered.execute("task", {}, undefined, undefined, undefined as never)).rejects.toThrow( + /requires task-based execution/, + ); + } else { + await expect( + invokeRemoteMcpTool({ serverName: "tasks", serverUrl: server.url, toolName: "required", arguments: {} }), + ).rejects.toThrow(/requires task-based execution/); + } + expect(server.calls).toEqual([]); + expect(server.requests).toEqual([undefined, "1"]); + }, +); + +test("a failed refresh cannot replace the output validation of the last good tool", async () => { + const server = await catalogServer([{ tools: [outputTool("check", "integer", "old")] }]); + config.value = { mcp_servers: { output: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const registered = harness.extension.tools.get("output__check")!.definition; + server.setListHandler((cursor) => { + if (cursor === undefined) return { tools: [outputTool("check", "string", "new")], nextCursor: "1" }; + throw new Error("second page failed"); + }); + await server.notify(); + await vi.waitFor(() => + expect(harness.notify).toHaveBeenCalledWith(expect.stringContaining("second page failed"), "warning"), + ); + expect(harness.extension.tools.get("output__check")!.definition).toBe(registered); + server.setCallHandler(async () => ({ content: [], structuredContent: { value: "new schema only" } })); + await expect(registered.execute("old", {}, undefined, undefined, undefined as never)).rejects.toThrow( + /output schema/, + ); +}); + +test("a pending call validates against its captured schema after a catalog update reuses the schema ID", async () => { + const server = await catalogServer([{ tools: [outputTool("check", "integer", "old")] }]); + config.value = { mcp_servers: { output: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const pending = gate(); + server.setCallHandler(async () => { + await pending.promise; + return { content: [], structuredContent: { value: 2 } }; + }); + const original = harness.extension.tools.get("output__check")!.definition; + const execution = original.execute("pending", {}, undefined, undefined, undefined as never); + void execution.catch(() => undefined); + await vi.waitFor(() => expect(server.calls).toHaveLength(1)); + server.setPages([{ tools: [outputTool("check", "string", "new")] }]); + await server.notify(); + await vi.waitFor(() => expect(harness.extension.tools.get("output__check")!.definition.description).toBe("new")); + pending.release(); + await expect(execution).resolves.toMatchObject({ details: { structuredContent: { value: 2 } } }); + const updated = harness.extension.tools.get("output__check")!.definition; + await expect(updated.execute("updated", {}, undefined, undefined, undefined as never)).rejects.toThrow( + /output schema/, + ); +}); + +test.each(["session", "one-shot"])( + "%s still requires structured output on successful schema-bearing calls", + async (kind) => { + const server = await catalogServer([{ tools: [outputTool("check", "integer")] }]); + if (kind === "session") { + config.value = { mcp_servers: { output: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const registered = harness.extension.tools.get("output__check")!.definition; + await expect(registered.execute("output", {}, undefined, undefined, undefined as never)).rejects.toThrow( + /structured content/, + ); + } else { + await expect( + invokeRemoteMcpTool({ serverName: "output", serverUrl: server.url, toolName: "check", arguments: {} }), + ).rejects.toThrow(/structured content/); + } + }, +); + +test("find_tools sees a late MCP catalog and stops returning tools removed by the server", async () => { + const initial = gate(); + const server = await catalogServer([]); + server.setListHandler(async (_cursor, request) => { + if (request === 1) { + await initial.promise; + return { + tools: [ + tool("calendar", "Read calendar events", { + type: "object", + properties: { calendar_id: { type: "string" } }, + required: ["calendar_id"], + }), + ], + }; + } + return { tools: [] }; + }); + config.value = { mcp_servers: { live: { url: server.url } } }; + const cwd = await mkdtemp(join(tmpdir(), "mcp-discovery-")); + cleanups.push(() => rm(cwd, { recursive: true, force: true })); + const extensionsResult = await createTestExtensionsResult([createStepMcpExtension()], cwd); + const { session } = await createAgentSession({ + cwd, + agentDir: join(cwd, "agent"), + model: stepModel(), + sessionManager: SessionManager.inMemory(), + settingsManager: SettingsManager.inMemory(), + resourceLoader: createTestResourceLoader({ extensionsResult }), + customTools: createStepToolProfile(cwd), + }); + const runner = session.extensionRunner; + cleanups.push(async () => { + await runner.emit({ type: "session_shutdown", reason: "quit" }); + session.dispose(); + }); + await session.bindExtensions({ mode: "tui" }); + const find = session.agent.state.tools.find((entry) => entry.name === "find_tools")!; + const search = async () => { + const result = await find.execute("find", { query: "calendar" }); + return result.content.map((block) => (block.type === "text" ? block.text : "")).join("\n"); + }; + expect(await search()).toBe("(no matching tools)"); + initial.release(); + await vi.waitFor(() => expect(session.getActiveToolNames()).toContain("live__calendar")); + expect(await search()).toContain("live__calendar"); + expect(await search()).toContain('"calendar_id"'); + await server.notify(); + await vi.waitFor(() => expect(session.getActiveToolNames()).not.toContain("live__calendar")); + expect(await search()).toBe("(no matching tools)"); +}); + +test("input schema preservation retains the provider's existing coercion and format limits", async () => { + const schema: McpTool["inputSchema"] = { + type: "object", + properties: { + count: { type: "integer" }, + email: { type: "string", format: "email" }, + custom: { type: "string", format: "unregistered-catalog-test-format" }, + }, + required: ["count", "email", "custom"], + }; + const { definition } = await schemaTool(schema); + expect(JSON.parse(JSON.stringify(definition.parameters))).toEqual(schema); + expect(validate(definition, { count: "2", email: "test@example.org", custom: "arbitrary" })).toEqual({ + count: 2, + email: "test@example.org", + custom: "arbitrary", + }); + expect(() => validate(definition, { count: 2, email: "not-an-email", custom: "arbitrary" })).toThrow( + /Validation failed/, + ); +}); + +test("a refresh with an unusable output schema keeps the old executable catalog", async () => { + const server = await catalogServer([{ tools: [outputTool("check", "integer", "old")] }]); + config.value = { mcp_servers: { output: { url: server.url } } }; + const harness = await setup(); + await harness.start(); + const original = harness.extension.tools.get("output__check")!.definition; + server.setPages([ + { + tools: [ + { + ...tool("check", "invalid"), + outputSchema: { type: "object", properties: { value: { $ref: "#/missing" } } }, + }, + ], + }, + ]); + await server.notify(); + await vi.waitFor(() => + expect(harness.notify).toHaveBeenCalledWith(expect.stringContaining("catalog refresh failed"), "warning"), + ); + expect(harness.extension.tools.get("output__check")!.definition).toBe(original); + expect(harness.batches).toHaveLength(1); + server.setCallHandler(async () => ({ content: [], structuredContent: { value: 2 } })); + await expect(original.execute("old", {}, undefined, undefined, undefined as never)).resolves.toMatchObject({ + details: { structuredContent: { value: 2 } }, + }); +}); diff --git a/packages/coding-agent/src/step/mcp-client.ts b/packages/coding-agent/src/step/mcp-client.ts index 7431b1fb..daca2bd1 100644 --- a/packages/coding-agent/src/step/mcp-client.ts +++ b/packages/coding-agent/src/step/mcp-client.ts @@ -1,6 +1,15 @@ import { Client } from "@modelcontextprotocol/sdk/client/index.js"; import { StreamableHTTPClientTransport } from "@modelcontextprotocol/sdk/client/streamableHttp.js"; -import { type CallToolResult, CallToolResultSchema, type Tool as McpTool } from "@modelcontextprotocol/sdk/types.js"; +import type { RequestOptions } from "@modelcontextprotocol/sdk/shared/protocol.js"; +import { + type CallToolResult, + CallToolResultSchema, + ErrorCode, + ListToolsResultSchema, + McpError, + type Tool as McpTool, +} from "@modelcontextprotocol/sdk/types.js"; +import { AjvJsonSchemaValidator } from "@modelcontextprotocol/sdk/validation/ajv"; import { STEPCODE_VERSION } from "./version.ts"; export const DEFAULT_MCP_TIMEOUT_MS = 30_000; @@ -30,15 +39,33 @@ export async function invokeRemoteMcpTool(input: RemoteMcpToolInvocation): Promi requestInit: input.headers ? { headers: input.headers } : undefined, }); const client = new Client({ name: "stepcode", version: STEPCODE_VERSION.value }, { capabilities: {} }); - + const closeClient = async () => { + try { + await client.close(); + } catch { + try { + await transport.close(); + } catch { + // Best-effort cleanup only. + } + } + }; + // SDK connect() awaits notifications/initialized without forwarding the + // request signal. Closing the transport also aborts that HTTP request. + const closeOnAbort = () => { + void closeClient(); + }; + signal.addEventListener("abort", closeOnAbort, { once: true }); try { + signal.throwIfAborted(); await client.connect(transport, { timeout: timeoutMs, signal }); const tools = await listAllMcpTools(client, signal, timeoutMs); - if (!tools.some((tool) => tool.name === input.toolName)) { + const tool = tools.find((tool) => tool.name === input.toolName); + if (!tool) { throw new Error(`Remote MCP server '${input.serverName}' does not expose tool '${input.toolName}'.`); } - const result = await client.callTool({ name: input.toolName, arguments: input.arguments }, CallToolResultSchema, { + const result = await createMcpToolCaller(client, tool)(input.arguments, { timeout: timeoutMs, resetTimeoutOnProgress: true, signal, @@ -49,15 +76,8 @@ export async function invokeRemoteMcpTool(input: RemoteMcpToolInvocation): Promi `MCP tool ${input.serverName}.${input.toolName} failed: ${error instanceof Error ? error.message : String(error)}`, ); } finally { - try { - await client.close(); - } catch { - try { - await transport.close(); - } catch { - // Best-effort cleanup only. - } - } + signal.removeEventListener("abort", closeOnAbort); + await closeClient(); } } @@ -93,20 +113,77 @@ function normalizeMcpCallToolResult(result: RawMcpCallToolResult): CallToolResul }); } -async function listAllMcpTools(client: Client, signal: AbortSignal, timeoutMs: number): Promise { +/** Collect a complete snapshot without mutating the SDK's per-page tool metadata cache. */ +export async function listAllMcpTools(client: Client, signal: AbortSignal, timeoutMs: number): Promise { + const deadline = AbortSignal.any([signal, AbortSignal.timeout(timeoutMs)]); const tools: McpTool[] = []; + const cursors = new Set(); let cursor: string | undefined; do { - const result = await client.listTools(cursor ? { cursor } : undefined, { - timeout: timeoutMs, - signal, - }); + deadline.throwIfAborted(); + // SDK 1.27.1 listTools() replaces output validators and task metadata on + // every page, even if a later page fails. Stage raw pages instead. + const result = await client.request( + { method: "tools/list", params: cursor === undefined ? undefined : { cursor } }, + ListToolsResultSchema, + { timeout: timeoutMs, signal: deadline }, + ); + deadline.throwIfAborted(); tools.push(...result.tools); cursor = result.nextCursor; - } while (cursor); + if (cursor !== undefined) { + if (cursors.has(cursor)) throw new Error("MCP tools/list returned a repeated pagination cursor."); + cursors.add(cursor); + } + } while (cursor !== undefined); return tools; } +/** Bind execution policy and output validation to this catalog definition, including in-flight calls. */ +export function createMcpToolCaller(client: Client, tool: McpTool) { + const name = tool.name; + const requiresTask = tool.execution?.taskSupport === "required"; + // Isolate schema IDs between tools and catalog generations. The SDK's + // default validator otherwise reuses an old schema with the same $id. + const validateOutput = tool.outputSchema ? new AjvJsonSchemaValidator().getValidator(tool.outputSchema) : undefined; + return async (args: Record, options: RequestOptions): Promise => { + options.signal?.throwIfAborted(); + if (requiresTask) { + throw new McpError( + ErrorCode.InvalidRequest, + `Tool "${name}" requires task-based execution, which this MCP client does not support.`, + ); + } + // callTool() looks up mutable SDK metadata after awaiting the response. + // Use the public request/validator APIs so catalog refreshes cannot alter + // the contract of a call that has already started. + const result = await client.request( + { method: "tools/call", params: { name, arguments: args } }, + CallToolResultSchema, + options, + ); + options.signal?.throwIfAborted(); + if (validateOutput) { + if (!result.structuredContent && !result.isError) { + throw new McpError( + ErrorCode.InvalidRequest, + `Tool ${name} has an output schema but did not return structured content`, + ); + } + if (result.structuredContent) { + const validation = validateOutput(result.structuredContent); + if (!validation.valid) { + throw new McpError( + ErrorCode.InvalidParams, + `Structured content does not match the tool's output schema: ${validation.errorMessage}`, + ); + } + } + } + return result; + }; +} + function isRecord(value: unknown): value is Record { return typeof value === "object" && value !== null && !Array.isArray(value); } diff --git a/packages/coding-agent/src/step/mcp.ts b/packages/coding-agent/src/step/mcp.ts index dda3b840..4b0bef20 100644 --- a/packages/coding-agent/src/step/mcp.ts +++ b/packages/coding-agent/src/step/mcp.ts @@ -3,12 +3,12 @@ import { UnauthorizedError } from "@modelcontextprotocol/sdk/client/auth.js"; import { Client } from "@modelcontextprotocol/sdk/client/index.js"; import { StdioClientTransport } from "@modelcontextprotocol/sdk/client/stdio.js"; import { StreamableHTTPClientTransport, StreamableHTTPError } from "@modelcontextprotocol/sdk/client/streamableHttp.js"; -import { CallToolResultSchema, type Tool as McpTool } from "@modelcontextprotocol/sdk/types.js"; +import { type Tool as McpTool, ToolListChangedNotificationSchema } from "@modelcontextprotocol/sdk/types.js"; import type { AgentToolResult } from "@step-harness/agent-core"; -import { type TSchema, Type } from "typebox"; import type { ExtensionAPI, ExtensionFactory } from "../core/extensions/types.ts"; import { theme } from "../theme/theme.ts"; import { readGlobalStepConfig } from "./config-toml.ts"; +import { createMcpToolCaller, listAllMcpTools } from "./mcp-client.ts"; import { resolveStepMcpEnvironment } from "./mcp-environment.ts"; import { createStoredMcpOAuthProvider, hasStoredMcpOAuthCredential } from "./mcp-oauth.ts"; import { @@ -35,7 +35,10 @@ interface ConnectedServer { readonly name: string; readonly client: Client; readonly transport: StdioClientTransport | StreamableHTTPClientTransport; - readonly tools: McpTool[]; + tools: McpTool[]; + /** Ends when this connection closes or its session shuts down. */ + readonly signal: AbortSignal; + readonly catalogTimeoutMs: number; /** Per-call timeout for this server, from `tool_timeout_sec`. */ readonly callTimeoutMs: number; } @@ -142,22 +145,76 @@ export function createStepMcpExtension(): ExtensionFactory { await Promise.all( discovered.map(async (item, index) => { let connected: ConnectedServer | undefined; + let published = false; + let pending = false; + let refreshing = false; + const refreshCatalog = async () => { + if (refreshing || !published || !connected || connected.signal.aborted) return; + const server = connected; + refreshing = true; + try { + while (pending && !server.signal.aborted) { + // Coalesce a notification burst before listing, with at most one + // more pass when notifications arrive during an in-flight list. + await yieldToEventLoop(); + pending = false; + try { + const tools = selectDeclaredTools( + await listAllMcpTools(server.client, server.signal, server.catalogTimeoutMs), + item.declaration, + ); + const remoteTools = tools.map((tool) => createRemoteTool(server, tool)); + await yieldToEventLoop(); + if (server.signal.aborted || !published) return; + pi.registerTools(remoteTools, { + remove: server.tools.map((tool) => remoteToolName(server, tool)), + }); + server.tools = tools; + statuses[index] = { name: item.name, status: "connected", toolCount: tools.length }; + } catch (error) { + if (server.signal.aborted || !published) return; + ctx.ui.notify( + `MCP server '${item.name}' catalog refresh failed; keeping the previous tools: ${error instanceof Error ? error.message : String(error)}`, + "warning", + ); + } + } + } finally { + refreshing = false; + } + }; + const onToolsChanged = () => { + pending = true; + void refreshCatalog(); + }; try { - const server = await connectStepMcpServer(item, controller.signal); + const server = await connectStepMcpServer(item, controller.signal, onToolsChanged); connected = server; const remoteTools = server.tools.map((tool) => createRemoteTool(server, tool)); // Publishing a server's catalog refreshes the registry and the // Step prompt once. Yield first so a server that finished while // the loop was busy cannot preempt input or rendering. await yieldToEventLoop(); - controller.signal.throwIfAborted(); + server.signal.throwIfAborted(); pi.registerTools(remoteTools); + published = true; servers.push(server); + server.signal.addEventListener( + "abort", + () => { + published = false; + if (controller.signal.aborted) return; + pi.registerTools([], { remove: server.tools.map((tool) => remoteToolName(server, tool)) }); + statuses[index] = { name: item.name, status: "failed", toolCount: 0 }; + }, + { once: true }, + ); statuses[index] = { name: item.name, status: "connected", toolCount: server.tools.length, }; + void refreshCatalog(); } catch (error) { if (connected) await closeStepMcpServer(connected); if (controller.signal.aborted) return; @@ -197,6 +254,9 @@ export function createStepMcpExtension(): ExtensionFactory { servers = []; startup = undefined; if (currentMcpStatuses === statuses) currentMcpStatuses = []; + pi.registerTools([], { + remove: closing.flatMap((server) => server.tools.map((tool) => remoteToolName(server, tool))), + }); await Promise.all(closing.map(closeStepMcpServer)); }); }; @@ -287,6 +347,7 @@ function normalizeDeclaration(value: Record): ServerDeclaration export async function connectStepMcpServer( input: DiscoveredServer, abortSignal?: AbortSignal, + onToolsChanged?: () => void, ): Promise { abortSignal?.throwIfAborted(); const timeout = timeoutMs(input.declaration.startup_timeout_sec, MCP_STARTUP_TIMEOUT_SEC); @@ -316,7 +377,11 @@ export async function connectStepMcpServer( throw new Error("MCP server must define command or url"); } const client = new Client(CLIENT_INFO, { capabilities: {} }); - const signal = AbortSignal.any([AbortSignal.timeout(timeout), ...(abortSignal ? [abortSignal] : [])]); + const closed = new AbortController(); + client.onclose = () => closed.abort(new Error("MCP connection closed")); + const lifetime = AbortSignal.any([closed.signal, ...(abortSignal ? [abortSignal] : [])]); + if (onToolsChanged) client.setNotificationHandler(ToolListChangedNotificationSchema, onToolsChanged); + const signal = AbortSignal.any([AbortSignal.timeout(timeout), lifetime]); const closeOnAbort = () => { void closeStepMcpServer({ client, transport }); }; @@ -324,13 +389,15 @@ export async function connectStepMcpServer( try { signal.throwIfAborted(); await client.connect(transport, { timeout, signal }); - const listed = await client.listTools(undefined, { timeout, signal }); + const tools = await listAllMcpTools(client, signal, timeout); signal.throwIfAborted(); return { name: input.name, client, transport, - tools: selectDeclaredTools(listed.tools, input.declaration), + tools: selectDeclaredTools(tools, input.declaration), + signal: lifetime, + catalogTimeoutMs: timeout, callTimeoutMs, }; } catch (error) { @@ -455,23 +522,30 @@ export function convertMcpCallResult(serverName: string, toolName: string, resul }; } +function remoteToolName(server: Pick, remote: McpTool): string { + return `${server.name}__${sanitizeName(remote.name)}`; +} + function createRemoteTool(server: ConnectedServer, remote: McpTool) { - const name = `${server.name}__${sanitizeName(remote.name)}`; + const name = remoteToolName(server, remote); + const call = createMcpToolCaller(server.client, remote); return { name, label: remote.title?.trim() || remote.name, description: remote.description?.trim() || `MCP tool '${remote.name}' from server '${server.name}'.`, - parameters: schemaFromJson(remote.inputSchema), + // The provider boundary accepts JSON Schema directly. Reconstructing it + // as TypeBox types loses constraints, local references and combinators. + parameters: remote.inputSchema, execute: async ( _toolCallId: string, params: unknown, signal: AbortSignal | undefined, ): Promise> => { - const result = (await server.client.callTool( - { name: remote.name, arguments: isRecord(params) ? params : {} }, - CallToolResultSchema, - { timeout: server.callTimeoutMs, resetTimeoutOnProgress: true, signal }, - )) as McpCallResult; + const result = await call(isRecord(params) ? params : {}, { + timeout: server.callTimeoutMs, + resetTimeoutOnProgress: true, + signal: signal ? AbortSignal.any([signal, server.signal]) : server.signal, + }); return convertMcpCallResult(server.name, remote.name, result); }, }; @@ -482,35 +556,6 @@ export function appendStepPageManagementHint(serverName: string, toolName: strin return `${text}\n\nTo manage your deployed pages, visit ${STEPPAGE_MANAGEMENT_URL}`; } -function schemaFromJson(schema: unknown): TSchema { - if (!isRecord(schema) || !isRecord(schema.properties)) return Type.Object({}, { additionalProperties: true }); - const properties: Record = {}; - for (const [key, value] of Object.entries(schema.properties)) { - const property = schemaValueToTypeBox(value); - properties[key] = - Array.isArray(schema.required) && schema.required.includes(key) ? property : Type.Optional(property); - } - return Type.Object(properties, { additionalProperties: true }); -} - -function schemaValueToTypeBox(value: unknown): TSchema { - if (!isRecord(value)) return Type.Unknown(); - if (Array.isArray(value.enum) && value.enum.length > 0) { - const literals = value.enum.filter( - (item): item is string | number | boolean => - typeof item === "string" || typeof item === "number" || typeof item === "boolean", - ); - if (literals.length === 1) return Type.Literal(literals[0]); - if (literals.length > 1) return Type.Union(literals.map((item) => Type.Literal(item))); - } - if (value.type === "array") return Type.Array(schemaValueToTypeBox(value.items)); - if (value.type === "object" && isRecord(value.properties)) return schemaFromJson(value); - if (value.type === "boolean") return Type.Boolean(); - if (value.type === "number" || value.type === "integer") return Type.Number(); - if (value.type === "string") return Type.String(); - return Type.Unknown(); -} - function sanitizeName(value: string): string { const normalized = value.replace(/[^a-zA-Z0-9_]+/gu, "_").replace(/^_+|_+$/gu, ""); return normalized || "tool"; diff --git a/packages/coding-agent/src/step/tool-profile.ts b/packages/coding-agent/src/step/tool-profile.ts index c28e01ea..708f514b 100644 --- a/packages/coding-agent/src/step/tool-profile.ts +++ b/packages/coding-agent/src/step/tool-profile.ts @@ -1113,7 +1113,7 @@ function countOccurrences(text: string, needle: string): number { } function createFindToolsDefinition( - definitions: readonly AnyToolDefinition[], + fallbackDefinitions: readonly AnyToolDefinition[], ): ToolDefinition { return { name: "find_tools", @@ -1121,7 +1121,8 @@ function createFindToolsDefinition( description: "Search registered tools by natural-language intent, tool name, description, and parameter names.", promptSnippet: "Find a tool by describing the operation you need", parameters: findToolsSchema, - execute: async (_toolCallId, args: FindToolsInput) => { + execute: async (_toolCallId, args: FindToolsInput, _signal, _onUpdate, ctx) => { + const definitions = ctx?.getToolCatalog?.() ?? fallbackDefinitions; const queryTokens = args.query.toLowerCase().match(/[\p{L}\p{N}_]+/gu) ?? []; const limit = Math.max(1, Math.min(20, args.limit ?? 8)); const matches = definitions @@ -1140,7 +1141,7 @@ function createFindToolsDefinition( ? matches .map( ({ definition, score }, index) => - `${index + 1}. ${definition.name} [score=${score}]\ndescription: ${definition.description}`, + `${index + 1}. ${definition.name} [score=${score}]\ndescription: ${definition.description}\nparameters: ${JSON.stringify(definition.parameters)}`, ) .join("\n\n") : "(no matching tools)"; diff --git a/packages/coding-agent/test/extensions-tool-catalog.test.ts b/packages/coding-agent/test/extensions-tool-catalog.test.ts new file mode 100644 index 00000000..bba206b2 --- /dev/null +++ b/packages/coding-agent/test/extensions-tool-catalog.test.ts @@ -0,0 +1,217 @@ +import { mkdtemp, rm } from "node:fs/promises"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { Type } from "typebox"; +import { afterEach, expect, test, vi } from "vitest"; +import type { ExtensionAPI, ExtensionContext, ExtensionFactory, ToolDefinition } from "../src/core/extensions/types.ts"; +import { createAgentSession } from "../src/core/sdk.ts"; +import { SessionManager } from "../src/core/session-manager.ts"; +import { SettingsManager } from "../src/core/settings-manager.ts"; +import { createStepToolProfile, stepToolNames } from "../src/step/tool-profile.ts"; +import { createTestExtensionsResult, createTestResourceLoader, stepModel } from "./utilities.ts"; + +const cleanups: Array<() => Promise | void> = []; +afterEach(async () => { + for (const cleanup of cleanups.splice(0).reverse()) await cleanup(); +}); + +function definition(name: string, description = name) { + return { + name, + label: name, + description, + parameters: Type.Object({}), + execute: async () => ({ content: [{ type: "text", text: description }], details: {} }), + } satisfies ToolDefinition; +} + +async function setup(factories: ExtensionFactory[], stepProfile = false) { + const cwd = await mkdtemp(join(tmpdir(), "extension-catalog-")); + cleanups.push(() => rm(cwd, { recursive: true, force: true })); + const extensionsResult = await createTestExtensionsResult(factories, cwd); + const { session } = await createAgentSession({ + cwd, + agentDir: join(cwd, "agent"), + model: stepModel(), + sessionManager: SessionManager.inMemory(), + settingsManager: SettingsManager.inMemory(), + resourceLoader: createTestResourceLoader({ extensionsResult }), + ...(stepProfile ? { customTools: createStepToolProfile(cwd) } : {}), + }); + cleanups.push(() => session.dispose()); + await session.bindExtensions({}); + if (stepProfile) session.setActiveToolsByName([...stepToolNames]); + return { session, ...extensionsResult }; +} + +function resultText(result: { content: Array<{ type: string; text?: string }> }) { + return result.content.map((block) => (block.type === "text" ? block.text : "")).join("\n"); +} + +test("a replacement batch removes only its owner's tools and refreshes the session once", async () => { + let owner!: ExtensionAPI; + const { session, runtime, extensions } = await setup([ + (pi) => { + owner = pi; + pi.registerTools([definition("keep", "old"), definition("remove")]); + }, + (pi) => { + pi.registerTool(definition("other")); + }, + ]); + const refresh = vi.fn(runtime.refreshTools); + runtime.refreshTools = refresh; + owner.registerTools([definition("keep", "updated"), definition("add")], { remove: ["remove", "other", "read"] }); + expect([...extensions[0].tools.keys()].sort()).toEqual(["add", "keep"]); + expect(extensions[1].tools.has("other")).toBe(true); + expect(session.getActiveToolNames()).toEqual(expect.arrayContaining(["keep", "add", "other", "read"])); + expect(session.getActiveToolNames()).not.toContain("remove"); + expect(session.getAllTools().find((tool) => tool.name === "keep")?.description).toBe("updated"); + expect(refresh).toHaveBeenCalledTimes(1); +}); + +test("removing the winning extension registration reveals the next extension overlay", async () => { + let first!: ExtensionAPI; + let second!: ExtensionAPI; + const { session, runtime } = await setup([ + (pi) => { + first = pi; + pi.registerTool(definition("shared", "first")); + }, + (pi) => { + second = pi; + pi.registerTool(definition("shared", "second")); + }, + ]); + expect(session.getAllTools().find((tool) => tool.name === "shared")?.description).toBe("first"); + const refresh = vi.fn(runtime.refreshTools); + runtime.refreshTools = refresh; + first.registerTools([], { remove: ["shared"] }); + expect(session.getAllTools().find((tool) => tool.name === "shared")?.description).toBe("second"); + expect(session.getActiveToolNames()).toContain("shared"); + const callable = session.agent.state.tools.find((tool) => tool.name === "shared")!; + expect(resultText(await callable.execute("overlay", {}))).toBe("second"); + expect(refresh).toHaveBeenCalledTimes(1); + second.registerTools([], { remove: ["shared"] }); + expect(session.getAllTools().some((tool) => tool.name === "shared")).toBe(false); + expect(session.getActiveToolNames()).not.toContain("shared"); +}); + +test("removing a shadowed registration leaves the current winner active", async () => { + let shadowed!: ExtensionAPI; + const { session, extensions } = await setup([ + (pi) => pi.registerTool(definition("shared", "winner")), + (pi) => { + shadowed = pi; + pi.registerTool(definition("shared", "shadowed")); + }, + ]); + shadowed.registerTools([], { remove: ["shared"] }); + expect(extensions[1].tools.has("shared")).toBe(false); + expect(session.getAllTools().find((tool) => tool.name === "shared")?.description).toBe("winner"); + expect(session.getActiveToolNames()).toContain("shared"); +}); + +test("removal can reveal a builtin override without removing another owner's definition", async () => { + let owner!: ExtensionAPI; + const { session } = await setup([ + (pi) => { + owner = pi; + pi.registerTool(definition("read", "overridden")); + }, + ]); + expect(session.getAllTools().find((tool) => tool.name === "read")?.description).toBe("overridden"); + owner.registerTools([], { remove: ["read"] }); + expect(session.getAllTools().find((tool) => tool.name === "read")?.sourceInfo.source).toBe("builtin"); + expect(session.getActiveToolNames()).toContain("read"); +}); + +test("empty or non-owned removals do not refresh, and stale runtimes cannot remove tools", async () => { + let owner!: ExtensionAPI; + const { runtime, extensions } = await setup([ + (pi) => { + owner = pi; + pi.registerTool(definition("owned")); + }, + ]); + const refresh = vi.fn(runtime.refreshTools); + runtime.refreshTools = refresh; + owner.registerTools([], { remove: ["missing", "read"] }); + owner.registerTools([], { remove: [] }); + expect(refresh).not.toHaveBeenCalled(); + runtime.invalidate(); + expect(() => owner.registerTools([], { remove: ["owned"] })).toThrow(/stale/); + expect(extensions[0].tools.has("owned")).toBe(true); +}); + +test("the optional context catalog accessor reads active tools live and rejects stale contexts", async () => { + let api!: ExtensionAPI; + let ctx!: ExtensionContext; + const { session, runtime } = await setup([ + (pi) => { + api = pi; + pi.on("session_start", (_event, context) => { + ctx = context; + }); + }, + ]); + expect(typeof ctx.getToolCatalog).toBe("function"); + expect(ctx.getToolCatalog!().map((tool) => tool.name)).toContain("read"); + api.registerTool(definition("late", "late catalog entry")); + expect(ctx.getToolCatalog!()).toContainEqual( + expect.objectContaining({ name: "late", description: "late catalog entry" }), + ); + session.setActiveToolsByName(["read"]); + expect(ctx.getToolCatalog!().map((tool) => tool.name)).toEqual(["read"]); + runtime.invalidate(); + expect(() => ctx.getToolCatalog!()).toThrow(/stale/); +}); + +test("find_tools discovers late tools, callable schemas, and updated descriptions through a real session", async () => { + let api!: ExtensionAPI; + const { session } = await setup( + [ + (pi) => { + api = pi; + }, + ], + true, + ); + const find = session.agent.state.tools.find((tool) => tool.name === "find_tools")!; + expect(resultText(await find.execute("before", { query: "calendar" }))).toBe("(no matching tools)"); + const parameters = Type.Object({ calendar_id: Type.String(), count: Type.Optional(Type.Integer({ minimum: 1 })) }); + api.registerTool({ ...definition("plugin__calendar", "Query calendar events"), parameters }); + const discovered = resultText(await find.execute("after", { query: "calendar_id" })); + expect(discovered).toContain("plugin__calendar"); + expect(discovered).toContain("Query calendar events"); + expect(discovered).toContain(JSON.stringify(parameters)); + api.registerTool({ + ...definition("plugin__calendar", "Query revised calendar events"), + parameters: Type.Object({ date: Type.String() }), + }); + const updated = resultText(await find.execute("updated", { query: "calendar" })); + expect(updated).toContain("Query revised calendar events"); + expect(updated).toContain('"date"'); + expect(updated).not.toContain("calendar_id"); +}); + +test("find_tools excludes inactive and removed tools without falling back to captured builtins", async () => { + let api!: ExtensionAPI; + const { session } = await setup( + [ + (pi) => { + api = pi; + }, + ], + true, + ); + api.registerTool(definition("plugin__calendar", "Query calendar events")); + const find = session.agent.state.tools.find((tool) => tool.name === "find_tools")!; + session.setActiveToolsByName(["find_tools"]); + expect(resultText(await find.execute("inactive", { query: "calendar" }))).toBe("(no matching tools)"); + expect(resultText(await find.execute("builtin-inactive", { query: "read_file" }))).toBe("(no matching tools)"); + session.setActiveToolsByName(["find_tools", "plugin__calendar"]); + expect(resultText(await find.execute("active", { query: "calendar" }))).toContain("plugin__calendar"); + api.registerTools([], { remove: ["plugin__calendar"] }); + expect(resultText(await find.execute("removed", { query: "calendar" }))).toBe("(no matching tools)"); +}); From ddb051be8dee1099735eb576fdb8a011726be6e3 Mon Sep 17 00:00:00 2001 From: longyongshen Date: Mon, 28 Sep 2026 18:53:20 +0800 Subject: [PATCH 2/2] fix(providers): keep local schema definitions in Anthropic tool input The non-strict Anthropic tool path keeps only `properties` and `required`, so a schema whose properties point at root `$defs` or `definitions` (as pydantic-based MCP servers emit for enums and nested models) was sent with dangling `$ref`s. MCP input schemas now reach the provider unchanged, so carry those local definitions along with the kept properties. --- .../providers/src/api/anthropic-messages.ts | 10 ++++++- .../anthropic-eager-tool-input-compat.test.ts | 27 +++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/packages/providers/src/api/anthropic-messages.ts b/packages/providers/src/api/anthropic-messages.ts index a5742e8b..f69af3d4 100644 --- a/packages/providers/src/api/anthropic-messages.ts +++ b/packages/providers/src/api/anthropic-messages.ts @@ -1254,11 +1254,19 @@ function convertTools( return tools.map((tool, index) => { const strict = resolveJsonSchemaStrictSampling(tool, supportsStrictTools); const parameters = getJsonSchemaToolParameters(tool, strict); - const schema = parameters as { properties?: unknown; required?: string[] }; + const schema = parameters as { + properties?: unknown; + required?: string[]; + $defs?: unknown; + definitions?: unknown; + }; const legacyInputSchema = { type: "object" as const, properties: schema.properties ?? {}, required: schema.required ?? [], + // Keep local definitions so `$ref`s inside the kept properties still resolve. + ...(schema.$defs !== undefined ? { $defs: schema.$defs } : {}), + ...(schema.definitions !== undefined ? { definitions: schema.definitions } : {}), }; const inputSchema = strict === true diff --git a/packages/providers/test/anthropic-eager-tool-input-compat.test.ts b/packages/providers/test/anthropic-eager-tool-input-compat.test.ts index 37280cb4..5aa5042e 100644 --- a/packages/providers/test/anthropic-eager-tool-input-compat.test.ts +++ b/packages/providers/test/anthropic-eager-tool-input-compat.test.ts @@ -37,6 +37,21 @@ const schemaCompatibilityTool: Tool = { parameters: Type.Object({ value: Type.String() }, { additionalProperties: false, title: "LookupInput" }), }; +// Shape emitted by pydantic-based MCP servers (for example FastMCP) for enum and nested-model parameters. +const localReferenceTool: Tool = { + ...tool, + parameters: { + type: "object", + properties: { + priority: { $ref: "#/$defs/Priority" }, + owner: { $ref: "#/definitions/Owner" }, + }, + required: ["priority"], + $defs: { Priority: { type: "string", enum: ["low", "high"] } }, + definitions: { Owner: { type: "object", properties: { name: { type: "string" } } } }, + } as unknown as Tool["parameters"], +}; + const strictTool: Tool = { ...tool, parameters: Type.Object( @@ -163,4 +178,16 @@ describe("Anthropic eager tool input streaming compatibility", () => { title: "StrictLookupInput", }); }); + + it("keeps local definitions referenced by legacy input schemas", async () => { + const request = await captureAnthropicRequest(undefined, createContext([localReferenceTool])); + const parameters = localReferenceTool.parameters as Record; + expect(getFirstToolInputSchema(request.body)).toEqual({ + type: "object", + properties: parameters.properties, + required: ["priority"], + $defs: parameters.$defs, + definitions: parameters.definitions, + }); + }); });