Handle upstream stream drops
All checks were successful
build / test (push) Successful in 29s

This commit is contained in:
2026-05-23 12:20:22 -04:00
parent 45e12324df
commit 2fd97d8d09
3 changed files with 224 additions and 29 deletions

View File

@@ -28,6 +28,10 @@ export interface GatewayOptions {
agents?: EnabledAgent[]; agents?: EnabledAgent[];
/** Override the bearer token (otherwise read from ~/.a2a/credentials.json). */ /** Override the bearer token (otherwise read from ~/.a2a/credentials.json). */
token?: string | null; token?: string | null;
/** Test/advanced override for one-shot upstream JSON-RPC calls. */
upstreamRequestTimeoutMs?: number;
/** Test/advanced override for idle SSE upstream tools/call streams. */
upstreamStreamIdleTimeoutMs?: number;
/** Optional override for the underlying MCP Server (tests). */ /** Optional override for the underlying MCP Server (tests). */
server?: Server; server?: Server;
} }
@@ -56,6 +60,8 @@ export async function buildGateway(opts: GatewayOptions = {}): Promise<GatewayHa
new UpstreamAgent({ new UpstreamAgent({
name: a.name, url: a.url, token, name: a.name, url: a.url, token,
cpJwt, cpUrl, bucket, cpJwt, cpUrl, bucket,
requestTimeoutMs: opts.upstreamRequestTimeoutMs,
streamIdleTimeoutMs: opts.upstreamStreamIdleTimeoutMs,
}), }),
); );
} }

View File

