diff --git a/src/api/providers/__tests__/openai-codex-native-tool-calls.spec.ts b/src/api/providers/__tests__/openai-codex-native-tool-calls.spec.ts index e799090f48..39afbbcca7 100644 --- a/src/api/providers/__tests__/openai-codex-native-tool-calls.spec.ts +++ b/src/api/providers/__tests__/openai-codex-native-tool-calls.spec.ts @@ -526,4 +526,113 @@ describe("OpenAiCodexHandler native tool calls", () => { }), ) }) + + describe("createMessage abort signal", () => { + it("should bridge the external abortSignal into the internal AbortController", async () => { + vi.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("test-token") + vi.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + + // The mock transport pauses mid-flight until the request-local signal aborts + const mockCreate = vi.fn().mockImplementation(async (_body: unknown, init?: { signal?: AbortSignal }) => { + return { + async *[Symbol.asyncIterator]() { + yield { type: "response.text.delta", delta: "test" } + await new Promise((resolve) => { + const signal = init?.signal + if (!signal || signal.aborted) { + resolve() + return + } + signal.addEventListener("abort", () => resolve(), { once: true }) + }) + yield { + type: "response.completed", + response: { + id: "resp_1", + status: "completed", + output: [{ type: "message", content: [{ type: "output_text", text: "test" }] }], + usage: { input_tokens: 1, output_tokens: 1 }, + }, + } + // Arrives after the abort resolves the pause: the loop guard must break before + // processing the completed event and this delta. + yield { type: "response.text.delta", delta: "post-abort" } + }, + } + }) + Object.assign(handler, { + client: { + responses: { create: mockCreate }, + }, + }) + + const controller = new AbortController() + const stream = handler.createMessage("system", [{ role: "user", content: "hello" }], { + taskId: "t", + abortSignal: controller.signal, + }) + + // Consume the stream (the mock transport pauses mid-flight) + const collected = collectStream(stream) + + // Wait until the request has started; the bridge listener is registered before + // the SDK call, so aborting now lands mid-flight + await vi.waitFor(() => expect(mockCreate).toHaveBeenCalled()) + + // Abort the external signal mid-flight; the bridge must abort the request-local controller + controller.abort() + + const chunks = await collected + // Exactly the pre-abort delta: the completed event and the post-abort delta only arrive + // once the abort resolves the transport's pause, so the loop guard must break before + // processing either. + expect(chunks).toEqual([{ type: "text", text: "test" }]) + + expect(mockCreate).toHaveBeenCalled() + const createCallArgs = mockCreate.mock.calls[0][1] as { signal?: AbortSignal } + // The captured (request-local) signal passed to the SDK must now be aborted + expect(createCallArgs.signal).toBeDefined() + expect(createCallArgs.signal).toBeInstanceOf(AbortSignal) + expect(createCallArgs.signal?.aborted).toBe(true) + }) + + it("should immediately abort when the external signal is already aborted", async () => { + vi.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("test-token") + vi.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + + const mockCreate = vi.fn().mockResolvedValue({ + async *[Symbol.asyncIterator]() { + yield { type: "response.text.delta", delta: "test" } + yield { + type: "response.completed", + response: { + id: "resp_1", + status: "completed", + output: [{ type: "message", content: [{ type: "output_text", text: "test" }] }], + usage: { input_tokens: 1, output_tokens: 1 }, + }, + } + }, + }) + Object.assign(handler, { + client: { + responses: { create: mockCreate }, + }, + }) + + const controller = new AbortController() + controller.abort() // Pre-abort + + const stream = handler.createMessage("system", [{ role: "user", content: "hello" }], { + taskId: "t", + abortSignal: controller.signal, + }) + + // A pre-aborted signal settles in the OAuth race before the SDK is reached: the abort + // contract must win over the request setup, so the stream rejects and the SDK is never + // called with a signal that is already aborted. + await expect(collectStream(stream)).rejects.toMatchObject({ name: "AbortError" }) + expect(mockCreate).not.toHaveBeenCalled() + }) + }) }) diff --git a/src/api/providers/__tests__/openai-codex.spec.ts b/src/api/providers/__tests__/openai-codex.spec.ts index 35c1f5de60..845e341681 100644 --- a/src/api/providers/__tests__/openai-codex.spec.ts +++ b/src/api/providers/__tests__/openai-codex.spec.ts @@ -9,6 +9,7 @@ vitest.mock("@roo-code/telemetry", () => ({ })) import { Anthropic } from "@anthropic-ai/sdk" +import { TelemetryService } from "@roo-code/telemetry" import { OPEN_AI_CODEX_SERVICE_TIER_KEY, OpenAiCodexServiceTier, SERVICE_TIER_KEY } from "@roo-code/types" import { OpenAiCodexHandler, transformResponsesLiteBody } from "../openai-codex" import { openAiCodexOAuthManager } from "../../../integrations/openai-codex/oauth" @@ -491,17 +492,73 @@ describe("OpenAiCodexHandler.completePrompt streaming", () => { expect(signalDuringRequest!.aborted).toBe(true) }) - it("rejects when the caller's signal is already aborted", async () => { - const handler = createHandler() - const create = injectStream(handler, [ - { type: "response.completed", response: { id: "r1", status: "completed", output: [] } }, - ]) + // A request cancelled before it starts must not spend the provider setup on it: the + // fast-fail rejects before the (here deferred, never resolving) token and account fetch + // and before the SDK request, so a pending OAuth flow cannot hold a cancelled completion + // hostage. + it("fast-fails before the OAuth setup when the caller's signal is already aborted", async () => { + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol" }) + const getAccessToken = vitest + .spyOn(openAiCodexOAuthManager, "getAccessToken") + .mockReturnValue(new Promise(() => {})) + const getAccountId = vitest + .spyOn(openAiCodexOAuthManager, "getAccountId") + .mockReturnValue(new Promise(() => {})) + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) await expect(handler.completePrompt("Hello", { abortSignal: AbortSignal.abort() })).rejects.toMatchObject({ name: "AbortError", + message: "This operation was aborted", }) - expect(create.mock.calls[0][1].signal.aborted).toBe(true) + expect(create).not.toHaveBeenCalled() + expect(getAccessToken).not.toHaveBeenCalled() + expect(getAccountId).not.toHaveBeenCalled() + }) + + // The abort lands while event2 is being processed: its `delta` getter fires it after + // executeRequest's top-of-loop check has passed. The pre-yield guard must end the request + // there, so the stream is never pulled a third time - the pull count distinguishes a flow + // that kept pulling (a mutated guard that lets the loop's continuation run) from the + // guarded one, and the buffered "post-abort" chunk must never be joined. + it("settles with the abort contract and stops pulling once an event's processing aborts the request", async () => { + const handler = createHandler() + const controller = new AbortController() + let sdkPulls = 0 + + const event1 = { type: "response.output_text.delta", delta: "pre-abort" } + const event2 = { + type: "response.output_text.delta", + get delta() { + controller.abort() + return "post-abort" + }, + } + + const create = vitest.fn().mockImplementation(() => { + return Promise.resolve({ + [Symbol.asyncIterator]() { + return { + next: async () => { + sdkPulls++ + return { value: sdkPulls === 1 ? event1 : event2, done: false } + }, + return: async () => ({ value: undefined, done: true }), + } + }, + }) + }) + Reflect.set(handler, "client", { responses: { create } }) + + await expect(handler.completePrompt("Hello", { abortSignal: controller.signal })).rejects.toMatchObject({ + name: "AbortError", + }) + + // event1 is pulled and joined; event2 is pulled (its getter fires the abort) and the + // pre-yield guard ends the request before its buffered chunk is joined - a third pull + // only happens when a mutated guard lets the event loop's continuation run. + expect(sdkPulls).toBe(2) }) // The SSE fallback is for an SDK that could not be used at all. Replaying the request after the @@ -591,6 +648,72 @@ describe("OpenAiCodexHandler.completePrompt streaming", () => { expect(mockFetch).not.toHaveBeenCalled() }) + // The cancellation wins over the auth retry: force-refreshing a token for a request that is + // already gone would spend a network round trip on a dead request and surface the + // cancellation as an authentication failure. + it("fails fast with the abort contract instead of force-refreshing the token after cancellation", async () => { + const handler = createHandler() + const refresh = vitest.spyOn(openAiCodexOAuthManager, "forceRefreshAccessToken") + const controller = new AbortController() + // The SDK keeps the request open until the signal it was handed aborts, then rejects the way + // it rejects aborted requests. + const create = vitest.fn().mockImplementation( + (_body: unknown, options: { signal: AbortSignal }) => + new Promise((_resolve, reject) => { + options.signal.addEventListener("abort", () => reject(new Error("Request was aborted")), { + once: true, + }) + }), + ) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const completion = handler.completePrompt("Hello", { abortSignal: controller.signal }) + await vitest.waitFor(() => expect(create).toHaveBeenCalled()) + controller.abort() + + await expect(completion).rejects.toMatchObject({ name: "AbortError" }) + expect(refresh).not.toHaveBeenCalled() + expect(create).toHaveBeenCalledTimes(1) + expect(mockFetch).not.toHaveBeenCalled() + }) + + // The abort lands while the fallback fetch is in flight, so the cancellation must come out as + // the shared abort contract - not a telemetry event and not a wrapped connection error. + it("keeps the abort contract and skips telemetry when the fallback fetch is cancelled", async () => { + const handler = createHandler() + // The module mock keeps the spy across tests, so clear it before asserting on this request + const captureException = vitest.mocked(TelemetryService.instance.captureException) + captureException.mockClear() + // The SDK path is unusable, so the request falls back to the SSE transport. + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + // Reject the way fetch rejects once the signal it was handed aborts. + const mockFetch = vitest.fn((_url: unknown, init?: { signal?: AbortSignal }) => { + const signal = init?.signal + if (!signal || signal.aborted) { + return Promise.reject(new DOMException("The operation was aborted", "AbortError")) + } + return new Promise((_resolve, reject) => { + signal.addEventListener( + "abort", + () => reject(new DOMException("The operation was aborted", "AbortError")), + { once: true }, + ) + }) + }) + vitest.stubGlobal("fetch", mockFetch) + + const completion = handler.completePrompt("Hello", { abortSignal: controller.signal }) + await vitest.waitFor(() => expect(mockFetch).toHaveBeenCalled()) + controller.abort() + + await expect(completion).rejects.toMatchObject({ name: "AbortError" }) + expect(captureException).not.toHaveBeenCalled() + }) + it("wraps failures from both transports as a completion error", async () => { const handler = createHandler() const create = vitest.fn().mockRejectedValue(new Error("sdk down")) @@ -931,6 +1054,48 @@ describe("OpenAiCodexHandler Responses Lite requests", () => { }) }) + it("settles with the abort contract when cancellation lands during the auth retry refresh", async () => { + // The first request fails authentication and the forced token refresh is in flight + // when the caller aborts. The refresh must be raced against the signal: without the + // race the handler stays suspended until the refresh settles and then surfaces the + // authentication failure of a request that no longer exists instead of the abort. + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.6-luna" }) + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("expired-token") + vitest.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + // Object.assign sidesteps the SDK client's structural type so the double stays cast-free. + Object.assign(handler, { + client: { responses: { create: vitest.fn().mockRejectedValue(new Error("SDK unavailable")) } }, + }) + // The OAuth operation settles (with no token) after the test aborts, mirroring a + // slow refresh that loses the race. + const refresh = vitest + .spyOn(openAiCodexOAuthManager, "forceRefreshAccessToken") + .mockImplementation(() => new Promise((resolve) => setTimeout(() => resolve(null), 50))) + const mockFetch = vitest.fn().mockResolvedValueOnce({ + ok: false, + status: 401, + text: vitest.fn().mockResolvedValue('{"error":{"message":"Codex API invalid token"}}'), + }) + vitest.stubGlobal("fetch", mockFetch) + + const abortController = new AbortController() + const result = collectStream( + handler.createMessage("Instructions", [{ role: "user", content: "Abort during refresh" }], { + taskId: "task-abort-refresh", + tools: [], + abortSignal: abortController.signal, + }), + ) + // Let the first request fail and the refresh start, then cancel the caller. + await new Promise((resolve) => setTimeout(resolve, 10)) + abortController.abort() + + await expect(result).rejects.toMatchObject({ name: "AbortError" }) + expect(refresh).toHaveBeenCalledTimes(1) + // No retry request goes out for a request that was aborted mid-refresh. + expect(mockFetch).toHaveBeenCalledTimes(1) + }) + it.each(["gpt-5.5", "gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna-alias"])( "does not apply Luna behavior to %s", async (apiModelId) => { @@ -1012,3 +1177,780 @@ describe("OpenAiCodexHandler Responses Lite requests", () => { }) }) }) + +describe("OpenAiCodexHandler.completePrompt timeout", () => { + function createHandler() { + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.1-codex" }) + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("test-token") + vitest.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + return handler + } + + function injectStream(handler: OpenAiCodexHandler, events: unknown[]) { + const create = vitest.fn().mockResolvedValue(asyncStreamFrom(events)) + Reflect.set(handler, "client", { responses: { create } }) + return create + } + + afterEach(() => { + vitest.restoreAllMocks() + vitest.unstubAllGlobals() + }) + + // timeoutMs <= 0 must not install a timer: the completion runs to the end and the + // request signal never aborts on its own. + it("treats timeoutMs=0 as no timeout", async () => { + const handler = createHandler() + const create = injectStream(handler, [ + { type: "response.output_text.delta", delta: "response" }, + { type: "response.completed", response: { id: "r1", status: "completed", output: [] } }, + ]) + + await expect(handler.completePrompt("test prompt", { timeoutMs: 0 })).resolves.toBe("response") + + expect(create.mock.calls[0][1].signal.aborted).toBe(false) + }) + + // A timeout and an external abort must cancel the same request: the signal the transport + // receives aborts when either of them fires. + it("merges abortSignal and timeoutMs into the request signal", async () => { + const handler = createHandler() + const controller = new AbortController() + let signalDuringRequest: AbortSignal | undefined + + const create = vitest.fn().mockImplementation((_body: unknown, options: { signal: AbortSignal }) => { + signalDuringRequest = options.signal + // Abort mid-flight, while the listener linking the caller signal is attached. + controller.abort() + return Promise.resolve( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "feat: half a" }, + { type: "response.completed", response: { id: "r1", status: "completed", output: [] } }, + ]), + ) + }) + Reflect.set(handler, "client", { responses: { create } }) + + await expect( + handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 5000 }), + ).rejects.toMatchObject({ name: "AbortError" }) + + expect(signalDuringRequest).toBeInstanceOf(AbortSignal) + expect(signalDuringRequest!.aborted).toBe(true) + }) + + // Native AbortSignal.timeout self-manages its timer, so a transport that rejects on abort + // is all the timeout needs to be observed end to end. + it("rejects with an AbortError when the timeout elapses", async () => { + const handler = createHandler() + let signalDuringRequest: AbortSignal | undefined + + const create = vitest.fn().mockImplementation( + (_body: unknown, options: { signal: AbortSignal }) => + new Promise((_resolve, reject) => { + signalDuringRequest = options.signal + // Reject the way the SDK does once the signal it was handed aborts. + options.signal.addEventListener("abort", () => reject(new Error("The operation was aborted")), { + once: true, + }) + }), + ) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + await expect(handler.completePrompt("test prompt", { timeoutMs: 50 })).rejects.toMatchObject({ + name: "AbortError", + }) + + // The timeout must have fired on the request signal, and a timed-out completion is not a + // transport failure, so it must not be replayed over SSE. + expect(signalDuringRequest!.aborted).toBe(true) + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("rejects with the abort contract when the timeout fires while the token lookup is pending", async () => { + const handler = createHandler() + let resolveToken!: (token: string) => void + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockReturnValue( + new Promise((resolve) => { + resolveToken = resolve + }), + ) + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) + + // No abort signal: the timeout is the only cancellation source, so the timeout path must + // race the OAuth lookup itself. + await expect(handler.completePrompt("test prompt", { timeoutMs: 20 })).rejects.toMatchObject({ + name: "AbortError", + }) + + // The lookup settling afterwards must be ignored: the request is already gone. + resolveToken("late-token") + expect(create).not.toHaveBeenCalled() + }) +}) + +describe("OpenAiCodexHandler.createMessage abort bridging", () => { + // These tests drive createMessage directly (not completePrompt): completePrompt re-normalizes + // any error to the shared abort contract once the caller's signal has fired, which would mask + // regressions in the abort checks inside the transports. + + function createHandler() { + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol" }) + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("test-token") + vitest.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + return handler + } + + afterEach(() => { + vitest.restoreAllMocks() + vitest.unstubAllGlobals() + }) + + it("hands back the abort contract instead of force-refreshing after the caller cancels", async () => { + const handler = createHandler() + const refresh = vitest + .spyOn(openAiCodexOAuthManager, "forceRefreshAccessToken") + .mockResolvedValue("refreshed-token") + // The SDK is wired to fail with exactly the auth-failure wording the retry path would act + // on: the cancellation must win over the refresh-and-retry logic at every level. + const create = vitest.fn().mockRejectedValue(new Error("401 invalid token")) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + await expect( + collectStream( + handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: AbortSignal.abort(), + }), + ), + ).rejects.toMatchObject({ name: "AbortError" }) + + // A pre-aborted signal settles in the OAuth race before the SDK is even reached, so there + // is no SDK attempt, no refresh, and no SSE fallback. + expect(refresh).not.toHaveBeenCalled() + expect(create).not.toHaveBeenCalled() + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("settles with the abort contract before a pending token lookup resolves", async () => { + const handler = createHandler() + let resolveToken!: (token: string) => void + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockReturnValue( + new Promise((resolve) => { + resolveToken = resolve + }), + ) + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) + + const controller = new AbortController() + const stream = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + // The first pull runs the generator up to the token lookup, which is still pending. + const firstPull = stream.next() + controller.abort() + + await expect(firstPull).rejects.toMatchObject({ name: "AbortError" }) + + // The lookup settling afterwards must be ignored: the request is already gone. + resolveToken("late-token") + await new Promise((resolve) => setTimeout(resolve, 0)) + expect(create).not.toHaveBeenCalled() + }) + + it("settles with the abort contract before a pending account lookup resolves", async () => { + const handler = createHandler() + let resolveAccount!: (id: string) => void + vitest.spyOn(openAiCodexOAuthManager, "getAccountId").mockReturnValue( + new Promise((resolve) => { + resolveAccount = resolve + }), + ) + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) + + const controller = new AbortController() + const stream = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + const firstPull = stream.next() + // Let the token lookup settle so the generator reaches the account lookup. + await new Promise((resolve) => setTimeout(resolve, 0)) + controller.abort() + + await expect(firstPull).rejects.toMatchObject({ name: "AbortError" }) + + resolveAccount("late-account") + await new Promise((resolve) => setTimeout(resolve, 0)) + expect(create).not.toHaveBeenCalled() + }) + + it("settles with the abort contract before a pending account lookup resolves on the fallback", async () => { + const handler = createHandler() + let resolveAccount!: (id: string) => void + // The executeRequest-level lookup settles; the fallback's own lookup stays pending. + vitest + .spyOn(openAiCodexOAuthManager, "getAccountId") + .mockImplementationOnce(() => Promise.resolve("acct_sdk")) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveAccount = resolve + }), + ) + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const controller = new AbortController() + const stream = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + const firstPull = stream.next() + // Let the SDK fail and the fallback reach its own account lookup. + await new Promise((resolve) => setTimeout(resolve, 0)) + controller.abort() + + await expect(firstPull).rejects.toMatchObject({ name: "AbortError" }) + + resolveAccount("late-account") + await new Promise((resolve) => setTimeout(resolve, 0)) + expect(create).toHaveBeenCalledTimes(1) + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("never yields the buffered chunks of an event whose processing aborts the request", async () => { + const handler = createHandler() + const controller = new AbortController() + const event1 = { type: "response.output_text.delta", delta: "pre-abort" } + const event2 = { + type: "response.output_text.delta", + get delta() { + controller.abort() + return "post-abort" + }, + } + const create = vitest.fn().mockImplementation(() => { + return Promise.resolve({ + [Symbol.asyncIterator]() { + let pulls = 0 + return { + next: async () => { + pulls++ + return { value: pulls === 1 ? event1 : event2, done: false } + }, + return: async () => ({ value: undefined, done: true }), + } + }, + }) + }) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + // event1's chunk streams; the abort lands while event2 is being processed, so its buffered + // "post-abort" chunk must not reach the caller - createMessage has no consumer-side guard. + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "pre-abort" } }) + await expect(iter.next()).rejects.toMatchObject({ name: "AbortError" }) + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("never yields a buffered SSE line after a cancellation that lands between lines", async () => { + const handler = createHandler() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + const encoder = new TextEncoder() + const body = new ReadableStream({ + start(streamController) { + // A single chunk: both lines land in one read so the second is buffered in the + // handler's line loop when the cancellation lands. + streamController.enqueue( + encoder.encode( + 'data: {"type":"response.output_text.delta","delta":"one"}\n\ndata: {"type":"response.output_text.delta","delta":"two"}\n\n', + ), + ) + streamController.close() + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "one" } }) + controller.abort() + // The second line was enqueued before the cancellation but must not stream: the delegate + // guard hands back the abort contract instead of yielding its chunk. + await expect(iter.next()).rejects.toMatchObject({ name: "AbortError" }) + }) + + it("never yields the remaining complete-response chunks once the caller cancels", async () => { + const handler = createHandler() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + const encoder = new TextEncoder() + const complete = { + response: { + id: "resp_test", + output: [ + { type: "text", content: [{ type: "text", text: "first" }] }, + { type: "text", content: [{ type: "text", text: "second" }] }, + ], + }, + } + const body = new ReadableStream({ + start(streamController) { + streamController.enqueue(encoder.encode(`data: ${JSON.stringify(complete)}\n\n`)) + streamController.close() + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "first" } }) + controller.abort() + // The second text item was buffered before the cancellation but must not stream: the + // guard before the yield breaks out of the item loop, and the while-top guard ends the read. + const chunks = await collectStream(iter) + expect(chunks).toEqual([]) + }) + + it("streams every direct SSE chunk shape the fallback supports", async () => { + const handler = createHandler() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const encoder = new TextEncoder() + const complete = { + response: { + output: [ + { type: "text", content: [{ type: "text", text: "complete-text" }] }, + { type: "reasoning", summary: [{ type: "summary_text", text: "complete-summary" }] }, + ], + usage: { input_tokens: 7, output_tokens: 11 }, + }, + } + const lines = [ + `data: ${JSON.stringify(complete)}`, + 'data: {"type":"response.output_text.delta","delta":"delegated"}', + 'data: {"choices":[{"delta":{"content":"choices-text"}}]}', + 'data: {"item":{"type":"text","text":"item-text"}}', + 'data: {"usage":{"input_tokens":3,"output_tokens":4}}', + '{"content":"plain-text"}', + ] + const body = new ReadableStream({ + start(streamController) { + for (const line of lines) { + streamController.enqueue(encoder.encode(`${line}\n\n`)) + } + streamController.close() + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const chunks = await collectStream( + handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + }), + ) + + // One stream, every direct SSE shape: the complete-response block (text, reasoning, + // usage), the delegated delta, the legacy choices/item/usage shapes, and a plain JSON + // line. The sequence also pins the order the fallback emits them in. + expect(chunks).toEqual([ + { type: "text", text: "complete-text" }, + { type: "reasoning", text: "complete-summary" }, + { type: "usage", inputTokens: 7, outputTokens: 11, cacheWriteTokens: 0, cacheReadTokens: 0, totalCost: 0 }, + { type: "text", text: "delegated" }, + { type: "text", text: "choices-text" }, + { type: "text", text: "item-text" }, + { type: "usage", inputTokens: 3, outputTokens: 4, cacheWriteTokens: 0, cacheReadTokens: 0, totalCost: 0 }, + { type: "text", text: "plain-text" }, + ]) + }) + + it("propagates a non-abort token failure as-is when no cancellation is involved", async () => { + const handler = createHandler() + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockRejectedValue(new Error("oauth down")) + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) + + // No abort signal: a genuine lookup failure must reach the caller unchanged, not be + // normalized into the abort contract. + await expect( + collectStream( + handler.createMessage("System", [{ role: "user", content: "Hello" }], { taskId: "task-test" }), + ), + ).rejects.toThrow("oauth down") + + expect(create).not.toHaveBeenCalled() + }) + + it("settles with the abort contract when the SDK stream fails at the moment the caller cancels", async () => { + const handler = createHandler() + const refresh = vitest + .spyOn(openAiCodexOAuthManager, "forceRefreshAccessToken") + .mockResolvedValue("refreshed-token") + const controller = new AbortController() + const event1 = { type: "response.output_text.delta", delta: "pre-abort" } + const create = vitest.fn().mockImplementation(() => { + return Promise.resolve({ + [Symbol.asyncIterator]() { + let pulls = 0 + return { + next: async () => { + pulls++ + if (pulls === 1) { + return { value: event1, done: false } + } + // The cancellation lands while the second pull is in flight: the top-of-loop + // check has already passed, so the failure must settle via the retry-catch + // abort check, not the auth-refresh path. + controller.abort() + throw new Error("401 invalid token") + }, + return: async () => ({ value: undefined, done: true }), + } + }, + }) + }) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "pre-abort" } }) + await expect(iter.next()).rejects.toMatchObject({ name: "AbortError" }) + + // The cancellation wins over the auth-failure wording: no refresh, no SSE fallback. + expect(refresh).not.toHaveBeenCalled() + expect(create).toHaveBeenCalledTimes(1) + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("settles with the abort contract when the signal aborts as the token lookup settles", async () => { + const handler = createHandler() + const controller = new AbortController() + const create = vitest.fn() + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + // The first pull runs the generator up to the token race. The token is already resolved, so + // the race settles on the next microtask and detaches its abort listener; queueing the abort + // after that microtask lands it in the window where only the request-local bridge can still + // see it. + const firstPull = iter.next() + queueMicrotask(() => controller.abort()) + + await expect(firstPull).rejects.toMatchObject({ name: "AbortError" }) + + expect(create).not.toHaveBeenCalled() + expect(mockFetch).not.toHaveBeenCalled() + }) + + // Each case streams one line before the cancellation and buffers the target line in the + // handler's line loop: the target's pre-yield guard must break before its chunk streams. + const sseData = (obj: unknown) => "data: " + JSON.stringify(obj) + const completeTextLine = (t: string) => + sseData({ response: { output: [{ type: "text", content: [{ type: "text", text: t }] }] } }) + const completeReasoningLine = (t: string) => + sseData({ response: { output: [{ type: "reasoning", summary: [{ type: "summary_text", text: t }] }] } }) + const completeUsageLine = (t: string) => + sseData({ + response: { + output: [{ type: "text", content: [{ type: "text", text: t }] }], + usage: { input_tokens: 1, output_tokens: 1 }, + }, + }) + const choicesLine = (t: string) => sseData({ choices: [{ delta: { content: t } }] }) + const itemLine = (t: string) => sseData({ item: { type: "text", text: t } }) + const usageLine = sseData({ usage: { input_tokens: 1, output_tokens: 1 } }) + const plainLine = (t: string) => JSON.stringify({ content: t }) + + const sseGuardCases: [string, string, string, Record][] = [ + [ + "a complete-response reasoning summary", + completeTextLine("a"), + completeReasoningLine("b"), + { type: "text", text: "a" }, + ], + ["a complete-response usage chunk", completeTextLine("a"), completeUsageLine("b"), { type: "text", text: "a" }], + ["a legacy choices delta", itemLine("a"), choicesLine("b"), { type: "text", text: "a" }], + ["a legacy item text", choicesLine("a"), itemLine("b"), { type: "text", text: "a" }], + ["a legacy usage object", itemLine("a"), usageLine, { type: "text", text: "a" }], + ["a plain JSON line", usageLine, plainLine("b"), { type: "usage", inputTokens: 1, outputTokens: 1 }], + ] + for (const [name, firstLine, targetLine, firstChunk] of sseGuardCases) { + it(`never yields ${name} once the caller cancels before the handler reads it`, async () => { + const handler = createHandler() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + const encoder = new TextEncoder() + const body = new ReadableStream({ + start(streamController) { + // A single chunk: both lines land in one read so the target is buffered in the + // handler's line loop when the cancellation lands. + streamController.enqueue(encoder.encode(`${firstLine}\n\n${targetLine}\n\n`)) + streamController.close() + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: firstChunk }) + controller.abort() + // The target line was buffered before the cancellation, but its pre-yield guard + // breaks out of the line loop and the while-top guard ends the read. + expect(await collectStream(iter)).toEqual([]) + }) + } + + it("emits nothing once the caller cancels before the first SDK event", async () => { + const handler = createHandler() + const controller = new AbortController() + const create = vitest.fn().mockImplementation(() => { + // The caller cancels while the SDK is still delivering, so the loop's abort check must + // stop every event from reaching the caller. + controller.abort() + return Promise.resolve( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "feat: half a" }, + { type: "response.output_text.delta", delta: "and the rest" }, + { type: "response.completed", response: { id: "r1", status: "completed", output: [] } }, + ]), + ) + }) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + const chunks = await collectStream( + handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }), + ) + + // The generators end quietly on abort - they break rather than throw - so an empty stream + // is the observable proof that nothing was processed after the cancellation. + expect(chunks).toEqual([]) + expect(mockFetch).not.toHaveBeenCalled() + }) + + it("clears the request controller once the request ends", async () => { + const handler = createHandler() + const create = vitest.fn().mockResolvedValue( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "done" }, + { type: "response.completed", response: { id: "r1", status: "completed", output: [] } }, + ]), + ) + Reflect.set(handler, "client", { responses: { create } }) + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + await collectStream(handler.createMessage("System", [{ role: "user", content: "Hello" }])) + + // A later request must be able to tell that no request is in flight. + expect(Reflect.get(handler, "abortController")).toBeUndefined() + }) + + it("keeps the later request's controller when the earlier request finishes", async () => { + const handler = createHandler() + const mockFetch = vitest.fn() + vitest.stubGlobal("fetch", mockFetch) + + // The earlier request holds its stream open until the later request has installed its own + // controller, so the earlier cleanup runs while a request is still in flight. + let releaseEarlier: (() => void) | undefined + const releaseGate = new Promise((resolve) => { + releaseEarlier = resolve + }) + const earlierStream = (async function* () { + yield { type: "response.output_text.delta", delta: "a" } + await releaseGate + yield { type: "response.completed", response: { id: "ra", status: "completed", output: [] } } + })() + const create = vitest + .fn() + .mockImplementationOnce(() => Promise.resolve(earlierStream)) + .mockImplementationOnce(() => + Promise.resolve( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "b" }, + { type: "response.completed", response: { id: "rb", status: "completed", output: [] } }, + ]), + ), + ) + Reflect.set(handler, "client", { responses: { create } }) + + const earlier = handler.createMessage("System", [{ role: "user", content: "Hello" }]) + expect(await earlier.next()).toMatchObject({ value: { type: "text", text: "a" } }) + + const later = handler.createMessage("System", [{ role: "user", content: "Hello" }]) + expect(await later.next()).toMatchObject({ value: { type: "text", text: "b" } }) + + // The earlier request finishes while the later one is still in flight. + releaseEarlier!() + await earlier.next() + + // The earlier request's cleanup must not clear the controller the later request installed. + const controller = Reflect.get(handler, "abortController") as AbortController | undefined + expect(controller).toBeDefined() + expect(controller?.signal).toBe(create.mock.calls[1][1].signal) + + // Let the later request finish and clear its own controller. + await later.next() + }) + + it("stops reading the fallback stream once the request aborts", async () => { + const handler = createHandler() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + const encoder = new TextEncoder() + const body = new ReadableStream({ + start(streamController) { + streamController.enqueue( + encoder.encode('data: {"type":"response.output_text.delta","delta":"one"}\n\n'), + ) + streamController.enqueue( + encoder.encode('data: {"type":"response.output_text.delta","delta":"two"}\n\n'), + ) + streamController.close() + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "one" } }) + + // The cancellation lands between two stream reads, exactly where the loop's check runs. + controller.abort() + + const chunks = await collectStream(iter) + // Everything enqueued after the cancellation must stay unread. + expect(chunks).toEqual([]) + }) + + it("hands back the abort contract when the fallback stream tears down while the caller cancels", async () => { + const handler = createHandler() + const captureException = vitest.mocked(TelemetryService.instance.captureException) + captureException.mockClear() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const controller = new AbortController() + const encoder = new TextEncoder() + let pullStartedResolve: (() => void) | undefined + let failRead: (() => void) | undefined + const pullStarted = new Promise((resolve) => { + pullStartedResolve = resolve + }) + const failGate = new Promise((resolve) => { + failRead = resolve + }) + const body = new ReadableStream({ + start(streamController) { + streamController.enqueue( + encoder.encode('data: {"type":"response.output_text.delta","delta":"one"}\n\n'), + ) + }, + pull(streamController) { + // The second read is pending; tear the stream down on the test's signal. + pullStartedResolve!() + return failGate.then(() => { + streamController.error(new Error("stream torn down")) + }) + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }], { + taskId: "task-test", + abortSignal: controller.signal, + }) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "one" } }) + + const pending = collectStream(iter) + await pullStarted + // The cancellation lands while the second read is in flight, i.e. after the loop's check. + controller.abort() + failRead!() + + await expect(pending).rejects.toMatchObject({ name: "AbortError" }) + // Cancellation is the caller's own doing, so neither catch may report it to telemetry. + expect(captureException).not.toHaveBeenCalled() + }) + + it("wraps a torn-down fallback stream as a stream error when nothing was aborted", async () => { + const handler = createHandler() + const captureException = vitest.mocked(TelemetryService.instance.captureException) + captureException.mockClear() + const create = vitest.fn().mockRejectedValue(new Error("sdk down")) + Reflect.set(handler, "client", { responses: { create } }) + const encoder = new TextEncoder() + const body = new ReadableStream({ + start(streamController) { + streamController.enqueue( + encoder.encode('data: {"type":"response.output_text.delta","delta":"one"}\n\n'), + ) + }, + pull(streamController) { + streamController.error(new Error("stream torn down")) + }, + }) + const mockFetch = vitest.fn().mockResolvedValue({ ok: true, body }) + vitest.stubGlobal("fetch", mockFetch) + + const iter = handler.createMessage("System", [{ role: "user", content: "Hello" }]) + expect(await iter.next()).toMatchObject({ value: { type: "text", text: "one" } }) + + // The wrap chain surfaces the connection-failure key in every case, so the message alone + // cannot prove the innermost check classified this as a stream failure: an always-true + // check would swap in the shared abort contract, and the request-level catch would wrap + // that in the same key. The telemetry count is the witness: the stream-processing catch + // and the request catch both report the failure (twice); the abort path skips the first. + await expect(collectStream(iter)).rejects.toThrow(/connectionFailed|stream torn down/) + expect(captureException).toHaveBeenCalledTimes(2) + }) +}) diff --git a/src/api/providers/openai-codex.ts b/src/api/providers/openai-codex.ts index 1498aafc66..470efb8a9b 100644 --- a/src/api/providers/openai-codex.ts +++ b/src/api/providers/openai-codex.ts @@ -29,6 +29,13 @@ import { isMcpTool } from "../../utils/mcp-name" import { sanitizeOpenAiCallId } from "../../utils/tool-id" import { openAiCodexOAuthManager } from "../../integrations/openai-codex/oauth" import { t } from "../../i18n" +import { + createAbortError, + isRequestAborted, + mergeAbortSignalAndTimeout, + rejectOnAbort, + throwIfAborted, +} from "./utils/abort-signal" export type OpenAiCodexModel = ReturnType @@ -250,7 +257,29 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion this.streamedToolCallIds.clear() // Get access token from OAuth manager - let accessToken = await openAiCodexOAuthManager.getAccessToken() + // The lookup can await credential loading, token refresh, and persistence: race it against + // the caller's signal so a cancelled request settles immediately instead of hanging on a + // lookup that no longer matters. + let accessToken: string | null + try { + if (abortSignal) { + accessToken = await rejectOnAbort( + openAiCodexOAuthManager.getAccessToken(), + abortSignal, + this.providerName, + ) + } else { + accessToken = await openAiCodexOAuthManager.getAccessToken() + } + } catch (error) { + // An abort landing while the lookup is pending must surface as the shared abort + // contract, not as an authentication failure. + // Stryker disable next-line ConditionalExpression: the false mutant differs from the original only if an abort lands between the OAuth race's settle and this catch — a microtask window (the abort event is a microtask) no test can schedule deterministically, because the test's own microtasks queue after the settle that precedes the window; the observable contract (a genuine failure with no cancellation propagates as-is, an abort surfaces as the shared contract) is pinned by the adjacent regressions and the true mutant is killed by the non-abort-failure regression. + if (isRequestAborted(error, abortSignal)) { + throw createAbortError(this.providerName) + } + throw error + } if (!accessToken) { throw new Error( t("common:errors.openAiCodex.notAuthenticated", { @@ -290,6 +319,13 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion yield* this.executeRequest(requestBody, model, accessToken, effectiveSessionId, abortSignal) return } catch (error) { + // The caller's cancellation wins over the retry: force-refreshing a token for a + // request that is already gone would spend a network round trip on a dead request + // and surface the cancellation as an authentication failure. + if (abortSignal?.aborted) { + throw createAbortError(this.providerName) + } + const message = error instanceof Error ? error.message : String(error) const isAuthFailure = /unauthorized|invalid token|not authenticated|authentication|401/i.test(message) @@ -297,8 +333,15 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion // the service has accepted the request, so a refreshed-token retry would replay it // and append a second generation to output the caller already has. if (attempt === 0 && isAuthFailure && !this.sawSdkEventInCurrentResponse) { - // Force refresh the token for retry - const refreshed = await openAiCodexOAuthManager.forceRefreshAccessToken() + // Force refresh the token for retry. Race the refresh against the caller's + // signal: an abort landing while the OAuth operation is pending must settle + // the request with the shared abort contract instead of leaving createMessage + // suspended until the refresh settles and then surfacing the authentication + // failure of a request that no longer exists. + const refreshPromise = openAiCodexOAuthManager.forceRefreshAccessToken() + const refreshed = abortSignal + ? await rejectOnAbort(refreshPromise, abortSignal, this.providerName) + : await refreshPromise if (!refreshed) { throw new Error( t("common:errors.openAiCodex.notAuthenticated", { @@ -456,16 +499,19 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion effectiveSessionId: string, abortSignal?: AbortSignal, ): ApiStream { - // Create AbortController for cancellation - this.abortController = new AbortController() + // Request-local controller: a stale abort arriving after this request finishes must never + // reach the controller of a later request, so the bridge listener captures this controller + // directly instead of reading `this.abortController` at abort time. + const abortController = new AbortController() + this.abortController = abortController // A caller's signal has to be linked rather than used directly, since both transports below - // abort through `this.abortController`. Without this the signal never reaches the wire. - const abortFromCaller = () => this.abortController?.abort() + // abort through the request controller. Without this the signal never reaches the wire. + const abortFromCaller = () => abortController.abort() if (abortSignal) { if (abortSignal.aborted) { - this.abortController.abort() + abortController.abort() } else { abortSignal.addEventListener("abort", abortFromCaller, { once: true }) } @@ -476,7 +522,14 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion // is consistent across providers. try { // Get ChatGPT account ID for organization subscriptions - const accountId = await openAiCodexOAuthManager.getAccountId() + // The account lookup has the same gap as the token fetch: the request-local + // controller is already bridged to the caller's signal, so racing against it + // settles a cancelled request immediately instead of waiting on a dead lookup. + const accountId = await rejectOnAbort( + openAiCodexOAuthManager.getAccountId(), + abortController.signal, + this.providerName, + ) // Build Codex-specific headers. Authorization is provided by the SDK apiKey. const codexHeaders = this.buildCodexHeaders(model, effectiveSessionId, accountId) @@ -492,7 +545,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion }) const stream = (await (client as any).responses.create(requestBody, { - signal: this.abortController.signal, + signal: abortController.signal, // If the SDK supports per-request overrides, ensure headers are present. headers: codexHeaders, })) as AsyncIterable @@ -504,7 +557,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion } for await (const event of stream) { - if (this.abortController.signal.aborted) { + if (abortController.signal.aborted) { break } @@ -513,6 +566,11 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion this.sawSdkEventInCurrentResponse = true for await (const outChunk of this.processEvent(event, model)) { + // The top-of-loop check covers the wire; this one covers the buffered chunks + // processEvent still emits for an event whose processing raced the abort. A break would let + // the event loop pull the stream once more before its top check sees the abort, so hand + // back the abort contract immediately instead. + if (abortController.signal.aborted) throw createAbortError(this.providerName) if (outChunk.type === "text") { this.sawTextOutputInCurrentResponse = true } @@ -527,16 +585,26 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion // A cancellation is not a transport failure either. Falling back would spend a // second request on an already-aborted signal and report the cancellation as a // connection error. - if (this.sawSdkEventInCurrentResponse || this.abortController?.signal.aborted) { + if (this.sawSdkEventInCurrentResponse || abortController.signal.aborted) { throw sdkErr } // Fallback to manual SSE via fetch (Codex backend). - yield* this.makeCodexRequest(requestBody, model, accessToken, effectiveSessionId) + yield* this.makeCodexRequest( + requestBody, + model, + accessToken, + effectiveSessionId, + abortController.signal, + ) } } finally { abortSignal?.removeEventListener("abort", abortFromCaller) - this.abortController = undefined + // Only clear the field if this request still owns it: a concurrent request may have + // installed its own controller after this one started. + if (this.abortController === abortController) { + this.abortController = undefined + } } } @@ -628,12 +696,15 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion model: OpenAiCodexModel, accessToken: string, effectiveSessionId: string, + abortSignal: AbortSignal, ): ApiStream { // Per the implementation guide: route to Codex backend with Bearer token const url = `${CODEX_API_BASE_URL}/responses` // Get ChatGPT account ID for organization subscriptions - const accountId = await openAiCodexOAuthManager.getAccountId() + // The signal here is the request-local controller the caller's abort is bridged into, so a + // cancellation lands in the lookup and settles it before the fallback fetch would start. + const accountId = await rejectOnAbort(openAiCodexOAuthManager.getAccountId(), abortSignal, this.providerName) // Build headers with required Codex-specific fields const headers: Record = { @@ -647,7 +718,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion method: "POST", headers, body: JSON.stringify(requestBody), - signal: this.abortController?.signal, + signal: abortSignal, }) if (!response.ok) { @@ -707,8 +778,15 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion throw new Error(t("common:errors.openAiCodex.noResponseBody")) } - yield* this.handleStreamResponse(response.body, model) + yield* this.handleStreamResponse(response.body, model, abortSignal) } catch (error) { + // The caller cancelled, so this is not a transport fault: hand back the shared abort + // contract instead of reporting the caller's own cancellation to telemetry or wrapping it + // as a connection failure. + if (abortSignal.aborted) { + throw createAbortError(this.providerName) + } + const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "createMessage") TelemetryService.instance.captureException(apiError) @@ -723,7 +801,11 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion } } - private async *handleStreamResponse(body: ReadableStream, model: OpenAiCodexModel): ApiStream { + private async *handleStreamResponse( + body: ReadableStream, + model: OpenAiCodexModel, + abortSignal: AbortSignal, + ): ApiStream { const reader = body.getReader() const decoder = new TextDecoder() let buffer = "" @@ -731,7 +813,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion try { while (true) { - if (this.abortController?.signal.aborted) { + if (abortSignal.aborted) { break } @@ -790,6 +872,9 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion } for await (const outChunk of this.processEvent(parsed, model)) { + // A cancellation landing while processEvent is still emitting must not stream the + // event's remaining buffered chunks, so hand back the abort contract immediately. + if (abortSignal.aborted) throw createAbortError(this.providerName) if (outChunk.type === "text" || outChunk.type === "reasoning") { hasContent = true if (outChunk.type === "text") { @@ -809,6 +894,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion if (content.type === "text" && content.text) { hasContent = true this.sawTextOutputInCurrentResponse = true + if (abortSignal.aborted) break yield { type: "text", text: content.text } } } @@ -817,6 +903,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion for (const summary of outputItem.summary) { if (summary?.type === "summary_text" && typeof summary.text === "string") { hasContent = true + if (abortSignal.aborted) break yield { type: "reasoning", text: summary.text } } } @@ -825,6 +912,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion if (parsed.response.usage) { const usageData = this.normalizeUsage(parsed.response.usage, model) if (usageData) { + if (abortSignal.aborted) break yield usageData } } @@ -951,6 +1039,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion } else if (parsed.choices?.[0]?.delta?.content) { hasContent = true this.sawTextOutputInCurrentResponse = true + if (abortSignal.aborted) break yield { type: "text", text: parsed.choices[0].delta.content } } else if ( parsed.item && @@ -959,10 +1048,12 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion ) { hasContent = true this.sawTextOutputInCurrentResponse = true + if (abortSignal.aborted) break yield { type: "text", text: parsed.item.text } } else if (parsed.usage) { const usageData = this.normalizeUsage(parsed.usage, model) if (usageData) { + if (abortSignal.aborted) break yield usageData } } @@ -977,6 +1068,7 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion if (parsed.content || parsed.text || parsed.message) { hasContent = true this.sawTextOutputInCurrentResponse = true + if (abortSignal.aborted) break yield { type: "text", text: parsed.content || parsed.text || parsed.message } } } catch { @@ -986,6 +1078,13 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion } } } catch (error) { + // The caller cancelled while the fallback stream was being read: that is not a + // stream-processing failure, so hand back the shared abort contract instead of reporting + // the caller's own cancellation to telemetry. + if (abortSignal.aborted) { + throw createAbortError(this.providerName) + } + const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "createMessage") TelemetryService.instance.captureException(apiError) @@ -1329,6 +1428,15 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion * from having to be duplicated here. */ async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { + // Fast-fail if the caller's stop signal already fired before we started: a cancelled + // request must not spend the OAuth setup (token and account fetch) or an SDK request on + // a completion that is already gone. + throwIfAborted(options?.abortSignal) + + // Merge an optional timeout into the caller's abort signal so a timeout cancels the + // completion the same way an external abort does (timeoutMs <= 0 disables it). + const requestSignal = mergeAbortSignalAndTimeout(options?.abortSignal, options?.timeoutMs) + try { const model = this.getModel() @@ -1336,15 +1444,25 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion // the prompt enhancer writes this straight into the input box. let text = "" - for await (const chunk of this.handleResponsesApiMessage( + const stream = this.handleResponsesApiMessage( model, "", [{ role: "user", content: prompt }], // `taskId` is required, and resolves to the same session id this used to send // directly, so `prompt_cache_key` is unchanged. { taskId: this.sessionId }, - options?.abortSignal, - )) { + requestSignal, + ) + + for await (const chunk of stream) { + // A buffered chunk can still be pulled in the window between the abort and the + // inner generator's own stop, so break here: post-abort output must never be + // joined into the completion. + // Stryker disable next-line ConditionalExpression: the false mutant differs from the original only if an abort lands between the transport's last pre-yield check and this consumer's resumption — the yield-hop microtask window no test can schedule deterministically (the abort event is a microtask; the test's own microtasks queue after the chunk's yield); the guard is retained for that production window, the abort regressions pin its pre-abort contract, and the true mutant is killed by the multi-chunk happy paths. + if (requestSignal?.aborted) { + break + } + // Refusals are streamed as text for the chat, but they are not output: the // non-streaming request this replaced read `output_text`, which never carries them. // Keeping them would paste "[Refusal] ..." into the input box as if it were an answer. @@ -1356,20 +1474,18 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion // Both transports end quietly on abort - they break out of their loops rather than // throwing - so returning here would report a cancelled generation as a finished one and // hand the caller whatever partial text had arrived. - if (options?.abortSignal?.aborted) { - throw new DOMException("OpenAI Codex completion was aborted", "AbortError") + if (requestSignal?.aborted) { + throw createAbortError(this.providerName) } return text } catch (error) { // Cancelling is the caller's own doing, not a provider failure, so it is neither - // reported to telemetry nor relabelled as a completion error. A transport that rejects - // on abort reports it in its own words, so it is restated here: callers get one abort - // result whether the stream ended quietly or the request threw. - if (options?.abortSignal?.aborted) { - throw error instanceof DOMException && error.name === "AbortError" - ? error - : new DOMException("OpenAI Codex completion was aborted", "AbortError") + // reported to telemetry nor relabelled as a completion error: a timed-out or + // cancelled request is normalized to the shared abort contract whether the stream + // ended quietly or the request threw. + if (isRequestAborted(error, requestSignal)) { + throw createAbortError(this.providerName) } const errorModel = this.getModel() diff --git a/src/api/providers/utils/__tests__/abort-signal.spec.ts b/src/api/providers/utils/__tests__/abort-signal.spec.ts index 1692f71e63..7ca85d7782 100644 --- a/src/api/providers/utils/__tests__/abort-signal.spec.ts +++ b/src/api/providers/utils/__tests__/abort-signal.spec.ts @@ -3,9 +3,107 @@ import { isRequestAborted, mergeAbortSignalAndTimeout, mergeAbortSignals, + rejectOnAbort, throwIfAborted, } from "../abort-signal" +describe("rejectOnAbort", () => { + it("resolves with the pending value when it settles before the signal aborts", async () => { + const controller = new AbortController() + + await expect(rejectOnAbort(Promise.resolve("done"), controller.signal, "TestProvider")).resolves.toBe("done") + expect(controller.signal.aborted).toBe(false) + }) + + it("rejects with the provider abort error when the signal aborts first", async () => { + const controller = new AbortController() + // Never settles: the race must end purely via the abort. + const pending = new Promise(() => {}) + const race = rejectOnAbort(pending, controller.signal, "TestProvider") + controller.abort() + + // Bound the wait so a broken race fails this test fast (and fails the + // Stryker mutant) instead of hanging. + await expect( + new Promise((_resolve, reject) => { + const deadline = setTimeout(() => reject(new Error("race deadline exceeded")), 3000) + void race.then( + (value) => { + clearTimeout(deadline) + reject(new Error(`race unexpectedly resolved: ${String(value)}`)) + }, + (error) => { + clearTimeout(deadline) + reject(error) + }, + ) + }), + ).rejects.toMatchObject({ + name: "AbortError", + message: "The TestProvider request was aborted", + }) + }) + + it("rejects immediately when the signal is already aborted", async () => { + const controller = new AbortController() + controller.abort() + const pending = new Promise(() => {}) + + await expect(rejectOnAbort(pending, controller.signal, "TestProvider")).rejects.toMatchObject({ + name: "AbortError", + message: "The TestProvider request was aborted", + }) + }) + + it("propagates the pending rejection when the signal stays active", async () => { + const controller = new AbortController() + const boom = new Error("lookup failed") + + await expect(rejectOnAbort(Promise.reject(boom), controller.signal, "TestProvider")).rejects.toBe(boom) + }) + + it("detaches the abort listener once the pending settles", async () => { + const controller = new AbortController() + const addSpy = vi.spyOn(controller.signal, "addEventListener") + const removeSpy = vi.spyOn(controller.signal, "removeEventListener") + + await expect(rejectOnAbort(Promise.resolve("done"), controller.signal, "TestProvider")).resolves.toBe("done") + + // The settle path must remove the exact listener that was registered, not just any + // function: removing a different reference would leave the original abort listener + // attached to the signal. First require the registration to have happened at all, + // so a missing registration cannot silently degrade to an undefined comparison. + expect(addSpy).toHaveBeenCalledTimes(1) + const registeredListener = addSpy.mock.calls[0]?.[1] as EventListener | undefined + expect(typeof registeredListener).toBe("function") + expect(removeSpy).toHaveBeenCalledWith("abort", registeredListener) + addSpy.mockRestore() + removeSpy.mockRestore() + }) + + it("detaches the abort listener when the pending rejects", async () => { + const controller = new AbortController() + const addSpy = vi.spyOn(controller.signal, "addEventListener") + const removeSpy = vi.spyOn(controller.signal, "removeEventListener") + const lookupError = new Error("lookup failed") + + await expect(rejectOnAbort(Promise.reject(lookupError), controller.signal, "TestProvider")).rejects.toBe( + lookupError, + ) + + // The settle path must remove the exact listener that was registered, not just any + // function: removing a different reference would leave the original abort listener + // attached to the signal. First require the registration to have happened at all, + // so a missing registration cannot silently degrade to an undefined comparison. + expect(addSpy).toHaveBeenCalledTimes(1) + const registeredListener = addSpy.mock.calls[0]?.[1] as EventListener | undefined + expect(typeof registeredListener).toBe("function") + expect(removeSpy).toHaveBeenCalledWith("abort", registeredListener) + addSpy.mockRestore() + removeSpy.mockRestore() + }) +}) + describe("abort-signal utilities", () => { describe("mergeAbortSignalAndTimeout", () => { it("returns undefined when no signal or positive timeout is provided", () => { diff --git a/src/api/providers/utils/abort-signal.ts b/src/api/providers/utils/abort-signal.ts index 26f57c3e9a..bd7d579b00 100644 --- a/src/api/providers/utils/abort-signal.ts +++ b/src/api/providers/utils/abort-signal.ts @@ -93,3 +93,35 @@ export function createAbortError(providerName: string): Error { abortError.name = "AbortError" return abortError } + +/** + * Await `pending` but reject with the provider's abort error when `signal` + * aborts first. For async phases that have no native signal support (model + * discovery) yet must still settle promptly on cancellation. The underlying + * promise keeps running (its settlement is ignored) — cancellation is + * cooperative at this boundary. + * + * The abort listener is detached once `pending` settles (success or + * failure), so repeated calls on one signal do not accumulate listeners. + */ +export function rejectOnAbort(pending: Promise, signal: AbortSignal, providerName: string): Promise { + if (signal.aborted) { + return Promise.reject(createAbortError(providerName)) + } + + return new Promise((resolve, reject) => { + const onAbort = () => reject(createAbortError(providerName)) + // Stryker disable next-line ObjectLiteral,BooleanLiteral: a signal fires its abort event exactly once and the settle handler removes this listener, so the once flag is unobservable + signal.addEventListener("abort", onAbort, { once: true }) + void pending.then( + (value) => { + signal.removeEventListener("abort", onAbort) + resolve(value) + }, + (error) => { + signal.removeEventListener("abort", onAbort) + reject(error) + }, + ) + }) +}