Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 83 additions & 1 deletion packages/mcp/src/server.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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<void>((resolve) => {
notifyExecutorStarted = resolve
})
let notifyExecutorStopped: (() => void) | undefined
const executorStopped = new Promise<void>((resolve) => {
notifyExecutorStopped = resolve
})
const server = createPascalMcpServer({
bridge,
executeTool: async ({ name, signal, execute }) => {
if (name !== 'cancel_probe') return execute()
notifyExecutorStarted?.()
await new Promise<void>((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()
}
})
})
21 changes: 21 additions & 0 deletions packages/mcp/src/server.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import { version } from './version'

export type PascalMcpToolExecutor = <Result>(input: {
name: string
signal: AbortSignal
execute: () => Promise<Result>
}) => Promise<Result>

Expand Down Expand Up @@ -45,7 +46,27 @@ function installToolExecutor(server: McpServer, executeTool: PascalMcpToolExecut
registerTool(name, config, ((...args: Parameters<typeof callback>) =>
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
}
Loading