repositories / pi-ext
pi-ext
bugabingas pi extensions
owned by admin
extensions/web/__tests__/mcp.test.ts
Rawimport { 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,
);
});
},
);