Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/web/__tests__/mcp.test.ts

Raw
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,
			);
		});
	},
);