From 6ac73dac6a35f1de8cb9725c02be67666a6275e8 Mon Sep 17 00:00:00 2001 From: Aymeric Rabot Date: Wed, 9 Sep 2026 15:01:30 -0400 Subject: [PATCH] fix(mcp): propagate tool cancellation --- packages/mcp/src/server.test.ts | 84 ++++++++++++++++++++++++++++++++- packages/mcp/src/server.ts | 21 +++++++++ 2 files changed, 104 insertions(+), 1 deletion(-) diff --git a/packages/mcp/src/server.test.ts b/packages/mcp/src/server.test.ts index cadb9263b..232970725 100644 --- a/packages/mcp/src/server.test.ts +++ b/packages/mcp/src/server.test.ts @@ -11,8 +11,9 @@ describe('Pascal MCP tool execution', () => { const events: string[] = [] const server = createPascalMcpServer({ bridge, - executeTool: async ({ name, execute }) => { + executeTool: async ({ name, signal, execute }) => { events.push(`before:${name}`) + expect(signal).toBeInstanceOf(AbortSignal) const result = await execute() events.push(`after:${name}`) return result @@ -31,4 +32,85 @@ describe('Pascal MCP tool execution', () => { await server.close() } }) + + test('passes request cancellation through the executor before invoking a tool', async () => { + const bridge = new SceneBridge() + bridge.loadDefault() + let callbackCalls = 0 + let notifyExecutorStarted: (() => void) | undefined + const executorStarted = new Promise((resolve) => { + notifyExecutorStarted = resolve + }) + let notifyExecutorStopped: (() => void) | undefined + const executorStopped = new Promise((resolve) => { + notifyExecutorStopped = resolve + }) + const server = createPascalMcpServer({ + bridge, + executeTool: async ({ name, signal, execute }) => { + if (name !== 'cancel_probe') return execute() + notifyExecutorStarted?.() + await new Promise((resolve) => { + if (signal.aborted) resolve() + else signal.addEventListener('abort', () => resolve(), { once: true }) + }) + try { + signal.throwIfAborted() + return await execute() + } finally { + notifyExecutorStopped?.() + } + }, + }) + server.registerTool('cancel_probe', { inputSchema: {} }, async () => { + callbackCalls++ + return { content: [{ type: 'text', text: 'mutated' }] } + }) + const [serverTransport, clientTransport] = InMemoryTransport.createLinkedPair() + const client = new Client({ name: 'tool-cancellation-test', version: '0.0.0' }) + await Promise.all([server.connect(serverTransport), client.connect(clientTransport)]) + + try { + const controller = new AbortController() + const call = client.callTool({ name: 'cancel_probe', arguments: {} }, undefined, { + signal: controller.signal, + }) + await executorStarted + controller.abort() + await expect(call).rejects.toThrow() + await executorStopped + expect(callbackCalls).toBe(0) + } finally { + await client.close() + await server.close() + } + }) + + test('wraps tools registered through the deprecated tool surface', async () => { + const bridge = new SceneBridge() + bridge.loadDefault() + const executed: string[] = [] + const server = createPascalMcpServer({ + bridge, + executeTool: async ({ name, execute }) => { + executed.push(name) + return execute() + }, + }) + server.tool('legacy_probe', async () => ({ + content: [{ type: 'text', text: 'legacy result' }], + })) + const [serverTransport, clientTransport] = InMemoryTransport.createLinkedPair() + const client = new Client({ name: 'legacy-tool-executor-test', version: '0.0.0' }) + await Promise.all([server.connect(serverTransport), client.connect(clientTransport)]) + + try { + const result = await client.callTool({ name: 'legacy_probe', arguments: {} }) + expect(result.isError).toBeFalsy() + expect(executed).toEqual(['legacy_probe']) + } finally { + await client.close() + await server.close() + } + }) }) diff --git a/packages/mcp/src/server.ts b/packages/mcp/src/server.ts index 4424dec82..b03c530bf 100644 --- a/packages/mcp/src/server.ts +++ b/packages/mcp/src/server.ts @@ -10,6 +10,7 @@ import { version } from './version' export type PascalMcpToolExecutor = (input: { name: string + signal: AbortSignal execute: () => Promise }) => Promise @@ -45,7 +46,27 @@ function installToolExecutor(server: McpServer, executeTool: PascalMcpToolExecut registerTool(name, config, ((...args: Parameters) => executeTool({ name, + signal: toolRequestSignal(args), execute: () => Promise.resolve(Reflect.apply(callback, undefined, args)), })) as typeof callback) server.registerTool = wrappedRegisterTool + + const tool = server.tool.bind(server) + server.tool = ((name: string, ...args: unknown[]) => { + const callback = args.at(-1) + if (typeof callback !== 'function') { + return Reflect.apply(tool, undefined, [name, ...args]) + } + args[args.length - 1] = (...callbackArgs: unknown[]) => + executeTool({ + name, + signal: toolRequestSignal(callbackArgs), + execute: () => Promise.resolve(Reflect.apply(callback, undefined, callbackArgs)), + }) + return Reflect.apply(tool, undefined, [name, ...args]) + }) as McpServer['tool'] +} + +function toolRequestSignal(args: readonly unknown[]): AbortSignal { + return (args.at(-1) as { signal: AbortSignal }).signal }