diff --git a/sdk/src/codex/connection.test.ts b/sdk/src/codex/connection.test.ts index a1ba898..f0e5728 100644 --- a/sdk/src/codex/connection.test.ts +++ b/sdk/src/codex/connection.test.ts @@ -421,6 +421,121 @@ describe("CodexAxonConnection", () => { expect(handlerSignal?.aborted).toBe(true); }); + it("tracks external resolution until the handler response finishes publishing", async () => { + const { ctrl, mock, conn } = setup(); + let responseSignal: AbortSignal | undefined; + mock.axon.publish.mockImplementationOnce( + (...args: unknown[]) => + new Promise((_resolve, reject) => { + responseSignal = (args[1] as { signal?: AbortSignal } | undefined)?.signal; + responseSignal?.addEventListener("abort", () => reject(new Error("publish aborted")), { + once: true, + }); + }), + ); + conn.onApprovalRequest("mcpServer/elicitation/request", () => ({ + action: "accept", + content: { city: "Paris" }, + _meta: null, + })); + await conn.connect(); + ctrl.push( + makeAgentEvent("mcpServer/elicitation/request", { + method: "mcpServer/elicitation/request", + id: "elicit-publishing", + params: { mode: "form", message: "Which city?", requestedSchema: {} }, + }), + ); + await tick(); + expect(responseSignal).toBeDefined(); + + ctrl.push( + makeAgentEvent("serverRequest/resolved", { + method: "serverRequest/resolved", + params: { threadId: "thr-1", requestId: "elicit-publishing" }, + }), + ); + await tick(); + + expect(responseSignal?.aborted).toBe(true); + // Once Codex reports the request resolved, the in-flight response is + // canceled and its failure does not trigger a second, stale error response. + expect(mock.axon.publish).toHaveBeenCalledTimes(1); + }); + + it("cancels a fallback error response when app-server resolves the request", async () => { + const { ctrl, mock, conn } = setup(); + let responseSignal: AbortSignal | undefined; + mock.axon.publish.mockImplementationOnce( + (...args: unknown[]) => + new Promise((_resolve, reject) => { + responseSignal = (args[1] as { signal?: AbortSignal } | undefined)?.signal; + responseSignal?.addEventListener("abort", () => reject(new Error("publish aborted")), { + once: true, + }); + }), + ); + conn.onApprovalRequest("mcpServer/elicitation/request", () => { + throw new Error("handler failed"); + }); + await conn.connect(); + ctrl.push( + makeAgentEvent("mcpServer/elicitation/request", { + method: "mcpServer/elicitation/request", + id: "elicit-error-publishing", + params: { mode: "form", message: "Which city?", requestedSchema: {} }, + }), + ); + await tick(); + expect(responseSignal).toBeDefined(); + + ctrl.push( + makeAgentEvent("serverRequest/resolved", { + method: "serverRequest/resolved", + params: { threadId: "thr-1", requestId: "elicit-error-publishing" }, + }), + ); + await tick(); + + expect(responseSignal?.aborted).toBe(true); + expect(mock.axon.publish).toHaveBeenCalledTimes(1); + }); + + it("tracks default responses until they finish publishing", async () => { + const { ctrl, mock, conn } = setup(); + let responseSignal: AbortSignal | undefined; + mock.axon.publish.mockImplementationOnce( + (...args: unknown[]) => + new Promise((_resolve, reject) => { + responseSignal = (args[1] as { signal?: AbortSignal } | undefined)?.signal; + responseSignal?.addEventListener("abort", () => reject(new Error("publish aborted")), { + once: true, + }); + }), + ); + await conn.connect(); + ctrl.push( + makeAgentEvent("item/commandExecution/requestApproval", { + method: "item/commandExecution/requestApproval", + id: "default-publishing", + params: {}, + }), + ); + await tick(); + expect(responseSignal).toBeDefined(); + + ctrl.push( + makeAgentEvent("serverRequest/resolved", { + method: "serverRequest/resolved", + params: { threadId: "thr-1", requestId: "default-publishing" }, + }), + ); + await tick(); + + expect(responseSignal?.aborted).toBe(true); + expect(mock.axon.publish).toHaveBeenCalledTimes(1); + }); + it("replaces a stale parked handler when the same request is replayed", async () => { const { ctrl, mock, conn } = setup({ requestTimeoutMs: 1_000 }); const signals: AbortSignal[] = []; diff --git a/sdk/src/codex/connection.ts b/sdk/src/codex/connection.ts index 944838e..4b6c825 100644 --- a/sdk/src/codex/connection.ts +++ b/sdk/src/codex/connection.ts @@ -73,6 +73,14 @@ export type ApprovalHandler = ( context: ApprovalHandlerContext, ) => Promise | unknown; const SERVER_REQUEST_RESOLVED = Symbol("server-request-resolved"); +interface ServerRequestLifecycle { + /** Cancels any handler or outbound response when app-server clears the request. */ + signal: AbortSignal; + /** True when app-server cleared the request. */ + resolvedElsewhere: () => boolean; + /** Stop tracking the request after its final response has finished publishing. */ + release: () => void; +} /** * JSON-RPC error returned by the Codex app-server for a client request. * Preserves the wire `code` and `data` alongside the message. @@ -460,40 +468,71 @@ export class CodexAxonConnection { return { permissions: {}, scope: "turn" }; } } + private trackServerRequest(requestId: string | number): ServerRequestLifecycle { + let resolvedElsewhere = false; + const controller = new AbortController(); + const release = () => { + if (this.serverRequestResolutionWaiters.get(requestId) === onResolved) { + this.serverRequestResolutionWaiters.delete(requestId); + } + }; + const onResolved = () => { + resolvedElsewhere = true; + release(); + controller.abort(); + }; + this.serverRequestResolutionWaiters.get(requestId)?.(); + this.serverRequestResolutionWaiters.set(requestId, onResolved); + return { signal: controller.signal, resolvedElsewhere: () => resolvedElsewhere, release }; + } private approvalWithTimeout( request: ApprovalRequest, handler: ApprovalHandler, timeoutMs: number, + requestSignal: AbortSignal, ): Promise { return new Promise((resolve, reject) => { let settled = false; let timer: ReturnType | undefined; const handlerController = new AbortController(); - const finish = (value: unknown | typeof SERVER_REQUEST_RESOLVED, error?: unknown) => { + const finish = (result: unknown) => { if (settled) return; settled = true; if (timer) clearTimeout(timer); - if (this.serverRequestResolutionWaiters.get(request.id) === onResolved) { - this.serverRequestResolutionWaiters.delete(request.id); - } + requestSignal.removeEventListener("abort", onResolved); + handlerController.abort(); + resolve(result); + }; + const fail = (error: unknown) => { + if (settled) return; + settled = true; + if (timer) clearTimeout(timer); + requestSignal.removeEventListener("abort", onResolved); handlerController.abort(); - if (error !== undefined) reject(error); - else resolve(value); + reject(error); }; - const onResolved = () => finish(SERVER_REQUEST_RESOLVED); - this.serverRequestResolutionWaiters.get(request.id)?.(); - this.serverRequestResolutionWaiters.set(request.id, onResolved); + const onResolved = () => { + if (settled) return; + settled = true; + if (timer) clearTimeout(timer); + requestSignal.removeEventListener("abort", onResolved); + handlerController.abort(); + resolve(SERVER_REQUEST_RESOLVED); + }; + if (requestSignal.aborted) { + onResolved(); + return; + } + requestSignal.addEventListener("abort", onResolved, { once: true }); timer = setTimeout(() => finish(this.defaultDecline(request)), timeoutMs); Promise.resolve() .then(() => handler(request, { signal: handlerController.signal })) - .then( - (value) => finish(value), - (error) => finish(undefined, error), - ); + .then(finish, fail); }); } private async handleServerRequest(request: ServerRequest): Promise { if (this.fatal || !this.transport?.isReady()) return; + let lifecycle: ServerRequestLifecycle | undefined; try { if (!CODEX_APPROVAL_REQUEST_METHOD_SET.has(request.method)) { await this.transport.write({ @@ -502,21 +541,37 @@ export class CodexAxonConnection { }); return; } + lifecycle = this.trackServerRequest(request.id); const approval = request as ApprovalRequest; const handler = this.handlers.get(approval.method); const timeoutMs = this.options.requestTimeoutMs ?? 60_000; - const result = handler - ? await this.approvalWithTimeout(approval, handler, timeoutMs) + const outcome = handler + ? await this.approvalWithTimeout(approval, handler, timeoutMs, lifecycle.signal) : this.defaultApproval(approval); - if (result === SERVER_REQUEST_RESOLVED) return; - await this.transport?.write({ id: request.id, result }); + if (outcome === SERVER_REQUEST_RESOLVED || lifecycle.resolvedElsewhere()) return; + await this.transport?.write( + { id: request.id, result: outcome }, + { signal: lifecycle.signal }, + ); } catch (error) { - if (this.transport?.isReady()) { - await this.transport.write({ - id: request.id, - error: { code: -32000, message: error instanceof Error ? error.message : String(error) }, - }); + if (!lifecycle?.resolvedElsewhere() && this.transport?.isReady()) { + try { + await this.transport.write( + { + id: request.id, + error: { + code: -32000, + message: error instanceof Error ? error.message : String(error), + }, + }, + lifecycle ? { signal: lifecycle.signal } : undefined, + ); + } catch (publishError) { + if (!lifecycle?.resolvedElsewhere()) throw publishError; + } } + } finally { + lifecycle?.release(); } } /** diff --git a/sdk/src/codex/transport.ts b/sdk/src/codex/transport.ts index 886384b..1a63687 100644 --- a/sdk/src/codex/transport.ts +++ b/sdk/src/codex/transport.ts @@ -22,7 +22,7 @@ export interface CodexAxonTransportOptions { export interface CodexTransport { connect(): Promise; reconnect(): Promise; - write(frame: CodexFrame | string): Promise; + write(frame: CodexFrame | string, options?: { signal?: AbortSignal }): Promise; readMessages(): AsyncIterable; close(): Promise; abortStream(): void; @@ -76,8 +76,8 @@ export class CodexAxonTransport implements CodexTransport { async reconnect(): Promise { await this.inner.reconnect(); } - async write(frame: CodexFrame | string): Promise { - await this.inner.write(frame); + async write(frame: CodexFrame | string, options?: { signal?: AbortSignal }): Promise { + await this.inner.write(frame, options); } async *readMessages(): AsyncGenerator { yield* this.inner.readMessages(); diff --git a/sdk/src/shared/axon-frame-transport.ts b/sdk/src/shared/axon-frame-transport.ts index 487fea3..223a697 100644 --- a/sdk/src/shared/axon-frame-transport.ts +++ b/sdk/src/shared/axon-frame-transport.ts @@ -63,7 +63,7 @@ export class AxonFrameTransport { ); } - async write(data: TFrame | string): Promise { + async write(data: TFrame | string, options?: { signal?: AbortSignal }): Promise { if (!this.isReady()) throw new Error("Transport is not ready. Call connect() first or check isReady()."); const raw = typeof data === "string" ? data : JSON.stringify(data); @@ -71,13 +71,16 @@ export class AxonFrameTransport { if (frame === undefined && !this.options.allowInvalidOutbound) throw new Error("Cannot publish an invalid protocol frame"); if (frame !== undefined) this.options.validateOutbound?.(frame); - await this.axon.publish({ - event_type: this.options.resolveEventType(frame, raw), - origin: "USER_EVENT", - payload: raw, - source: - typeof this.options.source === "function" ? this.options.source() : this.options.source, - }); + await this.axon.publish( + { + event_type: this.options.resolveEventType(frame, raw), + origin: "USER_EVENT", + payload: raw, + source: + typeof this.options.source === "function" ? this.options.source() : this.options.source, + }, + options?.signal ? { signal: options.signal } : undefined, + ); } async *readMessages(): AsyncGenerator {