import { afterEach, describe, expect, it, vi } from "vitest"; import { parallelSearch } from "../parallel.js"; import { zaiSearch } from "../zai.js"; afterEach(() => { vi.restoreAllMocks(); vi.unstubAllEnvs(); }); // Exercise the shared transport through both real provider adapters. describe.each(["parallel", "zai"] as const)( "%s MCP transport contract", (provider) => { function search(signal?: AbortSignal, timeout = 1000) { vi.stubEnv("ZAI_API_KEY", "test-key"); return provider === "parallel" ? parallelSearch("docs", signal, 1, timeout) : zaiSearch("docs", undefined, signal, { count: 1, timeoutMs: timeout, }); } it.each([false, true])( "handles JSON/SSE=%s, negotiation, session rotation, and full text", async (sse) => { const snippet = "full excerpt ".repeat(3000); const text = provider === "parallel" ? JSON.stringify({ results: [ { title: "Docs", url: "https://example.com", excerpts: [snippet], }, ], }) : JSON.stringify([ { title: "Docs", link: "https://example.com", content: snippet, }, ]); const network = vi .spyOn(globalThis, "fetch") .mockImplementation(async (_url, init) => { const request = JSON.parse(String(init?.body)); const headers = new Headers(init?.headers); expect(init?.redirect).toBe("error"); expect(headers.get("Authorization")).toBe( provider === "zai" ? "Bearer test-key" : null, ); if (request.method !== "initialize") expect(headers.get("MCP-Protocol-Version")).toBe("2025-03-26"); if (request.method === "notifications/initialized") { expect(request).not.toHaveProperty("id"); expect(headers.get("mcp-session-id")).toBe("first"); return new Response(null, { status: 204, headers: { "mcp-session-id": "second" }, }); } if (request.method === "tools/call") expect(headers.get("mcp-session-id")).toBe("second"); const result = request.method === "initialize" ? { protocolVersion: "2025-03-26" } : { content: [{ type: "text", text }] }; const body = sse ? `: keepalive\r\n\r\ndata: {"method":"notifications/progress"}\r\n\r\ndata: {"id":99,"result":{}}\r\n\r\ndata: {"id":${request.id},\r\ndata: "result":${JSON.stringify(result)}}\r\n\r\ndata: [DONE]\r\n\r\n` : JSON.stringify({ id: request.id, result }); return new Response(body, { headers: { "mcp-session-id": "first" }, }); }); expect((await search())[0].snippet).toBe(snippet); expect(network).toHaveBeenCalledTimes(3); }, ); it.each([ ['{"id":99,"result":{}}', "no matching MCP response"], ['{"id":1,"result":{}}', "initialization failed"], ['{"id":1,"result":{"protocolVersion":42}}', "initialization failed"], ["data: {invalid}\n\n", "Malformed MCP response"], [ '{"id":1,"error":{"code":-32600,"message":"bad request"}}', "bad request", ], ])("rejects invalid initialization: %s", async (body, error) => { const network = vi .spyOn(globalThis, "fetch") .mockResolvedValue(new Response(body)); await expect(search()).rejects.toThrow(error); expect(network).toHaveBeenCalledTimes(1); }); it.each(["notifications/initialized", "tools/call"])( "stops on HTTP failures at %s", async (method) => { const network = vi .spyOn(globalThis, "fetch") .mockImplementation(async (_url, init) => { const request = JSON.parse(String(init?.body)); if (request.method === method) return new Response("unavailable", { status: 503 }); if (request.method === "initialize") return new Response( '{"id":1,"result":{"protocolVersion":"2025-03-26"}}', ); return new Response(null, { status: 202 }); }); await expect(search()).rejects.toThrow(/503/); expect(network).toHaveBeenCalledTimes(method === "tools/call" ? 3 : 2); }, ); it("rejects an invalid tool result", async () => { vi.spyOn(globalThis, "fetch").mockImplementation(async (_url, init) => { const request = JSON.parse(String(init?.body)); if (request.method === "initialize") return new Response( '{"id":1,"result":{"protocolVersion":"2025-03-26"}}', ); if (request.method === "notifications/initialized") return new Response(null, { status: 202 }); return new Response('{"id":2,"result":null}'); }); await expect(search()).rejects.toThrow("invalid result"); }); it("does not reuse a session across searches", async () => { const network = vi .spyOn(globalThis, "fetch") .mockImplementation(async (_url, init) => { const request = JSON.parse(String(init?.body)); if (request.method === "initialize") { expect(new Headers(init?.headers).has("mcp-session-id")).toBe( false, ); return new Response( JSON.stringify({ id: 1, result: { protocolVersion: "2025-03-26" }, }), { headers: { "mcp-session-id": "s" } }, ); } if (request.method === "notifications/initialized") return new Response(null, { status: 202 }); return new Response( JSON.stringify({ id: 2, result: { content: [ { type: "text", text: provider === "parallel" ? '{"results":[]}' : "[]", }, ], }, }), ); }); await search(); await search(); expect(network).toHaveBeenCalledTimes(6); }); it("stops before I/O on caller cancellation", async () => { const network = vi.spyOn(globalThis, "fetch"); await expect( search(AbortSignal.abort(new Error("cancelled"))), ).rejects.toThrow("cancelled"); expect(network).not.toHaveBeenCalled(); }); it("stops between initialization and notification on cancellation", async () => { const abort = new AbortController(); const network = vi .spyOn(globalThis, "fetch") .mockImplementation(async () => { abort.abort(new Error("cancelled")); return new Response( '{"id":1,"result":{"protocolVersion":"2025-03-26"}}', ); }); await expect(search(abort.signal)).rejects.toThrow("cancelled"); expect(network).toHaveBeenCalledTimes(1); }); it("enforces the deadline with a caller signal", async () => { vi.spyOn(globalThis, "fetch").mockImplementation( (_url, init) => new Promise((_resolve, reject) => { init?.signal?.addEventListener( "abort", () => reject(init.signal?.reason), { once: true }, ); }), ); await expect(search(new AbortController().signal, 10)).rejects.toThrow( /timeout/i, ); }); }, );