This commit is contained in:
@@ -9,7 +9,7 @@ import {
|
||||
ListToolsRequestSchema,
|
||||
} from "@modelcontextprotocol/sdk/types.js";
|
||||
|
||||
import { buildGateway, SEP } from "../src/gateway.js";
|
||||
import { buildGateway, SEP, type GatewayOptions } from "../src/gateway.js";
|
||||
import type { EnabledAgent } from "../src/config.js";
|
||||
|
||||
/**
|
||||
@@ -70,25 +70,65 @@ function startFakeAgent(opts: {
|
||||
});
|
||||
}
|
||||
|
||||
function startSseAgent(opts: {
|
||||
onCall: (body: any, res: http.ServerResponse) => void;
|
||||
}): Promise<{ url: string; close: () => Promise<void>; requests: any[] }> {
|
||||
const requests: any[] = [];
|
||||
return new Promise((resolve) => {
|
||||
const server = http.createServer((req, res) => {
|
||||
if (req.method !== "POST" || req.url !== "/mcp") {
|
||||
res.statusCode = 404;
|
||||
res.end();
|
||||
return;
|
||||
}
|
||||
const chunks: Buffer[] = [];
|
||||
req.on("data", (c) => chunks.push(c));
|
||||
req.on("end", () => {
|
||||
const body = JSON.parse(Buffer.concat(chunks).toString("utf8"));
|
||||
requests.push({ body });
|
||||
if (body.method !== "tools/call") {
|
||||
res.setHeader("content-type", "application/json");
|
||||
res.end(JSON.stringify({ jsonrpc: "2.0", id: body.id, result: {} }));
|
||||
return;
|
||||
}
|
||||
res.writeHead(200, { "content-type": "text/event-stream" });
|
||||
opts.onCall(body, res);
|
||||
});
|
||||
});
|
||||
server.listen(0, "127.0.0.1", () => {
|
||||
const { port } = server.address() as AddressInfo;
|
||||
resolve({
|
||||
url: `http://127.0.0.1:${port}`,
|
||||
requests,
|
||||
close: () => new Promise((r) => server.close(() => r())),
|
||||
});
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
/** Drive an in-process gateway Server via the SDK's request handlers. */
|
||||
async function callGateway(
|
||||
agents: EnabledAgent[],
|
||||
token: string | null,
|
||||
method: "tools/list" | "tools/call",
|
||||
params: any,
|
||||
opts: {
|
||||
gateway?: Partial<Omit<GatewayOptions, "agents" | "token" | "server">>;
|
||||
extra?: any;
|
||||
} = {},
|
||||
) {
|
||||
const mcpServer = new Server(
|
||||
{ name: "test", version: "0.0.0" },
|
||||
{ capabilities: { tools: { listChanged: false } } },
|
||||
);
|
||||
await buildGateway({ agents, token, server: mcpServer });
|
||||
await buildGateway({ agents, token, server: mcpServer, ...opts.gateway });
|
||||
// Reach into the server's registered handlers — the SDK exposes them via
|
||||
// a private map, so we route through the public schema-keyed lookup we
|
||||
// know is wired up.
|
||||
const schema = method === "tools/list" ? ListToolsRequestSchema : CallToolRequestSchema;
|
||||
const handler = (mcpServer as any)._requestHandlers.get(schema.shape.method.value);
|
||||
if (!handler) throw new Error("handler not registered");
|
||||
return handler({ method, params }, {} as any);
|
||||
return handler({ method, params }, opts.extra ?? ({} as any));
|
||||
}
|
||||
|
||||
test("tools/list aggregates upstream tools with agent prefix", async () => {
|
||||
@@ -209,3 +249,74 @@ test("tools/call surfaces upstream JSON-RPC error as soft error", async () => {
|
||||
await upstream.close();
|
||||
}
|
||||
});
|
||||
|
||||
test("tools/call forwards upstream SSE progress notifications", async () => {
|
||||
const upstream = await startSseAgent({
|
||||
onCall: (body, res) => {
|
||||
res.write(
|
||||
`data: ${JSON.stringify({
|
||||
jsonrpc: "2.0",
|
||||
method: "notifications/progress",
|
||||
params: { progressToken: "p1", message: "halfway" },
|
||||
})}\n\n`,
|
||||
);
|
||||
res.end(
|
||||
`data: ${JSON.stringify({
|
||||
jsonrpc: "2.0",
|
||||
id: body.id,
|
||||
result: {
|
||||
content: [{ type: "text", text: "done" }],
|
||||
isError: false,
|
||||
},
|
||||
})}\n\n`,
|
||||
);
|
||||
},
|
||||
});
|
||||
const notifications: any[] = [];
|
||||
try {
|
||||
const result = await callGateway(
|
||||
[{ name: "streamer", url: upstream.url, addedAt: "now" }],
|
||||
null,
|
||||
"tools/call",
|
||||
{ name: `streamer${SEP}build`, arguments: {} },
|
||||
{
|
||||
extra: {
|
||||
requestId: "req-1",
|
||||
sendNotification: async (notification: any) => {
|
||||
notifications.push(notification);
|
||||
},
|
||||
},
|
||||
},
|
||||
);
|
||||
assert.equal(result.isError, false);
|
||||
assert.deepEqual(notifications, [
|
||||
{
|
||||
method: "notifications/progress",
|
||||
params: { progressToken: "p1", message: "halfway" },
|
||||
},
|
||||
]);
|
||||
} finally {
|
||||
await upstream.close();
|
||||
}
|
||||
});
|
||||
|
||||
test("tools/call reports upstream SSE idle timeout", async () => {
|
||||
const upstream = await startSseAgent({
|
||||
onCall: (_body, res) => {
|
||||
res.write(": connected\n\n");
|
||||
},
|
||||
});
|
||||
try {
|
||||
const result = await callGateway(
|
||||
[{ name: "slow", url: upstream.url, addedAt: "now" }],
|
||||
null,
|
||||
"tools/call",
|
||||
{ name: `slow${SEP}build`, arguments: {} },
|
||||
{ gateway: { upstreamStreamIdleTimeoutMs: 25 } },
|
||||
);
|
||||
assert.equal(result.isError, true);
|
||||
assert.match(result.content[0].text, /stream idle timeout/);
|
||||
} finally {
|
||||
await upstream.close();
|
||||
}
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user