From e7f541f156377612dd65ea1ba12056dcbd950593 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Wed, 29 Jul 2026 00:17:35 +0200 Subject: [PATCH] fix(v2): enforce connection lifecycle --- src/examples/dual-version-agent.ts | 10 +- src/jsonrpc.ts | 13 +- src/protocol-router.test.ts | 56 ++ src/protocol-router.ts | 14 +- src/v2/acp.test.ts | 916 +++++++++++++++++++++++++++-- src/v2/acp.ts | 868 +++++++++++++++++++++++++-- 6 files changed, 1767 insertions(+), 110 deletions(-) diff --git a/src/examples/dual-version-agent.ts b/src/examples/dual-version-agent.ts index d67c865..e982811 100644 --- a/src/examples/dual-version-agent.ts +++ b/src/examples/dual-version-agent.ts @@ -105,7 +105,7 @@ const v2Agent = v2 await cancelV2Turn(session); session.active = false; }) - .onRequest(v2.methods.agent.session.prompt, ({ params, client }) => { + .onRequest(v2.methods.agent.session.prompt, ({ params, client, accept }) => { const session = requireActiveV2Session(params.sessionId); if (session.turn) { throw v2.RequestError.invalidRequest( @@ -118,13 +118,7 @@ const v2Agent = v2 controller, done: Promise.resolve(), }; - // The framework queues the prompt response after this handler returns. - // Start work in the next event-loop task so that response is queued before - // any session updates from the turn. - const responseQueued = new Promise((resolve) => { - setTimeout(resolve, 0); - }); - turn.done = responseQueued + turn.done = accept() .then(() => runV2Turn(params, client, session, controller.signal)) .catch((error) => { console.error("v2 example turn failed", error); diff --git a/src/jsonrpc.ts b/src/jsonrpc.ts index 2c76d07..f1f8008 100644 --- a/src/jsonrpc.ts +++ b/src/jsonrpc.ts @@ -684,6 +684,13 @@ export class RequestResponder { */ public readonly signal: AbortSignal = new AbortController().signal, private finishRequest?: () => void, + /** + * Number of entries in the containing wire batch, when this request was + * received in a batch. + * + * @internal + */ + public readonly batchSize?: number, ) {} /** @@ -1325,6 +1332,7 @@ export class Connection { const processing = this.receiveMessage( message, isRequestMessage(message) ? collectResponse : undefined, + batch.length, ); if (isNotificationMessage(message)) { void processing.finally(() => { @@ -1338,6 +1346,7 @@ export class Connection { private receiveMessage( message: AnyMessage, sendResponse?: (response: AnyResponse) => Promise, + batchSize?: number, ): Promise { if (this.abortController.signal.aborted) { return Promise.resolve(); @@ -1355,7 +1364,7 @@ export class Connection { this.handleProtocolNotification(message); } return this.processIncomingMessage( - this.toIncomingMessage(message, sendResponse), + this.toIncomingMessage(message, sendResponse, batchSize), ).catch((error) => this.close(error)); } else if ("id" in message) { this.handleResponse(message); @@ -1428,6 +1437,7 @@ export class Connection { private toIncomingMessage( message: AnyRequest | AnyNotification, sendResponse?: (response: AnyResponse) => Promise, + batchSize?: number, ): IncomingMessage { if ("id" in message) { const abortController = new AbortController(); @@ -1458,6 +1468,7 @@ export class Connection { }, abortController.signal, finishRequest, + batchSize, ), }; } diff --git a/src/protocol-router.test.ts b/src/protocol-router.test.ts index cb17933..dab1b12 100644 --- a/src/protocol-router.test.ts +++ b/src/protocol-router.test.ts @@ -62,6 +62,62 @@ describe("AgentProtocolRouter", () => { await second.close(); }); + it.each([ + [ + 1, + { + clientCapabilities: { + _future: { nullable: null, entries: [{ value: 1, extra: true }] }, + }, + clientInfo: implementation("v1-client"), + futureTopLevel: { empty: [], explicitNull: null }, + }, + ], + [ + 2, + { + info: implementation("v2-client"), + capabilities: { + _future: { nullable: null, entries: [{ value: 2, extra: true }] }, + }, + futureTopLevel: { empty: [], explicitNull: null }, + }, + ], + ])( + "validates and preserves same-version v%s initialize params", + async (version, params) => { + const selected = new MockAgentConnector(); + const router = + version === 1 + ? new AgentProtocolRouter().withV1(selected) + : new AgentProtocolRouter().withV2(selected); + const initialize = initializeRequest(version, params); + const connection = await openRoutedConnection(router, initialize); + + expect(await selected.nextMessage()).toEqual(initialize); + await connection.close(); + }, + ); + + it("still validates same-version initialize params before routing", async () => { + const v2 = new MockAgentConnector(); + const router = new AgentProtocolRouter().withV2(v2); + const { response, closed } = await rejectedConnection( + router, + initializeRequest(2, { capabilities: {} }), + ); + + expect(response).toMatchObject({ + id: 1, + error: { + code: -32602, + data: expect.stringContaining("invalid initialize params"), + }, + }); + expect(v2.connectionCount).toBe(0); + await closed; + }); + it("routes a future protocol version to v2 and normalizes initialize", async () => { const v1 = new MockAgentConnector(); const v2 = new MockAgentConnector(); diff --git a/src/protocol-router.ts b/src/protocol-router.ts index 4c2854a..6871d32 100644 --- a/src/protocol-router.ts +++ b/src/protocol-router.ts @@ -43,8 +43,9 @@ const MAX_PROTOCOL_VERSION = 0xffff; * The router consumes the first wire item, which must be an `initialize` * request or a single-entry batch containing one. It selects the highest * configured protocol version that does not exceed the client's requested - * version. Only the initialize params and optional batch framing are - * normalized; every later wire item is forwarded unchanged. + * version. Same-version initialize params are validated and forwarded + * unchanged. Version selection and downgrade may normalize initialize params; + * every later wire item is forwarded unchanged. * * @experimental */ @@ -560,6 +561,15 @@ function rewriteInitializeParams( requested: number, selected: AgentProtocol, ): JsonObject { + if (requested === selected) { + if (selected === PROTOCOL_V1) { + zV1InitializeRequest.parse(params); + } else { + zV2InitializeRequest.parse(params); + } + return params; + } + if (selected === PROTOCOL_V1) { return requested >= PROTOCOL_V2 ? v2InitializeToV1(zV2InitializeRequest.parse(params)) diff --git a/src/v2/acp.test.ts b/src/v2/acp.test.ts index 7da387d..2d97102 100644 --- a/src/v2/acp.test.ts +++ b/src/v2/acp.test.ts @@ -22,6 +22,7 @@ import type { ClientContext, DiffPatch, ExtensionMethod, + InitializeRequest, InitializeResponse, McpServer, NewSessionRequest, @@ -34,6 +35,13 @@ import type { const clientInfo = { name: "test-client", version: "1.0.0" }; const agentInfo = { name: "test-agent", version: "1.0.0" }; +function testAgent(): sdk.AgentApp { + return agent().onRequest(methods.agent.initialize, () => ({ + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + })); +} + function assertV2MethodTypes( agentContext: ClientContext, clientContext: AgentContext, @@ -250,7 +258,7 @@ describe("experimental v2 app API", () => { } satisfies NewSessionRequest; const expected = structuredClone(request); - await client().connectWith(agent(), (agentClient) => { + await client().connectWith(testAgent(), (agentClient) => { const builder = agentClient.buildSession(request); additionalDirectories[0] = "/mutated-input"; @@ -315,7 +323,7 @@ describe("experimental v2 app API", () => { } satisfies McpServer; const expected = structuredClone(mcpServer); - await client().connectWith(agent(), (agentClient) => { + await client().connectWith(testAgent(), (agentClient) => { const builder = agentClient .buildSession("/workspace") .withMcpServer(mcpServer); @@ -347,6 +355,711 @@ describe("experimental v2 app API", () => { } }); + it("initializes exactly once, queues later calls, and exposes snapshots", async () => { + const initializeGate = Promise.withResolvers(); + const events: string[] = []; + let newSessionCalls = 0; + let receivedInitialize: InitializeRequest | undefined; + + const initializeRequest = { + protocolVersion: PROTOCOL_VERSION, + info: { + ...clientInfo, + futureInfo: { value: "client" }, + }, + capabilities: { + futureCapability: { enabled: true }, + }, + futureRequest: { entries: [{ value: 1, extra: "preserved" }] }, + } as InitializeRequest & { + futureRequest: { entries: Array<{ value: number; extra: string }> }; + }; + const initializeResponse = { + protocolVersion: PROTOCOL_VERSION, + info: { + ...agentInfo, + futureInfo: { value: "agent" }, + }, + capabilities: { + session: {}, + futureCapability: { enabled: true }, + }, + authMethods: [ + { + type: "agent", + methodId: "agent-auth", + name: "Agent auth", + futureAuthField: { preserved: true }, + }, + ], + futureResponse: { entries: [{ value: 2, extra: "preserved" }] }, + } as InitializeResponse & { + futureResponse: { entries: Array<{ value: number; extra: string }> }; + }; + + const agentApp = agent() + .onConnect((connection) => { + events.push("agent-connect"); + expect(connection.initialization).toBeUndefined(); + expect(connection.clientCapabilities).toBeUndefined(); + }) + .onInitialized((connection, initialization) => { + events.push("agent-initialized"); + expect(connection.initialization).toEqual(initialization); + expect(connection.clientCapabilities).toMatchObject( + initializeRequest.capabilities!, + ); + }) + .onRequest(methods.agent.initialize, async ({ params, client }) => { + receivedInitialize = params; + expect(client.initialization).toBeUndefined(); + expect(client.clientCapabilities).toBeUndefined(); + await initializeGate.promise; + return initializeResponse; + }) + .onRequest(methods.agent.session.new, () => { + newSessionCalls += 1; + return { sessionId: "session-1" }; + }); + const clientApp = client() + .onConnect((connection) => { + events.push("client-connect"); + expect(connection.initialization).toBeUndefined(); + expect(connection.agentCapabilities).toBeUndefined(); + }) + .onInitialized((connection, initialization) => { + events.push("client-initialized"); + expect(connection.initialization).toEqual(initialization); + expect(connection.agentCapabilities).toMatchObject( + initializeResponse.capabilities!, + ); + }); + + await clientApp.connectWith(agentApp, async (agentContext) => { + await expect( + agentContext.request(methods.agent.session.new, { + cwd: "/workspace", + mcpServers: [], + }), + ).rejects.toMatchObject({ code: -32600 }); + await expect( + agentContext.notify(methods.agent.session.cancel, { + sessionId: "session-1", + }), + ).rejects.toMatchObject({ code: -32600 }); + await expect( + agentContext.request("_vendor/pre-initialize", {}), + ).rejects.toMatchObject({ code: -32600 }); + + const initialized = agentContext.request( + methods.agent.initialize, + initializeRequest, + ); + const queuedSession = agentContext.request(methods.agent.session.new, { + cwd: "/workspace", + mcpServers: [], + }); + await Promise.resolve(); + expect(newSessionCalls).toBe(0); + expect(() => + agentContext.request(methods.agent.initialize, initializeRequest), + ).toThrow("Invalid request"); + + initializeGate.resolve(); + await expect(initialized).resolves.toMatchObject(initializeResponse); + await expect(queuedSession).resolves.toEqual({ sessionId: "session-1" }); + + expect(agentContext.initialization).toMatchObject({ + request: initializeRequest, + response: initializeResponse, + }); + expect(agentContext.agentCapabilities).toMatchObject( + initializeResponse.capabilities!, + ); + expect(receivedInitialize).toMatchObject(initializeRequest); + expect(events.slice(0, 2)).toEqual(["agent-connect", "client-connect"]); + expect(events).toContain("agent-initialized"); + expect(events).toContain("client-initialized"); + expect(() => + agentContext.request(methods.agent.initialize, initializeRequest), + ).toThrow("Invalid request"); + }); + }); + + it("allows only a standalone initialize batch and does not retry failures", async () => { + let initializedHooks = 0; + const agentInitialized = Promise.withResolvers(); + const validAgent = agent() + .onInitialized(() => { + initializedHooks += 1; + agentInitialized.resolve(); + }) + .onRequest(methods.agent.initialize, () => ({ + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + })); + + await client().connectWith(validAgent, async (agentContext) => { + await expect( + agentContext.batch([ + batchRequest(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }), + batchNotification(methods.agent.session.cancel, { + sessionId: "session-1", + }), + ] as const), + ).rejects.toMatchObject({ code: -32600 }); + await expect( + agentContext.batch([ + batchRequest(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }), + batchRequest(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }), + ] as const), + ).rejects.toMatchObject({ code: -32600 }); + + await expect( + agentContext.batch([ + batchRequest(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }), + ] as const), + ).resolves.toEqual([ + expect.objectContaining({ protocolVersion: PROTOCOL_VERSION }), + ]); + await agentInitialized.promise; + expect(initializedHooks).toBe(1); + }); + + let failedHooks = 0; + const invalidAgent = agent() + .onInitialized(() => { + failedHooks += 1; + }) + .onRequest(methods.agent.initialize, () => ({ + protocolVersion: 1, + info: agentInfo, + })); + await client().connectWith(invalidAgent, async (agentContext) => { + const initialized = agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + const queued = agentContext.request(methods.agent.session.new, { + cwd: "/workspace", + mcpServers: [], + }); + const queuedRejected = expect(queued).rejects.toMatchObject({ + code: -32600, + }); + await expect(initialized).rejects.toMatchObject({ code: -32600 }); + await queuedRejected; + expect(agentContext.initialization).toBeUndefined(); + expect(() => + agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }), + ).toThrow("Invalid request"); + }); + expect(failedHooks).toBe(0); + }); + + it("rejects initialize after a pre-initialize request fails the connection", async () => { + let initializeCalls = 0; + let newSessionCalls = 0; + const agentApp = agent() + .onRequest(methods.agent.initialize, () => { + initializeCalls += 1; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }; + }) + .onRequest(methods.agent.session.new, () => { + newSessionCalls += 1; + return { sessionId: "session-1" }; + }); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + const reader = peerStream.readable.getReader(); + + try { + const preInitializeWrite = writer.write({ + jsonrpc: "2.0", + id: 1, + method: methods.agent.session.new, + params: { cwd: "/workspace", mcpServers: [] }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { + id: 1, + error: { code: -32600 }, + }, + }); + await preInitializeWrite; + + const initializeWrite = writer.write({ + jsonrpc: "2.0", + id: 2, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { + id: 2, + error: { code: -32600 }, + }, + }); + await initializeWrite; + + expect(initializeCalls).toBe(0); + expect(newSessionCalls).toBe(0); + expect(connection.initialization).toBeUndefined(); + } finally { + writer.releaseLock(); + reader.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("rejects mixed raw initialize batches before dispatching handlers", async () => { + let initializeCalls = 0; + let newSessionCalls = 0; + const agentApp = agent() + .onRequest(methods.agent.initialize, () => { + initializeCalls += 1; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }; + }) + .onRequest(methods.agent.session.new, () => { + newSessionCalls += 1; + return { sessionId: "session-1" }; + }); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + const reader = peerStream.readable.getReader(); + + try { + const batchWrite = writer.write([ + { + jsonrpc: "2.0", + id: 1, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }, + { + jsonrpc: "2.0", + id: 2, + method: methods.agent.session.new, + params: { cwd: "/workspace", mcpServers: [] }, + }, + ]); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: expect.arrayContaining([ + expect.objectContaining({ + id: 1, + error: expect.objectContaining({ code: -32600 }), + }), + expect.objectContaining({ + id: 2, + error: expect.objectContaining({ code: -32600 }), + }), + ]), + }); + await batchWrite; + + expect(initializeCalls).toBe(0); + expect(newSessionCalls).toBe(0); + expect(connection.initialization).toBeUndefined(); + } finally { + writer.releaseLock(); + reader.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("requires an initialize handler before opening a connection", async () => { + const agentApp = agent(); + const [firstStream] = memoryWireStreamPair(); + expect(() => agentApp.connect(firstStream)).toThrow( + "requires an initialize request handler", + ); + + agentApp.onRequest(methods.agent.initialize, () => ({ + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + })); + const [secondStream] = memoryWireStreamPair(); + const connection = agentApp.connect(secondStream); + connection.close(); + await connection.closed; + }); + + it("treats a cancellation notification before initialize as the first message", async () => { + let initializeCalls = 0; + const agentApp = agent().onRequest(methods.agent.initialize, () => { + initializeCalls += 1; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }; + }); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + const reader = peerStream.readable.getReader(); + + try { + await writer.write({ + jsonrpc: "2.0", + method: methods.protocol.cancelRequest, + params: { requestId: 999 }, + }); + const initializeWrite = writer.write({ + jsonrpc: "2.0", + id: 1, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { id: 1, error: { code: -32600 } }, + }); + await initializeWrite; + expect(initializeCalls).toBe(0); + } finally { + writer.releaseLock(); + reader.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("queues raw peer requests behind an in-flight initialize", async () => { + const initializeGate = Promise.withResolvers(); + const initializeStarted = Promise.withResolvers(); + let newSessionCalls = 0; + const agentApp = agent() + .onRequest(methods.agent.initialize, async () => { + initializeStarted.resolve(); + await initializeGate.promise; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }; + }) + .onRequest(methods.agent.session.new, () => { + newSessionCalls += 1; + return { sessionId: "session-1" }; + }); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + const reader = peerStream.readable.getReader(); + + try { + const initializeWrite = writer.write({ + jsonrpc: "2.0", + id: 1, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }); + await initializeStarted.promise; + + const malformedDuplicateWrite = writer.write({ + jsonrpc: "2.0", + id: 3, + method: methods.agent.initialize, + params: { protocolVersion: "invalid" }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { + id: 3, + error: { code: -32600 }, + }, + }); + await malformedDuplicateWrite; + + await writer.write({ + jsonrpc: "2.0", + method: methods.protocol.cancelRequest, + params: { requestId: 999 }, + }); + const queuedWrite = writer.write({ + jsonrpc: "2.0", + id: 2, + method: methods.agent.session.new, + params: { cwd: "/workspace", mcpServers: [] }, + }); + await Promise.resolve(); + expect(newSessionCalls).toBe(0); + + const initializeResponse = reader.read(); + initializeGate.resolve(); + await expect(initializeResponse).resolves.toMatchObject({ + done: false, + value: { + id: 1, + result: { protocolVersion: PROTOCOL_VERSION }, + }, + }); + await initializeWrite; + + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { + id: 2, + result: { sessionId: "session-1" }, + }, + }); + await queuedWrite; + expect(newSessionCalls).toBe(1); + + const malformedCancelWrite = writer.write({ + jsonrpc: "2.0", + id: 4, + method: methods.protocol.cancelRequest, + params: { requestId: 999 }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { id: 4, error: { code: -32601 } }, + }); + await malformedCancelWrite; + } finally { + initializeGate.resolve(); + writer.releaseLock(); + reader.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("rejects queued peer calls when an in-flight initialize connection closes", async () => { + const initializeGate = Promise.withResolvers(); + const initializeStarted = Promise.withResolvers(); + const agentApp = agent().onRequest(methods.agent.initialize, async () => { + initializeStarted.resolve(); + await initializeGate.promise; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }; + }); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + + try { + const initializeWrite = writer.write({ + jsonrpc: "2.0", + id: 1, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }); + await initializeStarted.promise; + await initializeWrite; + + const queuedCall = connection.client.request("_vendor/queued", {}); + connection.close(); + await expect(queuedCall).rejects.toMatchObject({ code: -32600 }); + } finally { + initializeGate.resolve(); + writer.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("rejects session/prompt in a mixed raw batch before its handler runs", async () => { + let promptCalls = 0; + const agentApp = agent() + .onRequest(methods.agent.initialize, () => ({ + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + })) + .onRequest( + methods.agent.session.prompt, + async ({ accept, client: agentClient }) => { + promptCalls += 1; + await accept(); + await agentClient.notify(methods.client.session.update, { + sessionId: "session-1", + update: { + sessionUpdate: "state_update", + state: "idle", + }, + }); + }, + ) + .onRequest(methods.agent.session.new, () => ({ + sessionId: "session-1", + })); + const [agentStream, peerStream] = memoryWireStreamPair(); + const connection = agentApp.connect(agentStream); + const writer = peerStream.writable.getWriter(); + const reader = peerStream.readable.getReader(); + + try { + const initializeWrite = writer.write({ + jsonrpc: "2.0", + id: 1, + method: methods.agent.initialize, + params: { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }, + }); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: { + id: 1, + result: { protocolVersion: PROTOCOL_VERSION }, + }, + }); + await initializeWrite; + + const promptBatchWrite = writer.write([ + { + jsonrpc: "2.0", + id: 2, + method: methods.agent.session.prompt, + params: { + sessionId: "session-1", + prompt: [{ type: "text", text: "Hello" }], + }, + }, + { + jsonrpc: "2.0", + id: 3, + method: methods.agent.session.new, + params: { cwd: "/workspace", mcpServers: [] }, + }, + ]); + await expect(reader.read()).resolves.toMatchObject({ + done: false, + value: expect.arrayContaining([ + expect.objectContaining({ + id: 2, + error: expect.objectContaining({ code: -32600 }), + }), + expect.objectContaining({ + id: 3, + result: { sessionId: "session-1" }, + }), + ]), + }); + await promptBatchWrite; + expect(promptCalls).toBe(0); + } finally { + writer.releaseLock(); + reader.releaseLock(); + connection.close(); + await connection.closed; + } + }); + + it("accepts prompts before waiting for turn work and updates", async () => { + const turnGate = Promise.withResolvers(); + const turnFinished = Promise.withResolvers(); + let updateClient: AgentContext | undefined; + + const agentApp = agent() + .onRequest(methods.agent.initialize, ({ client: agentClient }) => { + updateClient = agentClient; + return { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + capabilities: { session: {} }, + }; + }) + .onRequest(methods.agent.session.new, () => ({ sessionId: "session-1" })) + .onRequest( + methods.agent.session.prompt, + async ({ params, client: agentClient, accept }) => { + await accept(); + await turnGate.promise; + await agentClient.notify(methods.client.session.update, { + sessionId: params.sessionId, + update: { + sessionUpdate: "state_update", + state: "idle", + stopReason: "end_turn", + }, + }); + turnFinished.resolve(); + }, + ); + + await client().connectWith(agentApp, async (agentContext) => { + await agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + await expect( + agentContext.batch([ + batchRequest(methods.agent.session.prompt, { + sessionId: "session-1", + prompt: [{ type: "text", text: "Hello" }], + }), + batchRequest(methods.agent.session.new, { + cwd: "/workspace", + mcpServers: [], + }), + ] as const), + ).rejects.toMatchObject({ code: -32600 }); + const session = await agentContext.buildSession("/workspace").start(); + try { + const accepted = session.prompt("Hello"); + await expect(accepted).resolves.toEqual({}); + expect(updateClient).toBeDefined(); + + turnGate.resolve(); + await turnFinished.promise; + await expect(session.nextUpdate()).resolves.toMatchObject({ + kind: "stop", + stopReason: "end_turn", + }); + } finally { + turnGate.resolve(); + session.dispose(); + } + }); + }); + it("does not complete a prompt from an idle update received before it", async () => { let updateClient: AgentContext | undefined; @@ -674,15 +1387,89 @@ describe("experimental v2 app API", () => { ).rejects.toMatchObject({ code: -32600 }); }); - it("validates every built-in direct response before returning it", async () => { + it("preserves extensions from retained initialize array entries only", async () => { const [clientStream, peerStream] = memoryWireStreamPair(); - const response = client().connectWith(clientStream, (agentContext) => - agentContext.request(methods.agent.session.new, { - cwd: "/workspace", - mcpServers: [], + const initialized = client().connectWith(clientStream, (agentContext) => + agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, }), ); + await respondToNextRequest(peerStream, { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + authMethods: [ + { + type: "agent", + methodId: 42, + name: "Malformed", + fromMalformedMethod: true, + }, + { + type: "terminal", + methodId: "terminal-auth", + name: "Terminal", + env: [ + { + name: "TOKEN", + value: 42, + fromMalformedEnv: true, + }, + { + name: "TOKEN", + value: "good", + fromValidEnv: true, + }, + ], + fromValidMethod: true, + }, + ], + }); + + await expect(initialized).resolves.toMatchObject({ + authMethods: [ + { + type: "terminal", + methodId: "terminal-auth", + fromValidMethod: true, + env: [ + { + name: "TOKEN", + value: "good", + fromValidEnv: true, + }, + ], + }, + ], + }); + const response = await initialized; + expect(response.authMethods?.[0]).not.toHaveProperty("fromMalformedMethod"); + expect( + (response.authMethods?.[0] as { env?: unknown[] } | undefined)?.env?.[0], + ).not.toHaveProperty("fromMalformedEnv"); + }); + + it("validates every built-in direct response before returning it", async () => { + const [clientStream, peerStream] = memoryWireStreamPair(); + const response = client().connectWith( + clientStream, + async (agentContext) => { + await agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + return agentContext.request(methods.agent.session.new, { + cwd: "/workspace", + mcpServers: [], + }); + }, + ); + + await respondToNextRequest(peerStream, { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }); await respondToNextRequest(peerStream, { sessionId: 42 }); await expect(response).rejects.toThrow(); }); @@ -690,19 +1477,30 @@ describe("experimental v2 app API", () => { it("validates built-in batch responses before applying caller mappings", async () => { const [clientStream, peerStream] = memoryWireStreamPair(); let mapped = false; - const response = client().connectWith(clientStream, (agentContext) => - agentContext.batch([ - batchRequest( - methods.agent.session.new, - { cwd: "/workspace", mcpServers: [] }, - (session) => { - mapped = true; - return session.sessionId; - }, - ), - ] as const), + const response = client().connectWith( + clientStream, + async (agentContext) => { + await agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + return agentContext.batch([ + batchRequest( + methods.agent.session.new, + { cwd: "/workspace", mcpServers: [] }, + (session) => { + mapped = true; + return session.sessionId; + }, + ), + ] as const); + }, ); + await respondToNextRequest(peerStream, { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }); const reader = peerStream.readable.getReader(); const request = await reader.read(); reader.releaseLock(); @@ -733,22 +1531,43 @@ describe("experimental v2 app API", () => { it("rejects peer null for empty responses but preserves local void handlers", async () => { const [clientStream, peerStream] = memoryWireStreamPair(); - const invalidResponse = client().connectWith(clientStream, (agentContext) => - agentContext.request(methods.agent.session.delete, { - sessionId: "session-1", - }), + const invalidResponse = client().connectWith( + clientStream, + async (agentContext) => { + await agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + return agentContext.request(methods.agent.session.delete, { + sessionId: "session-1", + }); + }, ); + await respondToNextRequest(peerStream, { + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + }); await respondToNextRequest(peerStream, null); await expect(invalidResponse).rejects.toThrow(); await expect( client().connectWith( - agent().onRequest(methods.agent.session.delete, () => {}), - (agentContext) => - agentContext.request(methods.agent.session.delete, { + agent() + .onRequest(methods.agent.initialize, () => ({ + protocolVersion: PROTOCOL_VERSION, + info: agentInfo, + })) + .onRequest(methods.agent.session.delete, () => {}), + async (agentContext) => { + await agentContext.request(methods.agent.initialize, { + protocolVersion: PROTOCOL_VERSION, + info: clientInfo, + }); + return agentContext.request(methods.agent.session.delete, { sessionId: "session-1", - }), + }); + }, ), ).resolves.toEqual({}); }); @@ -790,6 +1609,7 @@ describe("experimental v2 app API", () => { ).toThrow("must start with '_'"); let notificationValue: string | undefined; + const notificationSent = Promise.withResolvers(); const clientApp = client().onNotification( "_vendor/acme/event", parseValue, @@ -799,27 +1619,32 @@ describe("experimental v2 app API", () => { ); const agentApp = agent() .onRequest("_vendor/acme/echo", parseValue, returnValue) - .onRequest( - methods.agent.initialize, - async ({ client: clientContext }) => { - expect(() => - clientContext.request( - uncheckedExtension("vendor/client-request"), - {}, - ), - ).toThrow("must start with '_'"); - expect(() => - clientContext.notify( - uncheckedExtension("vendor/client-notification"), - {}, - ), - ).toThrow("must start with '_'"); - await clientContext.notify("_vendor/acme/event", { + .onInitialized(async (connection) => { + try { + await connection.client.notify("_vendor/acme/event", { value: "notification", }); - return { protocolVersion: PROTOCOL_VERSION, info: agentInfo }; - }, - ); + notificationSent.resolve(); + } catch (error) { + notificationSent.reject(error); + throw error; + } + }) + .onRequest(methods.agent.initialize, ({ client: clientContext }) => { + expect(() => + clientContext.request( + uncheckedExtension("vendor/client-request"), + {}, + ), + ).toThrow("must start with '_'"); + expect(() => + clientContext.notify( + uncheckedExtension("vendor/client-notification"), + {}, + ), + ).toThrow("must start with '_'"); + return { protocolVersion: PROTOCOL_VERSION, info: agentInfo }; + }); await clientApp.connectWith(agentApp, async (agentContext) => { expect(() => @@ -847,6 +1672,7 @@ describe("experimental v2 app API", () => { { value: "response" }, ), ).resolves.toEqual({ value: "response" }); + await notificationSent.promise; }); expect(notificationValue).toBe("notification"); }); diff --git a/src/v2/acp.ts b/src/v2/acp.ts index 7a924cc..72c7b7d 100644 --- a/src/v2/acp.ts +++ b/src/v2/acp.ts @@ -309,8 +309,88 @@ function assertV2BatchMethods( } } +function rawContainsParsedValue(raw: unknown, parsed: unknown): boolean { + if (Array.isArray(parsed)) { + if (!Array.isArray(raw)) { + return false; + } + + let rawIndex = 0; + for (const value of parsed) { + while ( + rawIndex < raw.length && + !rawContainsParsedValue(raw[rawIndex], value) + ) { + rawIndex += 1; + } + if (rawIndex === raw.length) { + return false; + } + rawIndex += 1; + } + return true; + } + + if (typeof parsed === "object" && parsed !== null) { + if (typeof raw !== "object" || raw === null || Array.isArray(raw)) { + return false; + } + return Object.entries(parsed).every(([key, value]) => + rawContainsParsedValue((raw as Record)[key], value), + ); + } + + return Object.is(raw, parsed); +} + +function mergeParsedInitializeValue(raw: unknown, parsed: unknown): unknown { + if (Array.isArray(raw) && Array.isArray(parsed)) { + if (raw.length === parsed.length) { + return parsed.map((value, index) => + mergeParsedInitializeValue(raw[index], value), + ); + } + + let rawIndex = 0; + return parsed.map((value) => { + while ( + rawIndex < raw.length && + !rawContainsParsedValue(raw[rawIndex], value) + ) { + rawIndex += 1; + } + const matchingRaw = rawIndex < raw.length ? raw[rawIndex++] : undefined; + return mergeParsedInitializeValue(matchingRaw, value); + }); + } + + if ( + typeof raw !== "object" || + raw === null || + Array.isArray(raw) || + typeof parsed !== "object" || + parsed === null || + Array.isArray(parsed) + ) { + return parsed; + } + + const result = { ...(raw as Record) }; + for (const [key, parsedValue] of Object.entries(parsed)) { + result[key] = mergeParsedInitializeValue( + (raw as Record)[key], + parsedValue, + ); + } + return result; +} + function parseV2InitializeRequest(params: unknown): schema.InitializeRequest { - const request = validate.zInitializeRequest.parse(params); + const parsed = validate.zInitializeRequest.parse(params); + const request = mergeParsedInitializeValue( + params, + parsed, + ) as schema.InitializeRequest; if (request.protocolVersion !== schema.PROTOCOL_VERSION) { throw RequestError.invalidParams( { @@ -337,7 +417,11 @@ function normalizeOutgoingV2InitializeRequest( } function mapV2InitializeResponse(response: unknown): schema.InitializeResponse { - const parsed = validate.zInitializeResponse.parse(response); + const validated = validate.zInitializeResponse.parse(response); + const parsed = mergeParsedInitializeValue( + response, + validated, + ) as schema.InitializeResponse; if (parsed.protocolVersion !== schema.PROTOCOL_VERSION) { throw RequestError.invalidRequest( { @@ -350,6 +434,271 @@ function mapV2InitializeResponse(response: unknown): schema.InitializeResponse { return parsed; } +/** + * Validated initialize request and response for one ACP v2 connection. + * + * The snapshot becomes available only after the initialize response has been + * validated and sent or received successfully. + * + * @experimental + */ +export type InitializationSnapshot = Readonly<{ + request: schema.InitializeRequest; + response: schema.InitializeResponse; +}>; + +type InitializationPhase = + "uninitialized" | "initializing" | "initialized" | "failed"; + +function cloneInitialization( + initialization: InitializationSnapshot, +): InitializationSnapshot { + return structuredClone(initialization); +} + +function deferred(): { + promise: Promise; + resolve(value: T | PromiseLike): void; + reject(reason?: unknown): void; +} { + let resolve!: (value: T | PromiseLike) => void; + let reject!: (reason?: unknown) => void; + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise; + reject = rejectPromise; + }); + return { promise, resolve, reject }; +} + +class InitializationState { + private phase: InitializationPhase = "uninitialized"; + private request?: schema.InitializeRequest; + private value?: InitializationSnapshot; + private listeners: Array<(value: InitializationSnapshot) => void> = []; + private readonly barrier = deferred(); + + constructor() { + void this.barrier.promise.catch(() => {}); + } + + get status(): InitializationPhase { + return this.phase; + } + + get snapshot(): InitializationSnapshot | undefined { + return this.value ? cloneInitialization(this.value) : undefined; + } + + begin(request: schema.InitializeRequest): void { + if (this.phase !== "uninitialized") { + throw RequestError.invalidRequest( + "ACP v2 initialize may only be requested once per connection", + ); + } + + this.phase = "initializing"; + this.request = structuredClone(request); + } + + complete(response: schema.InitializeResponse): void { + if (this.phase !== "initializing" || !this.request) { + throw RequestError.invalidRequest( + "ACP v2 initialization is not in progress", + ); + } + + this.value = { + request: structuredClone(this.request), + response: structuredClone(response), + }; + this.phase = "initialized"; + this.barrier.resolve(); + for (const listener of this.listeners.splice(0)) { + listener(cloneInitialization(this.value)); + } + } + + fail(): void { + if (this.phase === "uninitialized" || this.phase === "initializing") { + this.phase = "failed"; + this.request = undefined; + this.barrier.reject( + RequestError.invalidRequest("ACP v2 connection initialization failed"), + ); + } + } + + waitUntilInitialized(method: string): Promise { + if (this.phase === "initialized") { + return Promise.resolve(); + } + if (this.phase === "initializing") { + return this.barrier.promise; + } + return Promise.reject(initializationUnavailable(method)); + } + + onInitialized(listener: (value: InitializationSnapshot) => void): void { + if (this.value) { + listener(cloneInitialization(this.value)); + return; + } + this.listeners.push(listener); + } +} + +const initializationStates = new WeakMap< + ConnectionContext, + InitializationState +>(); + +function initializationState(cx: ConnectionContext): InitializationState { + let state = initializationStates.get(cx); + if (!state) { + state = new InitializationState(); + initializationStates.set(cx, state); + cx.signal.addEventListener("abort", () => state?.fail(), { once: true }); + } + return state; +} + +function initializationUnavailable(method: string): RequestError { + return RequestError.invalidRequest( + `ACP v2 connection must be initialized before '${method}'`, + ); +} + +async function blockUninitializedIncoming( + message: IncomingMessage, + state: InitializationState, +): Promise { + state.fail(); + return rejectIncoming(message, initializationUnavailable(message.method)); +} + +async function rejectIncoming( + message: IncomingMessage, + error: RequestError, +): Promise { + if (message.kind === "request") { + await message.responder.respondWithError(error); + } + return Handled.yes(); +} + +function agentInitializationGuard(): JsonRpcHandler { + return { + async handleMessage(message, cx) { + const state = initializationState(cx); + if ( + message.kind === "notification" && + message.method === schema.PROTOCOL_METHODS.cancel_request + ) { + return state.status === "initializing" || state.status === "initialized" + ? Handled.yes() + : blockUninitializedIncoming(message, state); + } + + if ( + message.kind === "request" && + message.method === schema.AGENT_METHODS.initialize + ) { + if (state.status !== "uninitialized") { + return rejectIncoming( + message, + RequestError.invalidRequest( + "ACP v2 initialize may only be requested once per connection", + ), + ); + } + if ( + message.responder.batchSize !== undefined && + message.responder.batchSize !== 1 + ) { + state.fail(); + return rejectIncoming( + message, + RequestError.invalidRequest( + "ACP v2 initialize must be the only entry in its batch", + ), + ); + } + + let request: schema.InitializeRequest; + try { + request = parseV2InitializeRequest(message.params); + } catch (error) { + state.fail(); + throw error; + } + state.begin(request); + return Handled.no(message); + } + + if (state.status === "initialized") { + if ( + message.kind === "request" && + message.method === schema.AGENT_METHODS.session_prompt && + message.responder.batchSize !== undefined && + message.responder.batchSize !== 1 + ) { + return rejectIncoming( + message, + RequestError.invalidRequest( + "ACP v2 session/prompt must be the only entry in its batch", + ), + ); + } + return Handled.no(message); + } + if (state.status === "initializing") { + if ( + message.kind === "request" && + message.method === schema.AGENT_METHODS.session_prompt && + message.responder.batchSize !== undefined && + message.responder.batchSize !== 1 + ) { + return rejectIncoming( + message, + RequestError.invalidRequest( + "ACP v2 session/prompt must be the only entry in its batch", + ), + ); + } + await state.waitUntilInitialized(message.method); + return Handled.no(message); + } + return blockUninitializedIncoming(message, state); + }, + describe: () => "agent-initialization", + }; +} + +function clientInitializationGuard(): JsonRpcHandler { + return { + async handleMessage(message, cx) { + const state = initializationState(cx); + if ( + message.kind === "notification" && + message.method === schema.PROTOCOL_METHODS.cancel_request + ) { + return state.status === "initializing" || state.status === "initialized" + ? Handled.yes() + : blockUninitializedIncoming(message, state); + } + if (state.status === "initialized") { + return Handled.no(message); + } + if (state.status === "initializing") { + await state.waitUntilInitialized(message.method); + return Handled.no(message); + } + return blockUninitializedIncoming(message, state); + }, + describe: () => "client-initialization", + }; +} + function parseRequestResponse( spec: { response?: ParamsParser }, response: unknown, @@ -364,6 +713,7 @@ function normalizeV2Batch( { response?: ParamsParser } | undefined >, normalizeInitialize = false, + onInitializeResponse?: (response: schema.InitializeResponse) => void, ): Entries & { readonly 0: BatchEntry } { return entries.map((entry) => { if (entry.kind !== "request") { @@ -382,6 +732,12 @@ function normalizeV2Batch( mapResponse: spec ? (response: unknown) => { const parsed = parseRequestResponse(spec, response); + if ( + onInitializeResponse && + entry.method === schema.AGENT_METHODS.initialize + ) { + onInitializeResponse(parsed as schema.InitializeResponse); + } return mapResponse ? mapResponse(parsed) : parsed; } : mapResponse, @@ -495,6 +851,11 @@ const startActiveSession = Symbol("startActiveSession"); * @experimental */ export interface AcpConnection { + /** + * Validated initialize exchange, once this connection is initialized. + */ + readonly initialization: InitializationSnapshot | undefined; + /** * AbortSignal that aborts when the connection closes. */ @@ -525,6 +886,11 @@ export interface AgentConnection extends AcpConnection { * Context for calling client-side ACP methods. */ readonly client: AgentContext; + + /** + * Client capabilities negotiated during initialization. + */ + readonly clientCapabilities: schema.ClientCapabilities | undefined; } /** @@ -541,8 +907,33 @@ export interface ClientConnection extends AcpConnection { * Context for calling agent-side ACP methods. */ readonly agent: ClientContext; + + /** + * Agent capabilities negotiated during initialization. + */ + readonly agentCapabilities: schema.AgentCapabilities | undefined; } +/** + * Handler called after an agent connection completes initialization. + * + * @experimental + */ +export type AgentInitializedHandler = ( + connection: AgentConnection, + initialization: InitializationSnapshot, +) => MaybePromise; + +/** + * Handler called after a client connection completes initialization. + * + * @experimental + */ +export type ClientInitializedHandler = ( + connection: ClientConnection, + initialization: InitializationSnapshot, +) => MaybePromise; + /** * One batch entry sent to an ACP v2 agent. * @@ -614,6 +1005,18 @@ class AcpContext { private readonly currentRequestId?: JsonRpcId, ) {} + /** + * Validated initialize exchange, once this connection is initialized. + */ + get initialization(): InitializationSnapshot | undefined { + return initializationState(this.cx).snapshot; + } + + /** @internal */ + protected get initializationLifecycle(): InitializationState { + return initializationState(this.cx); + } + /** * JSON-RPC id of the request currently being handled. * @@ -675,6 +1078,13 @@ export class AgentContext extends AcpContext { return new AgentContext(cx, requestId); } + /** + * Client capabilities negotiated during initialization. + */ + get clientCapabilities(): schema.ClientCapabilities | undefined { + return this.initialization?.request.capabilities; + } + /** * Sends a request to the client by ACP method name. * @@ -699,12 +1109,16 @@ export class AgentContext extends AcpContext { assertV2Method(method, clientRequestSpecsByMethod, "request"); const spec = clientRequestSpecsByMethod[method] as AcpRequestSpec | undefined; - return this.sendRequest( - method, - params, - spec ? (response) => parseRequestResponse(spec, response) : undefined, - options, - ); + return this.initializationLifecycle + .waitUntilInitialized(method) + .then(() => + this.sendRequest( + method, + params, + spec ? (response) => parseRequestResponse(spec, response) : undefined, + options, + ), + ); } /** @@ -732,7 +1146,9 @@ export class AgentContext extends AcpContext { "notification", true, ); - return this.sendNotification(method, params); + return this.initializationLifecycle + .waitUntilInitialized(method) + .then(() => this.sendNotification(method, params)); } /** @@ -746,12 +1162,16 @@ export class AgentContext extends AcpContext { clientRequestSpecsByMethod, clientNotificationSpecsByMethod, ); - return this.sendBatch( - normalizeV2Batch( - entries as Entries & { readonly 0: BatchEntry }, - clientRequestSpecsByMethod, - ), - ); + return this.initializationLifecycle + .waitUntilInitialized("batch") + .then(() => + this.sendBatch( + normalizeV2Batch( + entries as Entries & { readonly 0: BatchEntry }, + clientRequestSpecsByMethod, + ), + ), + ); } } @@ -774,6 +1194,13 @@ export class ClientContext extends AcpContext { return new ClientContext(cx, requestId); } + /** + * Agent capabilities negotiated during initialization. + */ + get agentCapabilities(): schema.AgentCapabilities | undefined { + return this.initialization?.response.capabilities; + } + /** @internal */ [startActiveSession]( params: schema.NewSessionRequest, @@ -888,16 +1315,41 @@ export class ClientContext extends AcpContext { assertV2Method(method, agentRequestSpecsByMethod, "request"); const spec = agentRequestSpecsByMethod[method] as AcpRequestSpec | undefined; - const wireParams = - method === schema.AGENT_METHODS.initialize - ? normalizeOutgoingV2InitializeRequest(params) - : params; - return this.sendRequest( - method, - wireParams, - spec ? (response) => parseRequestResponse(spec, response) : undefined, - options, - ); + const state = this.initializationLifecycle; + if (method === schema.AGENT_METHODS.initialize) { + const request = normalizeOutgoingV2InitializeRequest(params); + state.begin(request); + + let response: Promise; + try { + response = this.sendRequest( + method, + request, + (value) => { + const parsed = parseRequestResponse(spec!, value); + state.complete(parsed as schema.InitializeResponse); + return parsed; + }, + options, + ); + } catch (error) { + state.fail(); + throw error; + } + void response.catch(() => state.fail()); + return response; + } + + return state + .waitUntilInitialized(method) + .then(() => + this.sendRequest( + method, + params, + spec ? (response) => parseRequestResponse(spec, response) : undefined, + options, + ), + ); } /** @@ -925,7 +1377,9 @@ export class ClientContext extends AcpContext { "notification", true, ); - return this.sendNotification(method, params); + return this.initializationLifecycle + .waitUntilInitialized(method) + .then(() => this.sendNotification(method, params)); } /** @@ -939,18 +1393,82 @@ export class ClientContext extends AcpContext { agentRequestSpecsByMethod, agentNotificationSpecsByMethod, ); - return this.sendBatch( - normalizeV2Batch( - entries as Entries & { readonly 0: BatchEntry }, - agentRequestSpecsByMethod, - true, - ), + if ( + entries.length !== 1 && + entries.some( + (entry) => + entry.kind === "request" && + entry.method === schema.AGENT_METHODS.session_prompt, + ) + ) { + return Promise.reject( + RequestError.invalidRequest( + "ACP v2 session/prompt must be the only entry in its batch", + ), + ); + } + const initializeEntries = entries.filter( + (entry) => + entry.kind === "request" && + entry.method === schema.AGENT_METHODS.initialize, ); + if (initializeEntries.length > 0) { + if (entries.length !== 1 || initializeEntries.length !== 1) { + return Promise.reject( + RequestError.invalidRequest( + "ACP v2 initialize must be the only entry in its batch", + ), + ); + } + + const state = this.initializationLifecycle; + const request = normalizeOutgoingV2InitializeRequest( + initializeEntries[0].params, + ); + state.begin(request); + + let response: Promise>; + try { + response = this.sendBatch( + normalizeV2Batch( + entries as Entries & { readonly 0: BatchEntry }, + agentRequestSpecsByMethod, + true, + (value) => state.complete(value), + ), + ); + } catch (error) { + state.fail(); + throw error; + } + void response.catch(() => state.fail()); + return response; + } + + return this.initializationLifecycle + .waitUntilInitialized("batch") + .then(() => + this.sendBatch( + normalizeV2Batch( + entries as Entries & { readonly 0: BatchEntry }, + agentRequestSpecsByMethod, + ), + ), + ); } } class AcpConnectionHandle implements AcpConnection { - constructor(private readonly connection: Connection) {} + constructor(protected readonly connection: Connection) {} + + get initialization(): InitializationSnapshot | undefined { + return initializationState(this.connection.getContext()).snapshot; + } + + /** @internal */ + protected get initializationLifecycle(): InitializationState { + return initializationState(this.connection.getContext()); + } get signal(): AbortSignal { return this.connection.signal; @@ -971,23 +1489,49 @@ class AgentConnectionHandle { readonly client: AgentContext; private didStartConnectHandlers = false; + private pendingInitialization?: InitializationSnapshot; constructor( connection: Connection, - private readonly connectHandlers: readonly AgentConnectHandler[] = [], + private connectHandlers: readonly AgentConnectHandler[] = [], + private initializedHandlers: readonly AgentInitializedHandler[] = [], ) { super(connection); this.client = AgentContext.create(connection.getContext()); + this.initializationLifecycle.onInitialized((initialization) => { + this.pendingInitialization = initialization; + this.startInitializedHandlers(); + }); + } + + get clientCapabilities(): schema.ClientCapabilities | undefined { + return this.initialization?.request.capabilities; } /** @internal */ - startConnectHandlers(): void { + startConnectHandlers( + connectHandlers?: readonly AgentConnectHandler[], + initializedHandlers?: readonly AgentInitializedHandler[], + ): void { if (this.didStartConnectHandlers) { return; } + this.connectHandlers = connectHandlers ?? this.connectHandlers; + this.initializedHandlers = initializedHandlers ?? this.initializedHandlers; this.didStartConnectHandlers = true; runConnectHandlers(this, this.connectHandlers); + this.startInitializedHandlers(); + } + + private startInitializedHandlers(): void { + if (!this.didStartConnectHandlers || !this.pendingInitialization) { + return; + } + + const initialization = this.pendingInitialization; + this.pendingInitialization = undefined; + runInitializedHandlers(this, initialization, this.initializedHandlers); } } @@ -997,38 +1541,74 @@ class ClientConnectionHandle { readonly agent: ClientContext; private didStartConnectHandlers = false; + private pendingInitialization?: InitializationSnapshot; constructor( connection: Connection, - private readonly connectHandlers: readonly ClientConnectHandler[] = [], + private connectHandlers: readonly ClientConnectHandler[] = [], + private initializedHandlers: readonly ClientInitializedHandler[] = [], ) { super(connection); this.agent = ClientContext.create(connection.getContext()); + this.initializationLifecycle.onInitialized((initialization) => { + this.pendingInitialization = initialization; + this.startInitializedHandlers(); + }); + } + + get agentCapabilities(): schema.AgentCapabilities | undefined { + return this.initialization?.response.capabilities; } /** @internal */ - startConnectHandlers(): void { + startConnectHandlers( + connectHandlers?: readonly ClientConnectHandler[], + initializedHandlers?: readonly ClientInitializedHandler[], + ): void { if (this.didStartConnectHandlers) { return; } + this.connectHandlers = connectHandlers ?? this.connectHandlers; + this.initializedHandlers = initializedHandlers ?? this.initializedHandlers; this.didStartConnectHandlers = true; runConnectHandlers(this, this.connectHandlers); + this.startInitializedHandlers(); + } + + private startInitializedHandlers(): void { + if (!this.didStartConnectHandlers || !this.pendingInitialization) { + return; + } + + const initialization = this.pendingInitialization; + this.pendingInitialization = undefined; + runInitializedHandlers(this, initialization, this.initializedHandlers); } } function agentConnection( connection: Connection, connectHandlers: readonly AgentConnectHandler[] = [], + initializedHandlers: readonly AgentInitializedHandler[] = [], ): AgentConnection { - return new AgentConnectionHandle(connection, connectHandlers); + return new AgentConnectionHandle( + connection, + connectHandlers, + initializedHandlers, + ); } function clientConnection( connection: Connection, connectHandlers: readonly ClientConnectHandler[] = [], + initializedHandlers: readonly ClientInitializedHandler[] = [], ): ClientConnection { - return new ClientConnectionHandle(connection, connectHandlers); + return new ClientConnectionHandle( + connection, + connectHandlers, + initializedHandlers, + ); } type AsyncQueueEntry = @@ -1623,6 +2203,27 @@ export type AgentRequestHandler = ( context: AgentRequestContext, ) => MaybePromise; +/** + * Context passed to `session/prompt` handlers. + * + * Call and await `accept()` before emitting turn updates when work continues + * after prompt acceptance. Returning without calling it preserves the + * implicit response-after-handler behavior. + */ +export type PromptRequestContext = AgentRequestContext & { + /** + * Sends the prompt-accepted response. + */ + accept(response?: schema.PromptResponse): Promise; +}; + +/** + * Handler for the v2 `session/prompt` request. + */ +export type PromptRequestHandler = ( + context: PromptRequestContext, +) => MaybePromise; + /** * Notification handler registered on an `AgentApp`. */ @@ -1709,21 +2310,46 @@ function registerAppRequest( cx: ConnectionContext, signal: AbortSignal, requestId: JsonRpcId, + respond: (response: HandlerResponse) => Promise, ) => Context, handler: (context: Context) => MaybePromise, + lifecycle?: { + afterResponse( + params: Params, + response: Response, + cx: ConnectionContext, + ): void; + onError(cx: ConnectionContext): void; + }, ): void { builder.onReceiveRequest( spec.method, (params) => parseParams(spec.params, params), async (params, responder, cx) => { - const response = await handler( - context(params, cx, responder.signal, responder.id), - ); - await responder.respond( - spec.serializeResponse + let sentResponse: Response | undefined; + let responseWrite: Promise | undefined; + const respond = (response: HandlerResponse): Promise => { + sentResponse = spec.serializeResponse ? spec.serializeResponse(response) - : (response as unknown as Response), - ); + : (response as unknown as Response); + responseWrite = responder.respond(sentResponse); + return responseWrite; + }; + + try { + const response = await handler( + context(params, cx, responder.signal, responder.id, respond), + ); + if (!responder.responded) { + await respond(response); + } else if (responseWrite) { + await responseWrite; + } + lifecycle?.afterResponse(params, sentResponse as Response, cx); + } catch (error) { + lifecycle?.onError(cx); + throw error; + } }, ); } @@ -2055,10 +2681,7 @@ export type AgentRequestHandlersByMethod = { schema.SetSessionConfigOptionRequest, schema.SetSessionConfigOptionResponse >; - [schema.AGENT_METHODS.session_prompt]: AgentRequestHandler< - schema.PromptRequest, - schema.PromptResponse | void - >; + [schema.AGENT_METHODS.session_prompt]: PromptRequestHandler; [schema.AGENT_METHODS.mcp_message]: AgentRequestHandler< schema.MessageMcpRequest, schema.MessageMcpResponse @@ -2281,6 +2904,19 @@ function agentRequestContext( }; } +function promptRequestContext( + params: schema.PromptRequest, + client: AgentContext, + signal: AbortSignal, + requestId: JsonRpcId, + accept: (response: schema.PromptResponse | void) => Promise, +): PromptRequestContext { + return { + ...agentRequestContext(params, client, signal, requestId), + accept, + }; +} + function agentNotificationContext( params: Params, client: AgentContext, @@ -2421,6 +3057,31 @@ function runConnectHandlers( } } +function runInitializedHandlers( + connection: ConnectionHandle, + initialization: InitializationSnapshot, + handlers: ReadonlyArray< + ( + connection: ConnectionHandle, + initialization: InitializationSnapshot, + ) => MaybePromise + >, +): void { + for (const handler of handlers) { + let result: MaybePromise; + try { + result = handler(connection, cloneInitialization(initialization)); + } catch (error) { + connection.close(error); + return; + } + + void Promise.resolve(result).catch((error) => { + connection.close(error); + }); + } +} + const appBuilder = Symbol("appBuilder"); const runAgentConnectHandlers = Symbol("runAgentConnectHandlers"); const runClientConnectHandlers = Symbol("runClientConnectHandlers"); @@ -2464,8 +3125,11 @@ export function agent(options?: AppOptions): AgentApp { export class AgentApp { private readonly builder = Connection.builder(); private readonly connectHandlers: AgentConnectHandler[] = []; + private readonly initializedHandlers: AgentInitializedHandler[] = []; + private hasInitializeHandler = false; constructor(options: AppOptions = {}) { + this.builder.withHandler(agentInitializationGuard()); if (options.name) { this.builder.name(options.name); } @@ -2473,12 +3137,16 @@ export class AgentApp { /** @internal */ [appBuilder](): ConnectionBuilder { + this.assertInitializeHandler(); return this.builder; } /** @internal */ [runAgentConnectHandlers](connection: AgentConnection): void { - runConnectHandlers(connection, this.connectHandlers); + (connection as AgentConnectionHandle).startConnectHandlers( + this.connectHandlers, + this.initializedHandlers, + ); } /** @@ -2537,6 +3205,17 @@ export class AgentApp { return this; } + /** + * Registers a handler that runs once initialization succeeds. + * + * Unlike `onConnect(...)`, this hook can safely call peer methods and inspect + * the negotiated initialization snapshot. + */ + onInitialized(handler: AgentInitializedHandler): this { + this.initializedHandlers.push(handler); + return this; + } + /** * Registers a request handler by ACP method name. * @@ -2574,6 +3253,20 @@ export class AgentApp { ); } + if (method === schema.AGENT_METHODS.session_prompt) { + return this.promptRequest( + spec as AcpRequestSpec< + schema.PromptRequest, + schema.PromptResponse | void, + schema.PromptResponse + >, + handlerOrParams as PromptRequestHandler, + ); + } + + if (method === schema.AGENT_METHODS.initialize) { + this.hasInitializeHandler = true; + } return this.request( spec as AcpRequestSpec, handlerOrParams as AgentRequestHandler, @@ -2643,6 +3336,40 @@ export class AgentApp { requestId, ), handler, + spec.method === schema.AGENT_METHODS.initialize + ? { + afterResponse: (_params, response, cx) => { + initializationState(cx).complete( + response as unknown as schema.InitializeResponse, + ); + }, + onError: (cx) => initializationState(cx).fail(), + } + : undefined, + ); + return this; + } + + private promptRequest( + spec: AcpRequestSpec< + schema.PromptRequest, + schema.PromptResponse | void, + schema.PromptResponse + >, + handler: PromptRequestHandler, + ): this { + registerAppRequest( + this.builder, + spec, + (params, cx, signal, requestId, respond) => + promptRequestContext( + params, + AgentContext.create(cx, requestId), + signal, + requestId, + respond, + ), + handler, ); return this; } @@ -2665,6 +3392,7 @@ export class AgentApp { target: Stream | ClientApp, options: AppConnectOptions = {}, ): AgentConnectionState { + this.assertInitializeHandler(); if (isStream(target)) { const state = this.openStreamConnection(target); if (!options.deferConnectHandlers) { @@ -2694,9 +3422,21 @@ export class AgentApp { const rawConnection = this.builder.connect(stream); return { rawConnection, - connection: agentConnection(rawConnection, this.connectHandlers), + connection: agentConnection( + rawConnection, + this.connectHandlers, + this.initializedHandlers, + ), }; } + + private assertInitializeHandler(): void { + if (!this.hasInitializeHandler) { + throw new Error( + "AgentApp requires an initialize request handler before connecting", + ); + } + } } /** @@ -2724,11 +3464,13 @@ export function client(options?: AppOptions): ClientApp { export class ClientApp { private readonly builder = Connection.builder(); private readonly connectHandlers: ClientConnectHandler[] = []; + private readonly initializedHandlers: ClientInitializedHandler[] = []; constructor(options: AppOptions = {}) { if (options.name) { this.builder.name(options.name); } + this.builder.withHandler(clientInitializationGuard()); this.builder.withHandler({ handleMessage: (message, cx) => sessionUpdateRouter(cx).handleMessage(message), @@ -2743,7 +3485,10 @@ export class ClientApp { /** @internal */ [runClientConnectHandlers](connection: ClientConnection): void { - runConnectHandlers(connection, this.connectHandlers); + (connection as ClientConnectionHandle).startConnectHandlers( + this.connectHandlers, + this.initializedHandlers, + ); } /** @@ -2797,6 +3542,17 @@ export class ClientApp { return this; } + /** + * Registers a handler that runs once initialization succeeds. + * + * Unlike `onConnect(...)`, this hook can safely call peer methods and inspect + * the negotiated initialization snapshot. + */ + onInitialized(handler: ClientInitializedHandler): this { + this.initializedHandlers.push(handler); + return this; + } + /** * Registers a client request handler by ACP method name. * @@ -2949,7 +3705,11 @@ export class ClientApp { const rawConnection = this.builder.connect(stream); return { rawConnection, - connection: clientConnection(rawConnection, this.connectHandlers), + connection: clientConnection( + rawConnection, + this.connectHandlers, + this.initializedHandlers, + ), }; } }