@@ -7,11 +7,15 @@
* header without coupling to a transport class. * header without coupling to a transport class.
*/ */
const JSON_RPC = "2.0"; const JSON_RPC = "2.0";
const DEFAULT_RPC_TIMEOUT_MS = 30_000;
const DEFAULT_STREAM_IDLE_TIMEOUT_MS = 120_000;
export interface UpstreamConfig { export interface UpstreamConfig {
name: string; name: string;
url: string; url: string;
token: string | null; token: string | null;
requestTimeoutMs?: number;
streamIdleTimeoutMs?: number;
/** Forwarded to the upstream as ``params._meta.cp_jwt`` on tools/call. /** Forwarded to the upstream as ``params._meta.cp_jwt`` on tools/call.
* Lets the upstream agent act on the caller's behalf (file/agent CRUD * Lets the upstream agent act on the caller's behalf (file/agent CRUD
* on /v1/me/*) — same path the platform orchestrator uses. */ * on /v1/me/*) — same path the platform orchestrator uses. */
@@ -63,32 +67,53 @@ export class UpstreamAgent {
private async rpc<T>(method: string, params: unknown): Promise<T> { private async rpc<T>(method: string, params: unknown): Promise<T> {
const id = _nextId++; const id = _nextId++;
const timeoutMs = this.cfg.requestTimeoutMs ?? DEFAULT_RPC_TIMEOUT_MS;
const controller = new AbortController();
let timedOut = false;
const timeout = setTimeout(() => {
timedOut = true;
controller.abort();
}, timeoutMs);
const headers: Record<string, string> = { const headers: Record<string, string> = {
"content-type": "application/json", "content-type": "application/json",
accept: "application/json", accept: "application/json",
}; };
if (this.cfg.token) headers["authorization"] = `Bearer ${this.cfg.token}`; if (this.cfg.token) headers["authorization"] = `Bearer ${this.cfg.token}`;
const resp = await fetch(`${this.cfg.url}/mcp`, { try {
method: "POST", const resp = await fetch(`${this.cfg.url}/mcp`, {
headers, method: "POST",
body: JSON.stringify({ jsonrpc: JSON_RPC, id, method, params }), headers,
}); body: JSON.stringify({ jsonrpc: JSON_RPC, id, method, params }),
signal: controller.signal,
});
if (!resp.ok) { if (!resp.ok) {
const text = await resp.text().catch(() => ""); const text = await resp.text().catch(() => "");
throw new UpstreamError(this.cfg.name, resp.status, text || resp.statusText); throw new UpstreamError(this.cfg.name, resp.status, text || resp.statusText);
} }
const env: any = await resp.json(); const env: any = await resp.json();
if (env?.error) { if (env?.error) {
throw new UpstreamError( throw new UpstreamError(
this.cfg.name, this.cfg.name,
500, 500,
`${env.error.code}: ${env.error.message}`, `${env.error.code}: ${env.error.message}`,
); );
}
return env?.result as T;
} catch (err) {
if (timedOut) {
throw new UpstreamError(
this.cfg.name,
504,
`request timed out after ${timeoutMs}ms`,
);
}
throw err;
} finally {
clearTimeout(timeout);
} }
return env?.result as T;
} }
async listTools(): Promise<UpstreamTool[]> { async listTools(): Promise<UpstreamTool[]> {
@@ -136,22 +161,51 @@ export class UpstreamAgent {
}) => void | Promise<void>, }) => void | Promise<void>,
): Promise<UpstreamCallResult> { ): Promise<UpstreamCallResult> {
const id = _nextId++; const id = _nextId++;
const idleTimeoutMs =
this.cfg.streamIdleTimeoutMs ?? DEFAULT_STREAM_IDLE_TIMEOUT_MS;
const controller = new AbortController();
let timedOut = false;
let idleTimer: ReturnType<typeof setTimeout> | undefined;
const resetIdleTimer = () => {
if (idleTimer) clearTimeout(idleTimer);
idleTimer = setTimeout(() => {
timedOut = true;
controller.abort();
}, idleTimeoutMs);
};
resetIdleTimer();
const headers: Record<string, string> = { const headers: Record<string, string> = {
"content-type": "application/json", "content-type": "application/json",
accept: "text/event-stream, application/json", accept: "text/event-stream, application/json",
}; };
if (this.cfg.token) headers["authorization"] = `Bearer ${this.cfg.token}`; if (this.cfg.token) headers["authorization"] = `Bearer ${this.cfg.token}`;
const resp = await fetch(`${this.cfg.url}/mcp`, { let resp: Response;
method: "POST", try {
headers, resp = await fetch(`${this.cfg.url}/mcp`, {
body: JSON.stringify({ method: "POST",
jsonrpc: JSON_RPC, id, method: "tools/call", params, headers,
}), body: JSON.stringify({
}); jsonrpc: JSON_RPC, id, method: "tools/call", params,
}),
signal: controller.signal,
});
} catch (err) {
clearTimeout(idleTimer);
if (timedOut) {
throw new UpstreamError(
this.cfg.name,
504,
`stream idle timeout after ${idleTimeoutMs}ms`,
);
}
throw err;
}
resetIdleTimer();
if (!resp.ok) { if (!resp.ok) {
const text = await resp.text().catch(() => ""); const text = await resp.text().catch(() => "");
clearTimeout(idleTimer);
throw new UpstreamError(this.cfg.name, resp.status, text || resp.statusText); throw new UpstreamError(this.cfg.name, resp.status, text || resp.statusText);
} }
@@ -160,15 +214,18 @@ export class UpstreamAgent {
// Server fell back to JSON (older a2a-pack or non-SSE upstream). // Server fell back to JSON (older a2a-pack or non-SSE upstream).
const env: any = await resp.json(); const env: any = await resp.json();
if (env?.error) { if (env?.error) {
clearTimeout(idleTimer);
throw new UpstreamError( throw new UpstreamError(
this.cfg.name, 500, this.cfg.name, 500,
`${env.error.code}: ${env.error.message}`, `${env.error.code}: ${env.error.message}`,
); );
} }
clearTimeout(idleTimer);
return env?.result as UpstreamCallResult; return env?.result as UpstreamCallResult;
} }
if (!resp.body) { if (!resp.body) {
clearTimeout(idleTimer);
throw new UpstreamError(this.cfg.name, 502, "empty stream body"); throw new UpstreamError(this.cfg.name, 502, "empty stream body");
} }
@@ -177,8 +234,24 @@ export class UpstreamAgent {
let buf = ""; let buf = "";
try { try {
while (true) { while (true) {
const { done, value } = await reader.read(); let done: boolean;
if (value) buf += decoder.decode(value, { stream: true }); let value: Uint8Array | undefined;
try {
({ done, value } = await reader.read());
} catch (err) {
if (timedOut) {
throw new UpstreamError(
this.cfg.name,
504,
`stream idle timeout after ${idleTimeoutMs}ms`,
);
}
throw err;
}
if (value) {
resetIdleTimer();
buf += decoder.decode(value, { stream: true });
}
let sep = buf.indexOf("\n\n"); let sep = buf.indexOf("\n\n");
while (sep !== -1) { while (sep !== -1) {
const frame = buf.slice(0, sep); const frame = buf.slice(0, sep);
@@ -189,6 +262,10 @@ export class UpstreamAgent {
} }
if (dataLines.length > 0) { if (dataLines.length > 0) {
const raw = dataLines.join("\n"); const raw = dataLines.join("\n");
if (raw === "[DONE]") {
sep = buf.indexOf("\n\n");
continue;
}
let msg: any; let msg: any;
try { msg = JSON.parse(raw); } catch { msg = null; } try { msg = JSON.parse(raw); } catch { msg = null; }
if (msg && msg.id === id) { if (msg && msg.id === id) {
@@ -221,6 +298,7 @@ export class UpstreamAgent {
if (done) break; if (done) break;
} }
} finally { } finally {
clearTimeout(idleTimer);
try { reader.releaseLock(); } catch { /* ignore */ } try { reader.releaseLock(); } catch { /* ignore */ }
} }
throw new UpstreamError(this.cfg.name, 502, "stream ended without result"); throw new UpstreamError(this.cfg.name, 502, "stream ended without result");

View File

@@ -9,7 +9,7 @@ import {
ListToolsRequestSchema, ListToolsRequestSchema,
} from "@modelcontextprotocol/sdk/types.js"; } 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"; 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. */ /** Drive an in-process gateway Server via the SDK's request handlers. */
async function callGateway( async function callGateway(
agents: EnabledAgent[], agents: EnabledAgent[],
token: string | null, token: string | null,
method: "tools/list" | "tools/call", method: "tools/list" | "tools/call",
params: any, params: any,
opts: {
gateway?: Partial<Omit<GatewayOptions, "agents" | "token" | "server">>;
extra?: any;
} = {},
) { ) {
const mcpServer = new Server( const mcpServer = new Server(
{ name: "test", version: "0.0.0" }, { name: "test", version: "0.0.0" },
{ capabilities: { tools: { listChanged: false } } }, { 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 // 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 // a private map, so we route through the public schema-keyed lookup we
// know is wired up. // know is wired up.
const schema = method === "tools/list" ? ListToolsRequestSchema : CallToolRequestSchema; const schema = method === "tools/list" ? ListToolsRequestSchema : CallToolRequestSchema;
const handler = (mcpServer as any)._requestHandlers.get(schema.shape.method.value); const handler = (mcpServer as any)._requestHandlers.get(schema.shape.method.value);
if (!handler) throw new Error("handler not registered"); 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 () => { 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(); 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();
}
});