diff --git a/src/api/providers/__tests__/openai-compatible.spec.ts b/src/api/providers/__tests__/openai-compatible.spec.ts new file mode 100644 index 0000000000..0ecf28758f --- /dev/null +++ b/src/api/providers/__tests__/openai-compatible.spec.ts @@ -0,0 +1,319 @@ +// npx vitest run api/providers/__tests__/openai-compatible.spec.ts + +import { OpenAICompatibleHandler } from "../openai-compatible" +import { makeApiHandlerOptions } from "../../../test-utils/api" +import { collectStream } from "../../../test-utils/stream" + +const mockGenerateText = vitest.fn() +const mockStreamText = vitest.fn() + +// The factory must not touch the mock bindings at factory-execution time (vi.mock is +// hoisted above the consts), so forward lazily through wrapper functions. +vitest.mock("ai", () => ({ + generateText: (...args: unknown[]) => mockGenerateText(...(args as [])), + streamText: (...args: unknown[]) => mockStreamText(...(args as [])), +})) + +// Concrete test implementation of the abstract OpenAI-compatible base class +class TestOpenAICompatibleHandler extends OpenAICompatibleHandler { + constructor(apiKey: string) { + super(makeApiHandlerOptions({ apiModelId: "test-model" }), { + providerName: "TestProvider", + baseURL: "https://test.example.com/v1", + apiKey, + modelId: "test-model", + modelInfo: { + maxTokens: 4096, + contextWindow: 128000, + supportsImages: false, + supportsPromptCache: false, + inputPrice: 0.5, + outputPrice: 1.5, + }, + }) + } + + override getModel() { + return { id: "test-model", info: this.config.modelInfo } + } +} + +function makeEmptyStreamResult() { + return { + fullStream: { + [Symbol.asyncIterator]: async function* () { + // Emit no parts + yield* [] + }, + }, + usage: Promise.resolve(undefined), + } +} + +describe("OpenAICompatibleHandler", () => { + let handler: TestOpenAICompatibleHandler + + beforeEach(() => { + vi.clearAllMocks() + handler = new TestOpenAICompatibleHandler("test-api-key") + }) + + describe("completePrompt", () => { + it("should return message content from successful response", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const result = await handler.completePrompt("test prompt") + + expect(result).toBe("response") + expect(mockGenerateText).toHaveBeenCalledTimes(1) + expect(mockGenerateText.mock.calls[0][0].prompt).toBe("test prompt") + }) + + it("should pass abortSignal through to generateText", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const controller = new AbortController() + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + + expect(mockGenerateText.mock.calls[0][0].abortSignal).toBe(controller.signal) + }) + + it("should pass timeoutMs through to generateText as a working timeout abort signal", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + await handler.completePrompt("test prompt", { timeoutMs: 50 }) + + const { abortSignal } = mockGenerateText.mock.calls[0][0] + expect(abortSignal).toBeInstanceOf(AbortSignal) + expect(abortSignal.aborted).toBe(false) + + // A never-expiring signal (or a pre-aborted one) would fail this check: + // the signal must fire on its own ~50ms timeout without any external abort. + let guardTimer: ReturnType | undefined + const fired = await Promise.race([ + new Promise((resolve) => { + abortSignal.addEventListener("abort", () => resolve(true), { once: true }) + }), + new Promise((resolve) => { + guardTimer = setTimeout(() => resolve(false), 1000) + }), + ]) + try { + expect(fired).toBe(true) + expect(abortSignal.aborted).toBe(true) + } finally { + // Clear the guard timer so a winning abort event does not leave an + // active timer behind that delays worker teardown. + clearTimeout(guardTimer) + } + }) + + it("should merge signal and timeout when both are provided", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const controller = new AbortController() + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 50 }) + + const mergedSignal = mockGenerateText.mock.calls[0][0].abortSignal as AbortSignal + expect(mergedSignal).toBeInstanceOf(AbortSignal) + expect(mergedSignal.aborted).toBe(false) + + // Aborting the external signal must abort the merged signal synchronously + controller.abort() + expect(mergedSignal.aborted).toBe(true) + }) + + it("should let the timeout component of a merged signal fire without the caller signal", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + await handler.completePrompt("test prompt", { abortSignal: new AbortController().signal, timeoutMs: 50 }) + + const mergedSignal = mockGenerateText.mock.calls[0][0].abortSignal as AbortSignal + expect(mergedSignal).toBeInstanceOf(AbortSignal) + expect(mergedSignal.aborted).toBe(false) + + // With the caller signal left untouched, only the timeout component can fire + let guardTimer: ReturnType | undefined + const fired = await Promise.race([ + new Promise((resolve) => { + mergedSignal.addEventListener("abort", () => resolve(true), { once: true }) + }), + new Promise((resolve) => { + guardTimer = setTimeout(() => resolve(false), 1000) + }), + ]) + try { + expect(fired).toBe(true) + expect(mergedSignal.aborted).toBe(true) + } finally { + // Clear the guard timer so a winning abort event does not leave an + // active timer behind that delays worker teardown. + clearTimeout(guardTimer) + } + }) + + it("should work without options (backward compatible)", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const result = await handler.completePrompt("test prompt") + + expect(result).toBe("response") + expect(mockGenerateText.mock.calls[0][0].abortSignal).toBeUndefined() + // The property must be absent entirely: an unconditional assignment would + // leave it present with an undefined value, which `toBeUndefined` cannot see. + expect("abortSignal" in mockGenerateText.mock.calls[0][0]).toBe(false) + }) + + it("should treat timeoutMs <= 0 as disabled", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + await handler.completePrompt("test prompt", { timeoutMs: 0 }) + expect(mockGenerateText.mock.calls[0][0].abortSignal).toBeUndefined() + + await handler.completePrompt("test prompt", { timeoutMs: -1 }) + expect(mockGenerateText.mock.calls[1][0].abortSignal).toBeUndefined() + }) + + it("should pass the caller signal unchanged when timeoutMs is 0", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const controller = new AbortController() + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 0 }) + + // A disabled timeout must not drop the caller's cancellation signal + expect(mockGenerateText.mock.calls[0][0].abortSignal).toBe(controller.signal) + }) + + it("should pass the caller signal unchanged when timeoutMs is negative", async () => { + mockGenerateText.mockResolvedValue({ text: "response" }) + + const controller = new AbortController() + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: -1 }) + + expect(mockGenerateText.mock.calls[0][0].abortSignal).toBe(controller.signal) + }) + + it("should reject with the real DOMException AbortError when the signal is pre-aborted", async () => { + // Emulate the real AI SDK: a pre-aborted signal makes the request reject + // with the fetch stack's DOMException abort error rather than a fabricated + // Error, so this exercises the provider's pass-through of a real SDK + // abort. openai-compatible.ts has no normalization layer, so the + // DOMException must surface unchanged (name and message). + mockGenerateText.mockImplementation((options: { abortSignal?: AbortSignal }) => { + if (options.abortSignal?.aborted) { + return Promise.reject(new DOMException("The operation was aborted.", "AbortError")) + } + return Promise.resolve({ text: "response" }) + }) + + const controller = new AbortController() + controller.abort() + + await expect( + handler.completePrompt("test prompt", { abortSignal: controller.signal }), + ).rejects.toMatchObject({ + name: "AbortError", + message: "The operation was aborted.", + }) + }) + + it("should throw handled error when API call fails", async () => { + mockGenerateText.mockRejectedValue(new Error("Network error")) + + await expect(handler.completePrompt("test prompt")).rejects.toThrow("Network error") + }) + }) + + describe("createMessage", () => { + it("should pass the external abortSignal to streamText", async () => { + mockStreamText.mockReturnValue(makeEmptyStreamResult()) + + const controller = new AbortController() + const stream = handler.createMessage("You are helpful.", [], { + taskId: "test", + abortSignal: controller.signal, + }) + await collectStream(stream) + + expect(mockStreamText).toHaveBeenCalledTimes(1) + expect(mockStreamText.mock.calls[0][0].abortSignal).toBe(controller.signal) + }) + + it("should not set an abortSignal when metadata has none", async () => { + mockStreamText.mockReturnValue(makeEmptyStreamResult()) + + const stream = handler.createMessage("You are helpful.", [], { taskId: "test" }) + await collectStream(stream) + + expect(mockStreamText.mock.calls[0][0].abortSignal).toBeUndefined() + // The property must be absent entirely (an unconditional assignment would + // leave it present with an undefined value). + expect("abortSignal" in mockStreamText.mock.calls[0][0]).toBe(false) + }) + + it("should complete without metadata and leave the abortSignal property unset", async () => { + mockStreamText.mockReturnValue(makeEmptyStreamResult()) + + const stream = handler.createMessage("You are helpful.", []) + await collectStream(stream) + + expect(mockStreamText).toHaveBeenCalledTimes(1) + // createMessage without metadata must not throw and must leave the + // property absent: metadata?.abortSignal is undefined when metadata is absent. + expect("abortSignal" in mockStreamText.mock.calls[0][0]).toBe(false) + }) + + it("should reject with AbortError when the external abortSignal is pre-aborted", async () => { + mockStreamText.mockImplementation((options: { abortSignal?: AbortSignal }) => ({ + fullStream: { + [Symbol.asyncIterator]: async function* () { + if (options.abortSignal?.aborted) { + const error = new Error("This operation was aborted") + error.name = "AbortError" + throw error + } + yield* [] + }, + }, + usage: Promise.resolve(undefined), + })) + + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage("You are helpful.", [], { + taskId: "test", + abortSignal: controller.signal, + }) + await expect(collectStream(stream)).rejects.toMatchObject({ name: "AbortError" }) + }) + + it("should abort the stream when the external abortSignal is aborted mid-request", async () => { + mockStreamText.mockImplementation((options: { abortSignal?: AbortSignal }) => ({ + fullStream: { + [Symbol.asyncIterator]: async function* () { + // Emulate a slow model response that ends when the request is aborted + await new Promise((resolve) => { + options.abortSignal?.addEventListener("abort", () => resolve(), { once: true }) + }) + yield* [] + const error = new Error("This operation was aborted") + error.name = "AbortError" + throw error + }, + }, + usage: Promise.resolve(undefined), + })) + + const controller = new AbortController() + const stream = handler.createMessage("You are helpful.", [], { + taskId: "test", + abortSignal: controller.signal, + }) + const collected = collectStream(stream) + setTimeout(() => controller.abort(), 10) + + await expect(collected).rejects.toMatchObject({ name: "AbortError" }) + }) + }) +}) diff --git a/src/api/providers/__tests__/openai-native.spec.ts b/src/api/providers/__tests__/openai-native.spec.ts index 8d5975c0d9..dd5d044070 100644 --- a/src/api/providers/__tests__/openai-native.spec.ts +++ b/src/api/providers/__tests__/openai-native.spec.ts @@ -16,9 +16,14 @@ import OpenAI from "openai" import { ApiProviderError, OpenAiServiceTier, SERVICE_TIER_KEY, serviceTiers } from "@roo-code/types" import { OpenAiNativeHandler } from "../openai-native" +import type { ApiStreamChunk, ApiStreamTextChunk } from "../../../api/transform/stream" import { ApiHandlerOptions } from "../../../shared/api" import { Package } from "../../../shared/package" -import { expectRequestObjectContaining, makeApiHandlerOptions } from "../../../test-utils/api" +import { + expectRequestObjectContaining, + makeApiHandlerOptions, + makeCreateMessageMetadata, +} from "../../../test-utils/api" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" import { deleteGlobalFetch } from "../../../test-utils/reset" @@ -132,6 +137,35 @@ describe("OpenAiNativeHandler", () => { }) describe("createMessage", () => { + function makeOpenStreamFetchMock() { + type OpenStream = { + controller?: ReadableStreamDefaultController + fetchSignal: AbortSignal + } + const openStreams: OpenStream[] = [] + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + const entry: OpenStream = { fetchSignal: options?.signal as AbortSignal } + const body = new ReadableStream({ + start: (controller) => { + entry.controller = controller + }, + }) + openStreams.push(entry) + return Promise.resolve({ + ok: true, + body, + }) + }) + const requireController = (index: number): ReadableStreamDefaultController => { + const entry = openStreams[index] + if (!entry?.controller) { + throw new Error("expected fallback fetch to have started") + } + return entry.controller + } + return { openStreams, mockFetch, requireController } + } + it("shapes GPT-6 Astra requests for Responses tool calling", () => { const astraHandler = new OpenAiNativeHandler({ ...mockOptions, @@ -426,6 +460,753 @@ describe("OpenAiNativeHandler", () => { } }).rejects.toThrow("OpenAI service error") }) + + it("should reject with AbortError when the external abortSignal is already aborted (fallback path)", async () => { + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + if (options?.signal?.aborted) { + const error = new Error("This operation was aborted") + error.name = "AbortError" + return Promise.reject(error) + } + return new Promise(() => {}) + }) + global.fetch = mockFetch as typeof fetch + + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + await expect(collectStream(stream)).rejects.toMatchObject({ name: "AbortError" }) + }) + + it("should abort the fallback fetch when the external abortSignal is aborted mid-request", async () => { + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + return new Promise((_resolve, reject) => { + options?.signal?.addEventListener( + "abort", + () => { + const error = new Error("This operation was aborted") + error.name = "AbortError" + reject(error) + }, + { once: true }, + ) + }) + }) + global.fetch = mockFetch as typeof fetch + + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + const controller = new AbortController() + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const collected = collectStream(stream) + setTimeout(() => controller.abort(), 10) + + await expect(collected).rejects.toMatchObject({ name: "AbortError" }) + }) + + it("should not let a late abort from an earlier request cancel a later request", async () => { + // Regression: the external-signal bridge must detach on request completion. + // With a lingering listener (or one reading the mutable this.abortController + // field), aborting the FIRST request's signal after completion would cancel + // the SECOND request's controller. + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + const { openStreams, mockFetch, requireController } = makeOpenStreamFetchMock() + global.fetch = mockFetch as typeof fetch + + const firstController = new AbortController() + const secondController = new AbortController() + + // First request: completes normally. + const firstStream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: firstController.signal }), + ) + const firstCollected = collectStream(firstStream) + await new Promise((resolve) => setTimeout(resolve, 10)) + requireController(0).enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"one"}\n\n'), + ) + requireController(0).enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + requireController(0).close() + + const firstChunks = await firstCollected + expect(firstChunks.some((chunk) => chunk.type === "text" && chunk.text === "one")).toBe(true) + + // Second request with a different external signal, left in-flight. + const secondStream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: secondController.signal }), + ) + const secondCollected = collectStream(secondStream) + await new Promise((resolve) => setTimeout(resolve, 10)) + expect(openStreams).toHaveLength(2) + + // Aborting the FIRST request's signal must not leak into the second request. + firstController.abort() + + // The second request's internal fetch signal must remain active... + expect(openStreams[1].fetchSignal.aborted).toBe(false) + + // ...and the second stream must still complete normally. + requireController(1).enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"two"}\n\n'), + ) + requireController(1).enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + requireController(1).close() + + const secondChunks = await secondCollected + expect(secondChunks.some((chunk) => chunk.type === "text" && chunk.text === "two")).toBe(true) + }) + + describe("abort-signal bridging", () => { + // The bedrock-pattern bridge in executeRequest and makeResponsesApiRequest forwards + // metadata?.abortSignal onto a request-local AbortController. These tests make every + // branch observable: a resolving SDK mock exercises the executeRequest bridge + // directly, a rejecting one exercises the fetch fallback bridge, and the + // request-local signal handed to the SDK/fetch is captured for assertions. + + function makeAbortError(): Error { + const error = new Error("This operation was aborted") + error.name = "AbortError" + return error + } + + const tick = (): Promise => new Promise((resolve) => setTimeout(resolve, 0)) + + function untilSignalAborted(signal: AbortSignal, timeoutMs = 200): Promise { + return new Promise((resolve) => { + if (signal.aborted) { + resolve() + return + } + const timer = setTimeout(() => resolve(), timeoutMs) + signal.addEventListener( + "abort", + () => { + clearTimeout(timer) + resolve() + }, + { once: true }, + ) + }) + } + + function textChunks(chunks: ApiStreamChunk[]): ApiStreamTextChunk[] { + return chunks.filter((chunk): chunk is ApiStreamTextChunk => chunk.type === "text") + } + + it("should register a once-only abort listener on the external signal and detach it when the SDK request completes", async () => { + const controller = new AbortController() + const addSpy = vi.spyOn(controller.signal, "addEventListener") + const removeSpy = vi.spyOn(controller.signal, "removeEventListener") + let sdkSignal: AbortSignal | undefined + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkSignal = options?.signal + return Promise.resolve( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "one" }, + { type: "response.output_text.delta", delta: " two" }, + ]), + ) + }) + + try { + const chunks = await collectStream( + handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ), + ) + + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one", " two"]) + expect(sdkSignal?.aborted).toBe(false) + // The bridge must listen for the "abort" event with { once: true } ... + expect(addSpy).toHaveBeenCalledTimes(1) + expect(addSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) + // ... and detach that exact listener when the request completes. + const registeredListener = addSpy.mock.calls.find(([type]) => type === "abort")?.[1] + expect(registeredListener).toBeDefined() + expect(removeSpy).toHaveBeenCalledTimes(1) + expect(removeSpy).toHaveBeenCalledWith("abort", registeredListener) + // The request-local controller is cleared once the request is done. + expect(handler["abortController"]).toBeUndefined() + } finally { + addSpy.mockRestore() + removeSpy.mockRestore() + } + }) + + it("should abort the SDK request immediately when the external signal is already aborted", async () => { + const controller = new AbortController() + controller.abort() + const addSpy = vi.spyOn(controller.signal, "addEventListener") + let sdkSignal: AbortSignal | undefined + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkSignal = options?.signal + // A real SDK rejects immediately when its request signal is pre-aborted. + if (options?.signal?.aborted) { + return Promise.reject(makeAbortError()) + } + return Promise.resolve(asyncStreamFrom([{ type: "response.output_text.delta", delta: "one" }])) + }) + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + if (options?.signal?.aborted) { + return Promise.reject(makeAbortError()) + } + return new Promise(() => {}) + }) + global.fetch = mockFetch as typeof fetch + + try { + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + await expect(collectStream(stream)).rejects.toMatchObject({ name: "AbortError" }) + + // The bridge must have pre-aborted the request-local controller ... + expect(sdkSignal?.aborted).toBe(true) + // ... instead of registering a listener on the already-aborted signal. + expect(addSpy).not.toHaveBeenCalled() + } finally { + addSpy.mockRestore() + } + }) + + it("should not pre-abort the SDK request for a pending external signal and should abort it mid-flight", async () => { + const controller = new AbortController() + let sdkSignal: AbortSignal | undefined + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkSignal = options?.signal + const signal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "one" } + if (signal) { + await untilSignalAborted(signal, 200) + } + })(), + ) + }) + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collected = collectStream(stream) + await tick() + + expect(sdkSignal).toBeDefined() + // A pending external signal must not abort the request up front. + expect(sdkSignal?.aborted).toBe(false) + + controller.abort() + await tick() + // ... but it must abort the request as soon as it fires. + expect(sdkSignal?.aborted).toBe(true) + + const chunks = await collected + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one"]) + }) + + it("should stop consuming the SDK stream once the external signal aborts the request", async () => { + const controller = new AbortController() + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + const signal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "first" } + if (signal) { + await untilSignalAborted(signal, 200) + } + yield { type: "response.output_text.delta", delta: "second" } + })(), + ) + }) + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collected = collectStream(stream) + await tick() + + controller.abort() + + const chunks = await collected + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["first"]) + }) + + it("should detach the external abort listener on completion so a late abort cannot abort the request signal", async () => { + const controller = new AbortController() + let openGate: (() => void) | undefined + const gate = new Promise((resolve) => { + openGate = resolve + }) + let sdkSignal: AbortSignal | undefined + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkSignal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "one" } + await gate + })(), + ) + }) + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collected = collectStream(stream) + await tick() + + // Let the request complete normally, then abort the external signal late. + if (!openGate) { + throw new Error("expected the stream gate to be ready") + } + openGate() + const chunks = await collected + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one"]) + + controller.abort() + await tick() + + // The bridging listener must have been detached: the late abort must + // not reach the already-completed request's controller. + expect(sdkSignal?.aborted).toBe(false) + expect(handler["abortController"]).toBeUndefined() + }) + + it("should not call removeEventListener on the external signal when no listener was registered", async () => { + const controller = new AbortController() + controller.abort() + const removeSpy = vi.spyOn(controller.signal, "removeEventListener") + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + if (options?.signal?.aborted) { + return Promise.reject(makeAbortError()) + } + return Promise.resolve(asyncStreamFrom([{ type: "response.output_text.delta", delta: "one" }])) + }) + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + if (options?.signal?.aborted) { + return Promise.reject(makeAbortError()) + } + return new Promise(() => {}) + }) + global.fetch = mockFetch as typeof fetch + + try { + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + await expect(collectStream(stream)).rejects.toMatchObject({ name: "AbortError" }) + + // A pre-aborted signal registers no listener, so nothing may be removed. + expect(removeSpy).not.toHaveBeenCalled() + } finally { + removeSpy.mockRestore() + } + }) + + it("should preserve a later fallback request's controller when an earlier SDK request completes", async () => { + // Request A: SDK path, in flight. Request B: SDK fails, so its fallback + // fetch installs the handler's controller. When A completes, its finally + // must not clear the controller owned by B's fallback. + let aGateOpen: (() => void) | undefined + const aGate = new Promise((resolve) => { + aGateOpen = resolve + }) + let aSdkSignal: AbortSignal | undefined + let sdkCalls = 0 + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkCalls += 1 + if (sdkCalls === 1) { + aSdkSignal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "a-one" } + await aGate + })(), + ) + } + return Promise.reject(new Error("SDK not available")) + }) + const { openStreams, mockFetch, requireController } = makeOpenStreamFetchMock() + global.fetch = mockFetch as typeof fetch + + const streamA = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ) + const collectedA = collectStream(streamA) + await tick() + expect(aSdkSignal?.aborted).toBe(false) + + const streamB = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ) + const collectedB = collectStream(streamB) + await tick() + expect(openStreams).toHaveLength(1) + + // Complete A while B's fallback owns the handler's controller. + if (!aGateOpen) { + throw new Error("expected the stream gate to be ready") + } + aGateOpen() + const chunksA = await collectedA + expect(textChunks(chunksA).map((chunk) => chunk.text)).toEqual(["a-one"]) + expect(handler["abortController"]?.signal).toBe(openStreams[0].fetchSignal) + + // Let B finish; its finally chain clears the controller. + requireController(0).enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"b-one"}\n\n'), + ) + requireController(0).enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + requireController(0).close() + const chunksB = await collectedB + expect(textChunks(chunksB).map((chunk) => chunk.text)).toEqual(["b-one"]) + expect(handler["abortController"]).toBeUndefined() + }) + + it("should not clear a later SDK request's controller when a fallback request completes", async () => { + // Mirror of the previous test: request B (fallback) starts first and + // request A (SDK) takes over the handler's controller. When B's fallback + // completes, its finally must not clear A's controller. + let aGateOpen: (() => void) | undefined + const aGate = new Promise((resolve) => { + aGateOpen = resolve + }) + let aSdkSignal: AbortSignal | undefined + let sdkCalls = 0 + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkCalls += 1 + if (sdkCalls === 1) { + return Promise.reject(new Error("SDK not available")) + } + aSdkSignal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "a-one" } + await aGate + })(), + ) + }) + const { openStreams, mockFetch, requireController } = makeOpenStreamFetchMock() + global.fetch = mockFetch as typeof fetch + + const streamB = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ) + const collectedB = collectStream(streamB) + await tick() + expect(openStreams).toHaveLength(1) + + const streamA = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ) + const collectedA = collectStream(streamA) + await tick() + expect(aSdkSignal?.aborted).toBe(false) + + // Let B's fallback complete while A owns the handler's controller. + requireController(0).enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"b-one"}\n\n'), + ) + requireController(0).enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + requireController(0).close() + const chunksB = await collectedB + expect(textChunks(chunksB).map((chunk) => chunk.text)).toEqual(["b-one"]) + expect(handler["abortController"]?.signal).toBe(aSdkSignal) + + // Let A finish; its finally clears the controller. + if (!aGateOpen) { + throw new Error("expected the stream gate to be ready") + } + aGateOpen() + const chunksA = await collectedA + expect(textChunks(chunksA).map((chunk) => chunk.text)).toEqual(["a-one"]) + expect(handler["abortController"]).toBeUndefined() + }) + + it("should clear the handler's abortController after a fallback request completes", async () => { + const mockFetch = vitest.fn().mockResolvedValue({ + ok: true, + body: new ReadableStream({ + start(controller) { + controller.enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"one"}\n\n'), + ) + controller.enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + controller.close() + }, + }), + }) + global.fetch = mockFetch as typeof fetch + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + const chunks = await collectStream(handler.createMessage(systemPrompt, messages)) + + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one"]) + // The fallback installs its own controller and must clear it when done. + expect(handler["abortController"]).toBeUndefined() + }) + + it("should register a once-only abort listener in the fallback path and detach it on completion", async () => { + const controller = new AbortController() + const addSpy = vi.spyOn(controller.signal, "addEventListener") + const removeSpy = vi.spyOn(controller.signal, "removeEventListener") + const mockFetch = vitest.fn().mockResolvedValue({ + ok: true, + body: new ReadableStream({ + start(controller) { + controller.enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"one"}\n\n'), + ) + controller.enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + controller.close() + }, + }), + }) + global.fetch = mockFetch as typeof fetch + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + try { + const chunks = await collectStream( + handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ), + ) + + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one"]) + // Both bridges (SDK path and fallback) listen for "abort" with + // { once: true }, and both detach their own listener on completion. + expect(addSpy).toHaveBeenCalledTimes(2) + expect(removeSpy).toHaveBeenCalledTimes(2) + for (const call of addSpy.mock.calls) { + expect(call[0]).toBe("abort") + expect(call[2]).toEqual({ once: true }) + } + const registered = addSpy.mock.calls.map(([, listener]) => listener) + const removed = removeSpy.mock.calls.map(([type, listener]) => { + expect(type).toBe("abort") + return listener + }) + // Each bridge must detach its own listener exactly once. + expect(new Set(removed).size).toBe(2) + expect(new Set(removed)).toEqual(new Set(registered)) + expect(handler["abortController"]).toBeUndefined() + } finally { + addSpy.mockRestore() + removeSpy.mockRestore() + } + }) + + it("should detach the fallback's external abort listener on completion so a late abort cannot abort the fetch signal", async () => { + const { openStreams, mockFetch, requireController } = makeOpenStreamFetchMock() + global.fetch = mockFetch as typeof fetch + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + + const controller = new AbortController() + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collected = collectStream(stream) + await tick() + expect(openStreams).toHaveLength(1) + + requireController(0).enqueue( + new TextEncoder().encode('data: {"type":"response.text.delta","delta":"one"}\n\n'), + ) + requireController(0).enqueue(new TextEncoder().encode("data: [DONE]\n\n")) + requireController(0).close() + + const chunks = await collected + expect(textChunks(chunks).map((chunk) => chunk.text)).toEqual(["one"]) + + // A late abort must not reach this request's own fetch signal. + controller.abort() + await tick() + expect(openStreams[0].fetchSignal.aborted).toBe(false) + }) + + it("should surface the contract AbortError once when the fallback stream read rejects on external abort", async () => { + // Force the fallback path and emulate undici: the body's reader.read() + // stays pending and rejects with a DOMException AbortError when the + // request signal aborts. Before the fix, handleStreamResponse wrapped + // that error in a plain Error (defeating the caller's AbortError guard) + // and the request was captured as an exception twice. + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + let fetchSignal: AbortSignal | undefined + const mockFetch = vitest.fn().mockImplementation((_url: string, options?: RequestInit) => { + const signal = options?.signal + if (!signal) { + return Promise.reject(new Error("expected the fallback fetch to carry a request signal")) + } + fetchSignal = signal + const body = new ReadableStream({ + pull: () => { + return new Promise((_resolve, reject) => { + if (signal.aborted) { + reject(new DOMException("This operation was aborted", "AbortError")) + return + } + signal.addEventListener( + "abort", + () => reject(new DOMException("This operation was aborted", "AbortError")), + { once: true }, + ) + }) + }, + }) + return Promise.resolve({ ok: true, body }) + }) + global.fetch = mockFetch as typeof fetch + + const controller = new AbortController() + const collected = collectStream( + handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ), + ) + await tick() + expect(fetchSignal).toBeDefined() + + // A user stop mid-stream must surface exactly one error, contract-named. + controller.abort() + await expect(collected).rejects.toMatchObject({ + name: "AbortError", + message: "The OpenAI Native request was aborted", + }) + // The provider must not report a user-triggered stop as an exception. + expect(mockCaptureException).not.toHaveBeenCalled() + }) + + it("should not convert a non-abort stream error into an AbortError", async () => { + mockResponsesCreate.mockRejectedValue(new Error("SDK not available")) + const mockFetch = vitest.fn().mockImplementation(() => { + const body = new ReadableStream({ + pull: () => Promise.reject(new Error("socket hang up")), + }) + return Promise.resolve({ ok: true, body }) + }) + global.fetch = mockFetch as typeof fetch + + await expect(collectStream(handler.createMessage(systemPrompt, messages))).rejects.toThrow( + "Error processing response stream: socket hang up", + ) + }) + + it("should not let an earlier request's external abort affect a later request on the same handler", async () => { + // Request 1 streams under external signal A while request 2 starts under + // external signal B. Aborting A while both are in flight must abort only + // request 1's request-local controller; request 2 completes normally. + const firstExternal = new AbortController() + const secondExternal = new AbortController() + let firstGateOpen: (() => void) | undefined + const firstGate = new Promise((resolve) => { + firstGateOpen = resolve + }) + let firstSdkSignal: AbortSignal | undefined + let secondSdkSignal: AbortSignal | undefined + let sdkCalls = 0 + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + sdkCalls += 1 + if (sdkCalls === 1) { + firstSdkSignal = options?.signal + return Promise.resolve( + (async function* () { + yield { type: "response.output_text.delta", delta: "one" } + await firstGate + })(), + ) + } + secondSdkSignal = options?.signal + return Promise.resolve( + asyncStreamFrom([ + { type: "response.output_text.delta", delta: "two-a" }, + { type: "response.output_text.delta", delta: " two-b" }, + ]), + ) + }) + + const collected1 = collectStream( + handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: firstExternal.signal }), + ), + ) + await tick() + expect(firstSdkSignal).toBeDefined() + + const collected2 = collectStream( + handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: secondExternal.signal }), + ), + ) + await tick() + expect(secondSdkSignal).toBeDefined() + + // Abort the first request's external signal while the second is in flight. + firstExternal.abort() + await tick() + expect(firstSdkSignal?.aborted).toBe(true) + expect(secondSdkSignal?.aborted).toBe(false) + + // The second request completes normally with its own content. + const chunks2 = await collected2 + expect(textChunks(chunks2).map((chunk) => chunk.text)).toEqual(["two-a", " two-b"]) + expect(secondExternal.signal.aborted).toBe(false) + + // Let the first request wind down; its stream simply ends. + if (!firstGateOpen) { + throw new Error("expected the first stream gate to be ready") + } + firstGateOpen() + const chunks1 = await collected1 + expect(textChunks(chunks1).map((chunk) => chunk.text)).toEqual(["one"]) + }) + }) }) describe("completePrompt", () => { @@ -534,6 +1315,180 @@ describe("OpenAiNativeHandler", () => { expect(result).toBe("") }) + it("should pass the external abort signal through to the SDK request", async () => { + mockResponsesCreate.mockResolvedValue({ + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + }) + + const controller = new AbortController() + await handler.completePrompt("Test prompt", { abortSignal: controller.signal }) + + // Without a timeout the merged signal is the external signal itself + expect(mockResponsesCreate.mock.calls[0][1].signal).toBe(controller.signal) + }) + + it("should work without options (backward compatible)", async () => { + mockResponsesCreate.mockResolvedValue({ + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + }) + + const result = await handler.completePrompt("Test prompt") + + expect(result).toBe("response") + expect(mockResponsesCreate.mock.calls[0][1].signal).toBeInstanceOf(AbortSignal) + }) + + it("completePrompt should abort its request signal when timeoutMs is reached", async () => { + // Node's AbortSignal.timeout() uses internal timers that vi.useFakeTimers() does not + // intercept, so this relies on a short real timeout instead of fake timers. + let requestSignal: AbortSignal | undefined + mockResponsesCreate.mockImplementationOnce(async (_body: unknown, options: { signal?: AbortSignal }) => { + requestSignal = options.signal + // Stay pending until the merged timeout signal aborts the request + await new Promise((resolve) => { + options.signal?.addEventListener("abort", () => resolve(), { once: true }) + }) + return { + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + } + }) + + const result = await handler.completePrompt("Test prompt", { timeoutMs: 50 }) + + expect(result).toBe("response") + expect(requestSignal).toBeInstanceOf(AbortSignal) + expect(requestSignal?.aborted).toBe(true) + }) + + it("completePrompt should merge the external signal and timeoutMs together", async () => { + const controller = new AbortController() + mockResponsesCreate.mockResolvedValue({ + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + }) + + await handler.completePrompt("Test prompt", { abortSignal: controller.signal, timeoutMs: 10000 }) + + const mergedSignal = mockResponsesCreate.mock.calls[0][1].signal as AbortSignal + expect(mergedSignal).toBeInstanceOf(AbortSignal) + + // Aborting the external signal must abort the merged signal synchronously + controller.abort() + expect(mergedSignal.aborted).toBe(true) + }) + + it("completePrompt should reject with AbortError when the abortSignal is already aborted", async () => { + mockResponsesCreate.mockImplementation((_body: unknown, options: { signal?: AbortSignal }) => { + if (options?.signal?.aborted) { + const error = new Error("This operation was aborted") + error.name = "AbortError" + return Promise.reject(error) + } + return Promise.resolve({ + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + }) + }) + + const controller = new AbortController() + controller.abort() + + await expect( + handler.completePrompt("Test prompt", { abortSignal: controller.signal }), + ).rejects.toMatchObject({ + name: "AbortError", + }) + }) + + it("completePrompt should rethrow non-Error failures after telemetry", async () => { + mockResponsesCreate.mockRejectedValue("string failure") + + await expect(handler.completePrompt("Test prompt")).rejects.toBe("string failure") + expect(mockCaptureException).toHaveBeenCalledWith( + expect.objectContaining({ + message: "string failure", + provider: "OpenAI Native", + modelId: "gpt-4.1", + operation: "completePrompt", + }), + ) + }) + + it("completePrompt should return direct response text fallback", async () => { + mockResponsesCreate.mockResolvedValue({ text: "fallback response" }) + + const result = await handler.completePrompt("Test prompt") + + expect(result).toBe("fallback response") + }) + + it("completePrompt should include supported service tier, reasoning, verbosity, and prompt cache retention", async () => { + const configuredHandler = new OpenAiNativeHandler({ + ...mockOptions, + apiModelId: "gpt-5.1", + openAiNativeServiceTier: "flex", + enableResponsesReasoningSummary: true, + }) + mockResponsesCreate.mockResolvedValue({ + output: [ + { + type: "message", + content: [{ type: "output_text", text: "response" }], + }, + ], + }) + + await configuredHandler.completePrompt("Test prompt") + + const requestBody = mockResponsesCreate.mock.calls[0][0] + expect(requestBody.service_tier).toBe("flex") + expect(requestBody.include).toEqual(["reasoning.encrypted_content"]) + expect(requestBody.reasoning).toEqual({ effort: "medium", summary: "auto" }) + expect(requestBody.text).toEqual({ verbosity: "medium" }) + expect(requestBody.prompt_cache_retention).toBe("24h") + }) + + it("should expose response id and encrypted reasoning content", () => { + handler["lastResponseId"] = "resp_123" + handler["lastResponseOutput"] = [ + { type: "message" }, + { type: "reasoning", encrypted_content: "encrypted", id: "reasoning_1" }, + ] + + expect(handler.getResponseId()).toBe("resp_123") + expect(handler.getEncryptedContent()).toEqual({ encrypted_content: "encrypted", id: "reasoning_1" }) + }) + + it("should return undefined when encrypted reasoning content is absent", () => { + expect(handler.getEncryptedContent()).toBeUndefined() + + handler["lastResponseOutput"] = [{ type: "reasoning" }] + + expect(handler.getEncryptedContent()).toBeUndefined() + }) }) describe("getModel", () => { diff --git a/src/api/providers/__tests__/request-config-builder.spec.ts b/src/api/providers/__tests__/request-config-builder.spec.ts index 977b09df6b..4ec2d3abc4 100644 --- a/src/api/providers/__tests__/request-config-builder.spec.ts +++ b/src/api/providers/__tests__/request-config-builder.spec.ts @@ -505,4 +505,20 @@ describe("RequestConfigBuilder", () => { expect(config.maxTokens).toBe(2000) }) }) + + describe("static merge helpers (canonical abort-signal entry points)", () => { + it("returns undefined from mergeAbortSignalAndTimeout when no external signal and no valid timeout", () => { + expect(RequestConfigBuilder.mergeAbortSignalAndTimeout(undefined, undefined)).toBeUndefined() + expect(RequestConfigBuilder.mergeAbortSignalAndTimeout(undefined, 0)).toBeUndefined() + expect(RequestConfigBuilder.mergeAbortSignalAndTimeout(undefined, -5)).toBeUndefined() + }) + + it("returns the external signal directly when no timeout is merged", () => { + const controller = new AbortController() + expect(RequestConfigBuilder.mergeAbortSignalAndTimeout(controller.signal, undefined)).toBe( + controller.signal, + ) + expect(RequestConfigBuilder.mergeAbortSignalAndTimeout(controller.signal, 0)).toBe(controller.signal) + }) + }) }) diff --git a/src/api/providers/config-builder/request-config-builder.ts b/src/api/providers/config-builder/request-config-builder.ts index 2201d735bc..38a6c36b08 100644 --- a/src/api/providers/config-builder/request-config-builder.ts +++ b/src/api/providers/config-builder/request-config-builder.ts @@ -163,4 +163,15 @@ export class RequestConfigBuilder { const languageModel = this.getLanguageModel() - const { text } = await generateText({ + const generateOptions: Parameters[0] & { abortSignal?: AbortSignal } = { model: languageModel, prompt, maxOutputTokens: this.getMaxOutputTokens(), temperature: this.config.temperature ?? 0, - }) + } + + const mergedAbortSignal = RequestConfigBuilder.mergeAbortSignalAndTimeout( + options?.abortSignal, + options?.timeoutMs, + ) + if (mergedAbortSignal) { + generateOptions.abortSignal = mergedAbortSignal + } + + const { text } = await generateText(generateOptions) return text } diff --git a/src/api/providers/openai-native.ts b/src/api/providers/openai-native.ts index e8d23a0c68..0ab0c644c7 100644 --- a/src/api/providers/openai-native.ts +++ b/src/api/providers/openai-native.ts @@ -32,6 +32,7 @@ import { NOT_PROVIDED } from "./constants" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" import { isMcpTool } from "../../utils/mcp-name" import { sanitizeOpenAiCallId } from "../../utils/tool-id" +import { RequestConfigBuilder } from "./config-builder/request-config-builder" export type OpenAiNativeModel = ReturnType @@ -409,6 +410,32 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio return body } + /** + * Bridges an external abort signal (e.g. task cancellation) onto a request-local + * controller using the Bedrock pattern: a pre-aborted signal aborts the controller + * immediately and registers no listener; otherwise a once-only listener forwards + * the abort. Returns a cleanup function that detaches the listener (or undefined + * when no listener was registered) so the caller can release it in its finally + * block and a late abort can never reach a later request's controller. + */ + private attachExternalAbort( + externalAbortSignal: AbortSignal | undefined, + requestController: AbortController, + ): (() => void) | undefined { + if (!externalAbortSignal) { + return undefined + } + if (externalAbortSignal.aborted) { + requestController.abort() + return undefined + } + const abortListener = () => requestController.abort() + externalAbortSignal.addEventListener("abort", abortListener, { once: true }) + return () => { + externalAbortSignal.removeEventListener("abort", abortListener) + } + } + private async *executeRequest( requestBody: any, model: OpenAiNativeModel, @@ -416,8 +443,17 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio systemPrompt?: string, messages?: Anthropic.Messages.MessageParam[], ): ApiStream { - // Create AbortController for cancellation - this.abortController = new AbortController() + // Create a request-local AbortController for cancellation. It is exposed via + // this.abortController so the stop-button path can observe it, but all bridging + // below captures the local reference so a late abort from an earlier request can + // never reach a later request's controller. + const requestController = new AbortController() + this.abortController = requestController + + // Bridge the external abort signal onto the request controller; the returned + // cleanup detaches the listener in the finally block so a late abort from this + // request cannot cancel a later request's controller. + const detach = this.attachExternalAbort(metadata?.abortSignal, requestController) // Build per-request headers using taskId when available, falling back to sessionId const taskId = metadata?.taskId @@ -431,7 +467,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio try { // Use the official SDK with per-request headers const stream = (await (this.client as any).responses.create(requestBody, { - signal: this.abortController.signal, + signal: requestController.signal, headers: requestHeaders, })) as AsyncIterable @@ -443,7 +479,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio for await (const event of stream) { // Check if request was aborted - if (this.abortController.signal.aborted) { + if (requestController.signal.aborted) { break } @@ -455,7 +491,14 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio // For errors, fallback to manual SSE via fetch yield* this.makeResponsesApiRequest(requestBody, model, metadata, systemPrompt, messages) } finally { - this.abortController = undefined + // Detach the bridging listener so a late abort from this request cannot + // cancel a later request's controller. + detach?.() + // Only clear the field if this request still owns it (the fallback path may + // have installed its own controller, which it clears itself). + if (this.abortController === requestController) { + this.abortController = undefined + } } } @@ -566,8 +609,17 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio const baseUrl = this.options.openAiNativeBaseUrl || "https://api.openai.com" const url = `${baseUrl}/v1/responses` - // Create AbortController for cancellation - this.abortController = new AbortController() + // Create a request-local AbortController for cancellation. It is exposed via + // this.abortController so the stop-button path can observe it, but the bridging + // listener captures the local reference so a late abort from an earlier request + // can never reach a later request's controller. + const requestController = new AbortController() + this.abortController = requestController + + // Bridge the external abort signal onto the request controller; the returned + // cleanup detaches the listener in the finally block so a late abort from this + // request cannot cancel a later request's controller. + const detach = this.attachExternalAbort(metadata?.abortSignal, requestController) // Build per-request headers using taskId when available, falling back to sessionId const taskId = metadata?.taskId @@ -584,7 +636,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio "User-Agent": userAgent, }, body: JSON.stringify(requestBody), - signal: this.abortController.signal, + signal: requestController.signal, }) if (!response.ok) { @@ -648,8 +700,13 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio } // Handle streaming response - yield* this.handleStreamResponse(response.body, model) + yield* this.handleStreamResponse(response.body, model, requestController) } catch (error) { + // Re-throw abort errors as-is so callers can identify cancellations + if (error instanceof Error && error.name === "AbortError") { + throw error + } + const model = this.getModel() const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "createMessage") @@ -666,7 +723,12 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio // Handle non-Error objects throw new Error(`Unexpected error connecting to Responses API`) } finally { - this.abortController = undefined + // Detach the bridging listener so a late abort from this request cannot + // cancel a later request's controller. + detach?.() + if (this.abortController === requestController) { + this.abortController = undefined + } } } @@ -676,8 +738,17 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio * This function iterates through the Server-Sent Events (SSE) stream, parses each event, * and yields structured data chunks (`ApiStream`). It handles a wide variety of event types, * including text deltas, reasoning, usage data, and various status/tool events. + * + * @param requestController - The request-local controller. When it aborts (external + * abort or request timeout), a stream error is surfaced as the contract AbortError + * instead of being wrapped, so a cancellation is never misreported as a provider + * error and is not captured twice. */ - private async *handleStreamResponse(body: ReadableStream, model: OpenAiNativeModel): ApiStream { + private async *handleStreamResponse( + body: ReadableStream, + model: OpenAiNativeModel, + requestController: AbortController, + ): ApiStream { const reader = body.getReader() const decoder = new TextDecoder() let buffer = "" @@ -1129,6 +1200,15 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio // If we didn't get any content, don't throw - the API might have returned an empty response // This can happen in certain edge cases and shouldn't break the flow } catch (error) { + // The request-local controller is only aborted on external abort or request + // timeout, both legitimate cancellations: surface the contract abort error + // before any wrapping so the caller's AbortError guard sees one correctly + // named error instead of a wrapped DOMException, and the stop is not + // captured as an exception here or again by the caller. + if (requestController.signal.aborted) { + throw new DOMException(`The ${this.providerName} request was aborted`, "AbortError") + } + const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "createMessage") TelemetryService.instance.captureException(apiError) @@ -1506,9 +1586,13 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio return this.lastResponseId } async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { - try { - this.abortController = new AbortController() + // Request-local abort signal: merges the external abort signal with an optional + // timeout without touching this.abortController (owned by streaming requests). + const requestSignal = + RequestConfigBuilder.mergeAbortSignalAndTimeout(options?.abortSignal, options?.timeoutMs) ?? + new AbortController().signal + try { const model = this.getModel() const { verbosity } = model @@ -1567,7 +1651,7 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio // Make the non-streaming request const response = await (this.client as any).responses.create(requestBody, { - signal: this.abortController.signal, + signal: requestSignal, }) // Extract text from the response @@ -1590,6 +1674,11 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio return "" } catch (error) { + // Re-throw abort errors as-is so callers can identify cancellations + if (error instanceof Error && error.name === "AbortError") { + throw error + } + const errorModel = this.getModel() const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, errorModel.id, "completePrompt") @@ -1599,8 +1688,6 @@ export class OpenAiNativeHandler extends BaseProvider implements SingleCompletio throw new Error(`OpenAI Native completion error: ${error.message}`) } throw error - } finally { - this.abortController = undefined } } }