diff --git a/packages/webmcp-bridge/src/index.ts b/packages/webmcp-bridge/src/index.ts index 255c62b..4d4ff5e 100644 --- a/packages/webmcp-bridge/src/index.ts +++ b/packages/webmcp-bridge/src/index.ts @@ -198,6 +198,14 @@ export async function createWebMcpBridge( >(); let closed = false; + async function close(): Promise { + if (closed) return; + closed = true; + for (const entry of registered.values()) entry.controller.abort(); + registered.clear(); + await client.close(); + } + async function register(tool: Tool): Promise { const taken = new Set( [...registered.values()].map((entry) => entry.bridged.name), @@ -263,7 +271,14 @@ export async function createWebMcpBridge( }); } - await scheduleSync(); + try { + await scheduleSync(); + } catch (error) { + // No bridge is returned on initialization failure, so the caller cannot + // release its resources. Preserve the original error if cleanup fails too. + await close().catch(() => {}); + throw error; + } return { get tools() { @@ -272,12 +287,6 @@ export async function createWebMcpBridge( get active() { return !closed; }, - async close() { - if (closed) return; - closed = true; - for (const entry of registered.values()) entry.controller.abort(); - registered.clear(); - await client.close(); - }, + close, }; } diff --git a/packages/webmcp-bridge/test/bridge.test.ts b/packages/webmcp-bridge/test/bridge.test.ts index 55bf10d..e383c55 100644 --- a/packages/webmcp-bridge/test/bridge.test.ts +++ b/packages/webmcp-bridge/test/bridge.test.ts @@ -1,6 +1,7 @@ -import { describe, expect, it } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"; import { InMemoryTransport } from "@modelcontextprotocol/sdk/inMemory.js"; +import { ListToolsRequestSchema } from "@modelcontextprotocol/sdk/types.js"; import { z } from "zod"; import { createWebMcpBridge } from "../src/index.js"; import type { ModelContextLike, ModelContextTool } from "../src/types.js"; @@ -145,6 +146,70 @@ describe("createWebMcpBridge", () => { expect(bridge.active).toBe(false); }); + it.each([false, true])( + "closes the connection when initial discovery fails (close rejects: %s)", + async (closeRejects) => { + const { server, clientTransport } = await makeServer(); + server.server.setRequestHandler(ListToolsRequestSchema, () => { + throw new Error("initial tools/list failed"); + }); + const originalClose = clientTransport.close.bind(clientTransport); + const close = vi.spyOn(clientTransport, "close"); + if (closeRejects) { + close.mockImplementationOnce(async () => { + await originalClose(); + throw new Error("transport cleanup failed"); + }); + } + const { mc, tools } = fakeModelContext(); + + try { + await expect( + createWebMcpBridge({ transport: clientTransport, modelContext: mc }), + ).rejects.toThrow("initial tools/list failed"); + expect(close).toHaveBeenCalled(); + await expect( + clientTransport.send({ jsonrpc: "2.0", method: "notifications/initialized" }), + ).rejects.toThrow("Not connected"); + expect(tools.size).toBe(0); + } finally { + close.mockRestore(); + await originalClose(); + await server.close(); + } + }, + ); + + it("unregisters partial tools when an initialization error handler throws", async () => { + const { server, clientTransport } = await makeServer(); + const { mc, tools } = fakeModelContext(); + const registerTool = mc.registerTool.bind(mc); + const registrationError = new Error("tool registration failed"); + mc.registerTool = async (tool, options) => { + if (tool.name === "delete_everything") throw registrationError; + await registerTool(tool, options); + }; + const close = vi.spyOn(clientTransport, "close"); + + try { + await expect( + createWebMcpBridge({ + transport: clientTransport, + modelContext: mc, + onRegisterError: (_name, error) => { + throw error; + }, + }), + ).rejects.toBe(registrationError); + expect(tools.size).toBe(0); + expect(close).toHaveBeenCalled(); + } finally { + close.mockRestore(); + await clientTransport.close(); + await server.close(); + } + }); + it("re-syncs when the server's tool list changes", async () => { const { server, clientTransport } = await makeServer(); const { mc, tools } = fakeModelContext();