From 2a58311c0b33b27843544491a04e4c535709395a Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 20 Aug 2026 10:11:42 +0800 Subject: [PATCH 01/20] feat(api): abort signal support for gemini, mistral, lite-llm (completePrompt + createMessage) --- .../__tests__/gemini-handler.spec.ts | 38 ++++ src/api/providers/__tests__/gemini.spec.ts | 209 ++++++++++++++++++ src/api/providers/__tests__/lite-llm.spec.ts | 104 +++++++++ src/api/providers/__tests__/mistral.spec.ts | 145 +++++++++++- src/api/providers/__tests__/vertex.spec.ts | 56 +++++ src/api/providers/gemini.ts | 56 ++++- src/api/providers/lite-llm.ts | 54 ++++- src/api/providers/mistral.ts | 151 ++++++++----- 8 files changed, 751 insertions(+), 62 deletions(-) diff --git a/src/api/providers/__tests__/gemini-handler.spec.ts b/src/api/providers/__tests__/gemini-handler.spec.ts index 110f60289c..364f62e23d 100644 --- a/src/api/providers/__tests__/gemini-handler.spec.ts +++ b/src/api/providers/__tests__/gemini-handler.spec.ts @@ -55,6 +55,44 @@ describe("GeminiHandler backend support", () => { expect(promptConfig.tools).toBeUndefined() }) + it("completePrompt should pass abort signal through to client via httpOptions", async () => { + const options = { + apiProvider: "gemini", + enableUrlContext: false, + enableGrounding: false, + } as ApiHandlerOptions + const handler = new GeminiHandler(options) + + const controller = new AbortController() + const stub = vi.fn().mockResolvedValue({ text: "response" }) + handler["client"].models.generateContent = stub + + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + + expect(stub).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ + abortSignal: controller.signal, + }), + }), + ) + }) + + it("completePrompt should work without options (backward compatible)", async () => { + const options = { + apiProvider: "gemini", + enableUrlContext: false, + enableGrounding: false, + } as ApiHandlerOptions + const handler = new GeminiHandler(options) + + const stub = vi.fn().mockResolvedValue({ text: "response" }) + handler["client"].models.generateContent = stub + + const result = await handler.completePrompt("test prompt") + expect(result).toBe("response") + }) + describe("error scenarios", () => { it("should handle grounding metadata extraction failure gracefully", async () => { const options = { diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 701c453e3e..59968060dc 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -12,14 +12,22 @@ vitest.mock("@roo-code/telemetry", () => ({ import { Anthropic } from "@anthropic-ai/sdk" +import type { GenerateContentResponse } from "@google/genai" + import { type ModelInfo, geminiDefaultModelId, ApiProviderError } from "@roo-code/types" import { t } from "i18next" import { GeminiHandler } from "../gemini" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" +import { makeCreateMessageMetadata } from "../../../test-utils/api" const GEMINI_MODEL_NAME = geminiDefaultModelId +// @google/genai's GenerateContentResponse exposes `text` via a getter backed by +// `candidates`, so the stub only carries the field the provider reads; the double +// cast is the least-friction way to satisfy the class type in mocks. +const stubGenerateContentResponse = (text: string) => ({ text }) as unknown as GenerateContentResponse + describe("GeminiHandler", () => { let handler: GeminiHandler @@ -342,6 +350,71 @@ describe("GeminiHandler", () => { const result = await handler.completePrompt("Test prompt") expect(result).toBe("") }) + + it("should pass abort signal through to client via config.abortSignal", async () => { + const controller = new AbortController() + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + expect(handler["client"].models.generateContent).toHaveBeenCalledWith({ + model: GEMINI_MODEL_NAME, + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: { + abortSignal: controller.signal, + httpOptions: undefined, + temperature: 1, + }, + }) + }) + + it("should work without options (backward compatible)", async () => { + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + const result = await handler.completePrompt("test prompt") + expect(result).toBe("response") + expect(handler["client"].models.generateContent).toHaveBeenCalledWith({ + model: GEMINI_MODEL_NAME, + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: { + httpOptions: undefined, + temperature: 1, + }, + }) + }) + + it("should pass timeoutMs through to client via httpOptions with abortSignal on config", async () => { + const controller = new AbortController() + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 10000 }) + expect(handler["client"].models.generateContent).toHaveBeenCalledWith({ + model: GEMINI_MODEL_NAME, + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: { + abortSignal: controller.signal, + httpOptions: { timeout: 10000 }, + temperature: 1, + }, + }) + }) + + it("should pass only timeoutMs when no signal is provided", async () => { + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + await handler.completePrompt("test prompt", { timeoutMs: 5000 }) + expect(handler["client"].models.generateContent).toHaveBeenCalledWith({ + model: GEMINI_MODEL_NAME, + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: { + httpOptions: { timeout: 5000 }, + temperature: 1, + }, + }) + }) }) describe("getModel", () => { @@ -475,6 +548,142 @@ describe("GeminiHandler", () => { }) }) + describe("completePrompt request options", () => { + it("should pass timeout and baseUrl through httpOptions", async () => { + const handlerWithBaseUrl = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "https://gemini.example.test", + }) + handlerWithBaseUrl["client"] = handler["client"] + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + const result = await handlerWithBaseUrl.completePrompt("Test prompt", { timeoutMs: 1234 }) + + expect(result).toBe("Response") + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ + httpOptions: { + timeout: 1234, + baseUrl: "https://gemini.example.test", + }, + }), + }), + ) + }) + + it("should pass abortSignal on config instead of httpOptions", async () => { + const controller = new AbortController() + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + await handler.completePrompt("Test prompt", { abortSignal: controller.signal }) + + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ + abortSignal: controller.signal, + httpOptions: undefined, + }), + }), + ) + }) + + it("should omit httpOptions when timeoutMs and baseUrl are not provided", async () => { + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + await handler.completePrompt("Test prompt") + + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ + httpOptions: undefined, + }), + }), + ) + }) + }) + + describe("createMessage abort signal (bridging)", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + it("should reject immediately with AbortError when the external signal is pre-aborted", async () => { + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage( + "You are a helpful assistant", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(handler["client"].models.generateContentStream).not.toHaveBeenCalled() + }) + + it("should abort the in-flight request when the external signal is triggered", async () => { + const controller = new AbortController() + let capturedSignal: AbortSignal | undefined + const stub = vi.fn().mockImplementation(async (params: { config?: { abortSignal?: AbortSignal } }) => { + capturedSignal = params.config?.abortSignal + return (async function* () { + yield { text: "partial" } + if (capturedSignal?.aborted) { + throw new DOMException("aborted", "AbortError") + } + await new Promise((_resolve, reject) => { + capturedSignal?.addEventListener( + "abort", + () => reject(new DOMException("aborted", "AbortError")), + { once: true }, + ) + }) + })() + }) + handler["client"].models.generateContentStream = stub + + const stream = handler.createMessage( + "You are a helpful assistant", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const collector = collectStream(stream).catch((e: unknown) => e) + await new Promise((resolve) => setTimeout(resolve, 10)) + controller.abort() + + const error = await collector + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(capturedSignal).toBeDefined() + expect(capturedSignal?.aborted).toBe(true) + }) + + it("should not set config.abortSignal when no external signal is provided", async () => { + const stub = vi.fn().mockReturnValue((async function* () {})()) + handler["client"].models.generateContentStream = stub + + await collectStream(handler.createMessage("You are a helpful assistant", messages)) + + const config = stub.mock.calls[0][0].config + expect(config.abortSignal).toBeUndefined() + }) + }) + describe("error telemetry", () => { const mockMessages: Anthropic.Messages.MessageParam[] = [ { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index eee5cf52bb..badab1e1a7 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -6,6 +6,7 @@ import { ApiHandlerOptions } from "../../../shared/api" import { litellmDefaultModelId, litellmDefaultModelInfo } from "@roo-code/types" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" import { clearAllMocks } from "../../../test-utils/reset" +import { makeCreateMessageMetadata } from "../../../test-utils/api" // Mock vscode first to avoid import errors vi.mock("vscode", () => ({ @@ -1235,4 +1236,107 @@ describe("LiteLLMHandler", () => { expect(requestHeaders).not.toHaveProperty("X-Zoo-Session-ID") }) }) + + describe("completePrompt", () => { + it("should pass abort signal through to client", async () => { + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) + const controller = new AbortController() + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: expect.any(String) }), + expect.objectContaining({ signal: controller.signal }), + ) + }) + + it("should pass timeout through to client", async () => { + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) + await handler.completePrompt("test prompt", { timeoutMs: 5000 }) + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: expect.any(String) }), + expect.objectContaining({ timeout: 5000 }), + ) + }) + + it("should merge signal and timeoutMs together", async () => { + const controller = new AbortController() + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 10000 }) + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: expect.any(String) }), + expect.objectContaining({ signal: controller.signal, timeout: 10000 }), + ) + }) + + it("should work without options (backward compatible)", async () => { + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) + const result = await handler.completePrompt("test prompt") + expect(result).toBe("response") + }) + }) + + describe("createMessage abort signal (bridging)", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + it("should reject immediately with AbortError when the external signal is pre-aborted", async () => { + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(mockCreate).not.toHaveBeenCalled() + }) + + it("should abort the in-flight stream when the external signal is triggered", async () => { + const controller = new AbortController() + let capturedSignal: AbortSignal | undefined + // The stream is built inside the mock implementation so that capturedSignal + // is already set before the abort-aware chunk is created. + mockCreate.mockImplementationOnce((_body: unknown, options?: { signal?: AbortSignal }) => { + capturedSignal = options?.signal + const mockStream = asyncStreamFrom([ + { + choices: [{ delta: { content: "partial" } }], + usage: undefined, + }, + new Promise((_resolve, reject) => { + const onAbort = () => reject(new DOMException("aborted", "AbortError")) + if (capturedSignal?.aborted) { + onAbort() + return + } + capturedSignal?.addEventListener("abort", onAbort, { once: true }) + }), + ]) + return { withResponse: vi.fn().mockResolvedValue({ data: mockStream }) } + }) + + const stream = handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const collector = collectStream(stream).catch((e: unknown) => e) + await new Promise((resolve) => setTimeout(resolve, 10)) + controller.abort() + + const error = await collector + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(capturedSignal).toBeDefined() + expect(capturedSignal?.aborted).toBe(true) + }) + }) }) diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index f2a7591bd8..aff33b6030 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -54,6 +54,7 @@ import { MistralHandler } from "../mistral" import type { ApiHandlerOptions } from "../../../shared/api" import type { ApiHandlerCreateMessageMetadata } from "../../index" import type { ApiStreamTextChunk, ApiStreamReasoningChunk, ApiStreamToolCallPartialChunk } from "../../transform/stream" +import { makeCreateMessageMetadata } from "../../../test-utils/api" describe("MistralHandler", () => { let handler: MistralHandler @@ -447,11 +448,14 @@ describe("MistralHandler", () => { const prompt = "Test prompt" const result = await handler.completePrompt(prompt) - expect(mockComplete).toHaveBeenCalledWith({ - model: mockOptions.apiModelId, - messages: [{ role: "user", content: prompt }], - temperature: 0, - }) + expect(mockComplete).toHaveBeenCalledWith( + { + model: mockOptions.apiModelId, + messages: [{ role: "user", content: prompt }], + temperature: 0, + }, + undefined, + ) expect(result).toBe("Test response") }) @@ -483,5 +487,136 @@ describe("MistralHandler", () => { mockComplete.mockRejectedValueOnce(new Error("API Error")) await expect(handler.completePrompt("Test prompt")).rejects.toThrow("Mistral completion error: API Error") }) + + it("should pass abort signal through to client", async () => { + const controller = new AbortController() + mockComplete.mockResolvedValueOnce({ + choices: [{ message: { content: "response" } }], + }) + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { + fetchOptions: { signal: controller.signal }, + }) + }) + + it("should work without options (backward compatible)", async () => { + mockComplete.mockResolvedValueOnce({ + choices: [{ message: { content: "response" } }], + }) + const result = await handler.completePrompt("test prompt") + expect(result).toBe("response") + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) + }) + + it("should pass timeout through to client", async () => { + const controller = new AbortController() + mockComplete.mockResolvedValueOnce({ + choices: [{ message: { content: "response" } }], + }) + await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 5000 }) + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { + fetchOptions: { signal: controller.signal }, + timeoutMs: 5000, + }) + }) + + it("should pass only timeoutMs when no signal provided", async () => { + mockComplete.mockResolvedValueOnce({ + choices: [{ message: { content: "response" } }], + }) + await handler.completePrompt("test prompt", { timeoutMs: 3000 }) + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { + timeoutMs: 3000, + }) + }) + + it("should still forward timeoutMs=0 (uses !== undefined check, not truthy check)", async () => { + mockComplete.mockResolvedValueOnce({ + choices: [{ message: { content: "response" } }], + }) + await handler.completePrompt("test prompt", { timeoutMs: 0 }) + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { + timeoutMs: 0, + }) + }) + }) + + describe("createMessage abort signal bridging", () => { + const systemPrompt = "You are a helpful assistant." + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: [{ type: "text", text: "Hello!" }], + }, + ] + + it("should reject immediately with AbortError when the external signal is pre-aborted", async () => { + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(mockCreate).not.toHaveBeenCalled() + }) + + it("should abort the in-flight stream when the external signal is triggered", async () => { + const controller = new AbortController() + let capturedSignal: AbortSignal | undefined + mockCreate.mockImplementationOnce( + async (_options: unknown, requestOptions?: { fetchOptions?: { signal?: AbortSignal } }) => { + capturedSignal = requestOptions?.fetchOptions?.signal + return asyncStreamFrom([ + { + data: { + choices: [ + { + delta: { content: "partial" }, + index: 0, + }, + ], + }, + }, + new Promise((_resolve, reject) => { + const onAbort = () => reject(new DOMException("aborted", "AbortError")) + if (capturedSignal?.aborted) { + onAbort() + return + } + capturedSignal?.addEventListener("abort", onAbort, { once: true }) + }), + ]) + }, + ) + + const stream = handler.createMessage( + systemPrompt, + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const collector = collectStream(stream).catch((e: unknown) => e) + await new Promise((resolve) => setTimeout(resolve, 10)) + controller.abort() + + const error = await collector + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(capturedSignal).toBeDefined() + expect(capturedSignal?.aborted).toBe(true) + }) + + it("should not pass a signal to the stream call when no external signal is provided", async () => { + const stream = handler.createMessage(systemPrompt, messages) + await collectStream(stream) + const streamOptions = mockCreate.mock.calls[0][1] + expect(streamOptions).toBeUndefined() + }) }) }) diff --git a/src/api/providers/__tests__/vertex.spec.ts b/src/api/providers/__tests__/vertex.spec.ts index a304518ca7..99c5787e0e 100644 --- a/src/api/providers/__tests__/vertex.spec.ts +++ b/src/api/providers/__tests__/vertex.spec.ts @@ -21,10 +21,19 @@ vitest.mock("@roo-code/telemetry", () => ({ import { Anthropic } from "@anthropic-ai/sdk" +import type { GenerateContentResponse } from "@google/genai" + import { ApiStreamChunk } from "../../transform/stream" import { t } from "i18next" import { VertexHandler } from "../vertex" +import { collectStream } from "../../../test-utils/stream" +import { makeCreateMessageMetadata } from "../../../test-utils/api" + +// @google/genai's GenerateContentResponse exposes `text` via a getter backed by +// `candidates`, so the stub only carries the field the provider reads; the double +// cast is the least-friction way to satisfy the class type in mocks. +const stubGenerateContentResponse = (text: string) => ({ text }) as unknown as GenerateContentResponse describe("VertexHandler", () => { let handler: VertexHandler @@ -137,6 +146,53 @@ describe("VertexHandler", () => { const result = await handler.completePrompt("Test prompt") expect(result).toBe("") }) + + it("should pass abort signal through to client via config.abortSignal", async () => { + const controller = new AbortController() + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + + await handler.completePrompt("test prompt", { abortSignal: controller.signal }) + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + model: expect.any(String), + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: expect.objectContaining({ + abortSignal: controller.signal, + httpOptions: undefined, + temperature: 1, + }), + }), + ) + }) + + it("should work without options (backward compatible)", async () => { + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + + const result = await handler.completePrompt("test prompt") + expect(result).toBe("response") + }) + }) + + describe("createMessage abort signal (inherited from GeminiHandler)", () => { + it("should reject immediately with AbortError when the external signal is pre-aborted", async () => { + const controller = new AbortController() + controller.abort() + + const stream = handler.createMessage( + "You are a helpful assistant", + [{ role: "user", content: "Hello" }], + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect(handler["client"].models.generateContentStream).not.toHaveBeenCalled() + }) }) describe("getModel", () => { diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index ec0d14e4c9..bc6bcfe2f1 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -344,7 +344,30 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } } - const params: GenerateContentParameters = { model, contents, config } + // Bridge the external abort signal from Task (metadata.abortSignal) into a + // request-local controller so the in-flight generateContentStream request + // can be cancelled. The @google/genai SDK merges this signal with its own + // timeout handling, which is preserved rather than replaced. + // A pre-aborted signal rejects immediately with an AbortError. + const externalAbortSignal = metadata?.abortSignal + let requestAbortController: AbortController | undefined + let externalAbortListener: (() => void) | undefined + if (externalAbortSignal) { + if (externalAbortSignal.aborted) { + throw new DOMException("Gemini request aborted", "AbortError") + } + const controller = new AbortController() + requestAbortController = controller + const onExternalAbort = () => controller.abort() + externalAbortListener = onExternalAbort + externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) + } + + const params: GenerateContentParameters = { + model, + contents, + config: requestAbortController ? { ...config, abortSignal: requestAbortController.signal } : config, + } try { const result = await this.client.models.generateContentStream(params) @@ -477,6 +500,11 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } } } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (metadata?.abortSignal?.aborted) { + throw new DOMException("Gemini request aborted", "AbortError") + } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "createMessage") TelemetryService.instance.captureException(apiError) @@ -486,6 +514,10 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } throw error + } finally { + if (externalAbortSignal && externalAbortListener) { + externalAbortSignal.removeEventListener("abort", externalAbortListener) + } } } @@ -585,14 +617,25 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl const temperatureConfig: number | undefined = supportsTemperature ? (this.options.modelTemperature ?? info.defaultTemperature ?? 1) : info.defaultTemperature + const httpOpts: { timeout?: number; baseUrl?: string } = {} + if (options?.timeoutMs !== undefined) { + httpOpts.timeout = options.timeoutMs + } + if (this.options.googleGeminiBaseUrl) { + httpOpts.baseUrl = this.options.googleGeminiBaseUrl + } const promptConfig: GenerateContentConfig = { - httpOptions: this.options.googleGeminiBaseUrl - ? { baseUrl: this.options.googleGeminiBaseUrl } - : undefined, + httpOptions: Object.keys(httpOpts).length > 0 ? httpOpts : undefined, temperature: temperatureConfig, } + // @google/genai expects request cancellation on config.abortSignal + // (not httpOptions.signal), so the signal is passed directly to the config. + if (options?.abortSignal) { + promptConfig.abortSignal = options.abortSignal + } + const request = { model, contents: [{ role: "user", parts: [{ text: prompt }] }], @@ -613,6 +656,11 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl return text } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (options?.abortSignal?.aborted) { + throw new DOMException("Gemini completion aborted", "AbortError") + } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "completePrompt") TelemetryService.instance.captureException(apiError) diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index 8cfe2d0a19..da82dd19cb 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -246,9 +246,31 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa requestHeaders["X-Zoo-Session-ID"] = metadata.taskId } + // Bridge the external abort signal from Task (metadata.abortSignal) into a + // request-local controller so the in-flight streaming request can be + // cancelled. A pre-aborted signal rejects immediately with an AbortError. + const externalAbortSignal = metadata?.abortSignal + let requestAbortController: AbortController | undefined + let externalAbortListener: (() => void) | undefined + if (externalAbortSignal) { + if (externalAbortSignal.aborted) { + throw new DOMException("LiteLLM streaming aborted", "AbortError") + } + const controller = new AbortController() + requestAbortController = controller + const onExternalAbort = () => controller.abort() + externalAbortListener = onExternalAbort + externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) + } + try { const { data: completion } = await this.client.chat.completions - .create(requestOptions, { headers: requestHeaders }) + .create( + requestOptions, + requestAbortController + ? { headers: requestHeaders, signal: requestAbortController.signal } + : { headers: requestHeaders }, + ) .withResponse() let lastUsage @@ -315,10 +337,19 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa yield usageData } } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (metadata?.abortSignal?.aborted) { + throw new DOMException("LiteLLM streaming aborted", "AbortError") + } if (error instanceof Error) { throw new Error(`LiteLLM streaming error: ${error.message}`) } throw error + } finally { + if (externalAbortSignal && externalAbortListener) { + externalAbortSignal.removeEventListener("abort", externalAbortListener) + } } } @@ -345,9 +376,28 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa requestOptions.max_tokens = info.maxTokens } - const response = await this.client.chat.completions.create(requestOptions) + // Build request options with abortSignal and/or timeout. The OpenAI SDK + // treats a timeout of 0 as an immediate timeout, so non-positive timeoutMs + // values disable the timeout instead of being forwarded. + const createOptions: OpenAI.RequestOptions = {} + if (options?.abortSignal) { + createOptions.signal = options.abortSignal + } + if (options?.timeoutMs !== undefined && options.timeoutMs > 0) { + createOptions.timeout = options.timeoutMs + } + + const response = await this.client.chat.completions.create( + requestOptions, + Object.keys(createOptions).length > 0 ? createOptions : undefined, + ) return response.choices[0]?.message.content || "" } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (options?.abortSignal?.aborted) { + throw new DOMException("LiteLLM completion aborted", "AbortError") + } if (error instanceof Error) { throw new Error(`LiteLLM completion error: ${error.message}`) } diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index c7816feaa2..9c304f7eb0 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -101,66 +101,98 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand // Temporary debug log for QA // console.log("[MISTRAL DEBUG] Raw API request body:", requestOptions) + // Bridge the external abort signal from Task (metadata.abortSignal) into a + // request-local controller so the in-flight streaming request can be + // cancelled. A pre-aborted signal rejects immediately with an AbortError. + const externalAbortSignal = metadata?.abortSignal + let requestAbortController: AbortController | undefined + let externalAbortListener: (() => void) | undefined + if (externalAbortSignal) { + if (externalAbortSignal.aborted) { + throw new DOMException("Mistral completion aborted", "AbortError") + } + const controller = new AbortController() + requestAbortController = controller + const onExternalAbort = () => controller.abort() + externalAbortListener = onExternalAbort + externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) + } + let response try { - response = await this.client.chat.stream(requestOptions) - } catch (error) { - const errorMessage = error instanceof Error ? error.message : String(error) - const apiError = new ApiProviderError(errorMessage, this.providerName, model, "createMessage") - TelemetryService.instance.captureException(apiError) - throw new Error(`Mistral completion error: ${errorMessage}`) - } + if (requestAbortController) { + response = await this.client.chat.stream(requestOptions, { + fetchOptions: { signal: requestAbortController.signal }, + }) + } else { + response = await this.client.chat.stream(requestOptions) + } - for await (const event of response) { - const delta = event.data.choices[0]?.delta - - if (delta?.content) { - if (typeof delta.content === "string") { - // Handle string content as text - yield { type: "text", text: delta.content } - } else if (Array.isArray(delta.content)) { - // Handle array of content chunks - // The SDK v1.9.18 supports ThinkChunk with type "thinking" - for (const chunk of delta.content as ContentChunkWithThinking[]) { - if (chunk.type === "thinking" && chunk.thinking) { - // Handle thinking content as reasoning chunks - // ThinkChunk has a 'thinking' property that contains an array of text/reference chunks - for (const thinkingPart of chunk.thinking) { - if (thinkingPart.type === "text" && thinkingPart.text) { - yield { type: "reasoning", text: thinkingPart.text } + for await (const event of response) { + const delta = event.data.choices[0]?.delta + + if (delta?.content) { + if (typeof delta.content === "string") { + // Handle string content as text + yield { type: "text", text: delta.content } + } else if (Array.isArray(delta.content)) { + // Handle array of content chunks + // The SDK v1.9.18 supports ThinkChunk with type "thinking" + for (const chunk of delta.content as ContentChunkWithThinking[]) { + if (chunk.type === "thinking" && chunk.thinking) { + // Handle thinking content as reasoning chunks + // ThinkChunk has a 'thinking' property that contains an array of text/reference chunks + for (const thinkingPart of chunk.thinking) { + if (thinkingPart.type === "text" && thinkingPart.text) { + yield { type: "reasoning", text: thinkingPart.text } + } } + } else if (chunk.type === "text" && chunk.text) { + // Handle text content normally + yield { type: "text", text: chunk.text } } - } else if (chunk.type === "text" && chunk.text) { - // Handle text content normally - yield { type: "text", text: chunk.text } } } } - } - // Handle tool calls in stream - // Mistral SDK provides tool_calls in delta similar to OpenAI format - const toolCalls = (delta as { toolCalls?: MistralToolCall[] })?.toolCalls - if (toolCalls) { - for (let i = 0; i < toolCalls.length; i++) { - const toolCall = toolCalls[i] - yield { - type: "tool_call_partial", - index: i, - id: toolCall.id, - name: toolCall.function?.name, - arguments: toolCall.function?.arguments, + // Handle tool calls in stream + // Mistral SDK provides tool_calls in delta similar to OpenAI format + const toolCalls = (delta as { toolCalls?: MistralToolCall[] })?.toolCalls + if (toolCalls) { + for (let i = 0; i < toolCalls.length; i++) { + const toolCall = toolCalls[i] + yield { + type: "tool_call_partial", + index: i, + id: toolCall.id, + name: toolCall.function?.name, + arguments: toolCall.function?.arguments, + } } } - } - if (event.data.usage) { - yield { - type: "usage", - inputTokens: event.data.usage.promptTokens || 0, - outputTokens: event.data.usage.completionTokens || 0, + if (event.data.usage) { + yield { + type: "usage", + inputTokens: event.data.usage.promptTokens || 0, + outputTokens: event.data.usage.completionTokens || 0, + } } } + } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (metadata?.abortSignal?.aborted) { + throw new DOMException("Mistral completion aborted", "AbortError") + } + const errorMessage = error instanceof Error ? error.message : String(error) + const apiError = new ApiProviderError(errorMessage, this.providerName, model, "createMessage") + TelemetryService.instance.captureException(apiError) + throw new Error(`Mistral completion error: ${errorMessage}`) + } finally { + if (externalAbortSignal && externalAbortListener) { + externalAbortSignal.removeEventListener("abort", externalAbortListener) + } } } @@ -196,11 +228,23 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand const { id: model, temperature } = this.getModel() try { - const response = await this.client.chat.complete({ - model, - messages: [{ role: "user", content: prompt }], - temperature, - }) + // Build Mistral SDK RequestOptions + const requestOptions: Parameters[1] = {} + if (options?.abortSignal) { + requestOptions.fetchOptions = { signal: options.abortSignal } + } + if (options?.timeoutMs !== undefined) { + requestOptions.timeoutMs = options.timeoutMs + } + + const response = await this.client.chat.complete( + { + model, + messages: [{ role: "user", content: prompt }], + temperature, + }, + Object.keys(requestOptions).length > 0 ? requestOptions : undefined, + ) const content = response.choices?.[0]?.message.content @@ -214,6 +258,11 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand return content || "" } catch (error) { + // User-initiated abort: surface a standard AbortError rather than the + // wrapped provider error so callers can distinguish cancellation. + if (options?.abortSignal?.aborted) { + throw new DOMException("Mistral completion aborted", "AbortError") + } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "completePrompt") TelemetryService.instance.captureException(apiError) From f6eba43d3992f48e657feea481333ee6976d48f0 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 20 Aug 2026 10:58:21 +0800 Subject: [PATCH 02/20] fix(api): unify timeoutMs:0 handling across gemini/mistral/lite-llm + fix test title --- .../__tests__/gemini-handler.spec.ts | 2 +- src/api/providers/__tests__/gemini.spec.ts | 15 ++++++++++ src/api/providers/__tests__/lite-llm.spec.ts | 6 ++++ src/api/providers/__tests__/mistral.spec.ts | 6 ++-- src/api/providers/gemini.ts | 9 ++++-- src/api/providers/lite-llm.ts | 13 ++++---- src/api/providers/mistral.ts | 8 +++-- .../utils/__tests__/request-timeout.spec.ts | 30 +++++++++++++++++++ src/api/providers/utils/request-timeout.ts | 10 +++++++ 9 files changed, 85 insertions(+), 14 deletions(-) create mode 100644 src/api/providers/utils/__tests__/request-timeout.spec.ts create mode 100644 src/api/providers/utils/request-timeout.ts diff --git a/src/api/providers/__tests__/gemini-handler.spec.ts b/src/api/providers/__tests__/gemini-handler.spec.ts index 364f62e23d..232849a091 100644 --- a/src/api/providers/__tests__/gemini-handler.spec.ts +++ b/src/api/providers/__tests__/gemini-handler.spec.ts @@ -55,7 +55,7 @@ describe("GeminiHandler backend support", () => { expect(promptConfig.tools).toBeUndefined() }) - it("completePrompt should pass abort signal through to client via httpOptions", async () => { + it("completePrompt should pass abort signal through to client via config.abortSignal", async () => { const options = { apiProvider: "gemini", enableUrlContext: false, diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 59968060dc..6350f3d395 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -415,6 +415,21 @@ describe("GeminiHandler", () => { }, }) }) + + it("should omit httpOptions entirely for timeoutMs=0 (0 disables the timeout)", async () => { + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("response"), + ) + await handler.completePrompt("test prompt", { timeoutMs: 0 }) + expect(handler["client"].models.generateContent).toHaveBeenCalledWith({ + model: GEMINI_MODEL_NAME, + contents: [{ role: "user", parts: [{ text: "test prompt" }] }], + config: { + httpOptions: undefined, + temperature: 1, + }, + }) + }) }) describe("getModel", () => { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index badab1e1a7..324c532735 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1272,6 +1272,12 @@ describe("LiteLLMHandler", () => { const result = await handler.completePrompt("test prompt") expect(result).toBe("response") }) + + it("should omit the timeout option for timeoutMs=0 (0 would abort immediately in the OpenAI SDK)", async () => { + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) + await handler.completePrompt("test prompt", { timeoutMs: 0 }) + expect(mockCreate).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) + }) }) describe("createMessage abort signal (bridging)", () => { diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index aff33b6030..2fb0c43f22 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -530,14 +530,12 @@ describe("MistralHandler", () => { }) }) - it("should still forward timeoutMs=0 (uses !== undefined check, not truthy check)", async () => { + it("should omit the timeout option for timeoutMs=0 (0 disables the timeout)", async () => { mockComplete.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }], }) await handler.completePrompt("test prompt", { timeoutMs: 0 }) - expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { - timeoutMs: 0, - }) + expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) }) diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index bc6bcfe2f1..434cb06a27 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -27,6 +27,7 @@ import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, Complete import { BaseProvider } from "./base-provider" import { NOT_PROVIDED } from "./constants" import { parseVertexJsonCredentials } from "./utils/vertex-credentials" +import { getRequestTimeoutMs } from "./utils/request-timeout" type GeminiHandlerOptions = ApiHandlerOptions & { isVertex?: boolean @@ -618,8 +619,12 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl ? (this.options.modelTemperature ?? info.defaultTemperature ?? 1) : info.defaultTemperature const httpOpts: { timeout?: number; baseUrl?: string } = {} - if (options?.timeoutMs !== undefined) { - httpOpts.timeout = options.timeoutMs + // Per the abort-signal series contract, timeoutMs <= 0 means 'no per-request + // timeout': the option is omitted entirely (some SDKs treat 0 as an + // immediate timeout). + const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) + if (timeoutMs !== undefined) { + httpOpts.timeout = timeoutMs } if (this.options.googleGeminiBaseUrl) { httpOpts.baseUrl = this.options.googleGeminiBaseUrl diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index da82dd19cb..af5152ef70 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -16,6 +16,7 @@ import { sanitizeOpenAiCallId } from "../../utils/tool-id" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" import { RouterProvider } from "./router-provider" import { extractReasoningFromDelta } from "./utils/extract-reasoning" +import { getRequestTimeoutMs } from "./utils/request-timeout" /** * LiteLLM provider handler @@ -376,15 +377,17 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa requestOptions.max_tokens = info.maxTokens } - // Build request options with abortSignal and/or timeout. The OpenAI SDK - // treats a timeout of 0 as an immediate timeout, so non-positive timeoutMs - // values disable the timeout instead of being forwarded. + // Build request options with abortSignal and/or timeout. Per the + // abort-signal series contract, timeoutMs <= 0 means 'no per-request + // timeout': the option is omitted entirely, because the OpenAI SDK treats + // a timeout of 0 as an immediate timeout. const createOptions: OpenAI.RequestOptions = {} if (options?.abortSignal) { createOptions.signal = options.abortSignal } - if (options?.timeoutMs !== undefined && options.timeoutMs > 0) { - createOptions.timeout = options.timeoutMs + const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) + if (timeoutMs !== undefined) { + createOptions.timeout = timeoutMs } const response = await this.client.chat.completions.create( diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index 9c304f7eb0..80748fda9f 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -16,6 +16,7 @@ import { ApiHandlerOptions } from "../../shared/api" import { convertToMistralMessages } from "../transform/mistral-format" import { ApiStream } from "../transform/stream" import { handleProviderError } from "./utils/error-handler" +import { getRequestTimeoutMs } from "./utils/request-timeout" import { BaseProvider } from "./base-provider" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" @@ -233,8 +234,11 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand if (options?.abortSignal) { requestOptions.fetchOptions = { signal: options.abortSignal } } - if (options?.timeoutMs !== undefined) { - requestOptions.timeoutMs = options.timeoutMs + // Per the abort-signal series contract, timeoutMs <= 0 means 'no per-request + // timeout': the option is omitted entirely. + const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) + if (timeoutMs !== undefined) { + requestOptions.timeoutMs = timeoutMs } const response = await this.client.chat.complete( diff --git a/src/api/providers/utils/__tests__/request-timeout.spec.ts b/src/api/providers/utils/__tests__/request-timeout.spec.ts new file mode 100644 index 0000000000..5e969a1036 --- /dev/null +++ b/src/api/providers/utils/__tests__/request-timeout.spec.ts @@ -0,0 +1,30 @@ +import { getRequestTimeoutMs } from "../request-timeout" + +describe("getRequestTimeoutMs", () => { + it("forwards positive timeout values unchanged", () => { + expect(getRequestTimeoutMs(5000)).toBe(5000) + expect(getRequestTimeoutMs(1)).toBe(1) + expect(getRequestTimeoutMs(1234)).toBe(1234) + }) + + it("returns undefined for zero (timeout disabled, not an immediate abort)", () => { + expect(getRequestTimeoutMs(0)).toBeUndefined() + }) + + it("returns undefined for negative values", () => { + expect(getRequestTimeoutMs(-1)).toBeUndefined() + expect(getRequestTimeoutMs(-5000)).toBeUndefined() + }) + + it("returns undefined when no value is provided", () => { + expect(getRequestTimeoutMs()).toBeUndefined() + expect(getRequestTimeoutMs(undefined)).toBeUndefined() + }) + + it("guards against non-number input at the runtime boundary", () => { + expect(getRequestTimeoutMs(NaN)).toBeUndefined() + // Non-number values can only reach this helper through untyped callers + // (e.g. user settings); the double cast exercises the typeof guard. + expect(getRequestTimeoutMs("5000" as unknown as number)).toBeUndefined() + }) +}) diff --git a/src/api/providers/utils/request-timeout.ts b/src/api/providers/utils/request-timeout.ts new file mode 100644 index 0000000000..a3d1e56d4c --- /dev/null +++ b/src/api/providers/utils/request-timeout.ts @@ -0,0 +1,10 @@ +/** + * Returns the value to pass as a client/SDK request timeout option, or undefined. + * + * Per the abort-signal series contract, timeoutMs <= 0 (or undefined) means + * 'no per-request timeout': the option is omitted entirely, because some SDKs + * (e.g. the OpenAI Node SDK) treat timeout: 0 as an IMMEDIATE timeout. + */ +export function getRequestTimeoutMs(timeoutMs?: number): number | undefined { + return typeof timeoutMs === "number" && timeoutMs > 0 ? timeoutMs : undefined +} From 4a1307248c2e8ace2d83bcc5a7558f75d58066fb Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 20 Aug 2026 19:32:03 +0800 Subject: [PATCH 03/20] test(api): close changed-line coverage gaps in gemini, mistral, lite-llm --- src/api/providers/__tests__/gemini.spec.ts | 13 +++++ src/api/providers/__tests__/lite-llm.spec.ts | 13 +++++ src/api/providers/__tests__/mistral.spec.ts | 56 +++++++++++++++++++- 3 files changed, 81 insertions(+), 1 deletion(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 6350f3d395..f3b497c73f 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -430,6 +430,19 @@ describe("GeminiHandler", () => { }, }) }) + + it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { + const controller = new AbortController() + controller.abort() + vi.mocked(handler["client"].models.generateContent).mockRejectedValue(new Error("Gemini API error")) + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(DOMException) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Gemini completion aborted") + }) }) describe("getModel", () => { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 324c532735..ee701b1134 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1278,6 +1278,19 @@ describe("LiteLLMHandler", () => { await handler.completePrompt("test prompt", { timeoutMs: 0 }) expect(mockCreate).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) + + it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { + mockCreate.mockRejectedValueOnce(new Error("LiteLLM API error")) + const controller = new AbortController() + controller.abort() + + const error = await handler + .completePrompt("test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(DOMException) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("LiteLLM completion aborted") + }) }) describe("createMessage abort signal (bridging)", () => { diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index 2fb0c43f22..a5bfc75fd0 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -53,7 +53,12 @@ import type OpenAI from "openai" import { MistralHandler } from "../mistral" import type { ApiHandlerOptions } from "../../../shared/api" import type { ApiHandlerCreateMessageMetadata } from "../../index" -import type { ApiStreamTextChunk, ApiStreamReasoningChunk, ApiStreamToolCallPartialChunk } from "../../transform/stream" +import type { + ApiStreamTextChunk, + ApiStreamReasoningChunk, + ApiStreamToolCallPartialChunk, + ApiStreamUsageChunk, +} from "../../transform/stream" import { makeCreateMessageMetadata } from "../../../test-utils/api" describe("MistralHandler", () => { @@ -234,6 +239,42 @@ describe("MistralHandler", () => { expect(results[1]).toEqual({ type: "reasoning", text: "Some reasoning" }) expect(results[2]).toEqual({ type: "text", text: "Second text" }) }) + + it("should yield a usage chunk when the stream event carries usage data", async () => { + // The final event carries usage without any delta content; the handler + // must translate it into a usage chunk with the reported token counts. + mockCreate.mockImplementationOnce(async (_options) => + asyncStreamFrom([ + { + data: { + choices: [ + { + delta: { content: "Test response" }, + index: 0, + }, + ], + }, + }, + { + data: { + choices: [], + usage: { promptTokens: 12, completionTokens: 34 }, + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: (ApiStreamTextChunk | ApiStreamUsageChunk)[] = [] + + for await (const chunk of iterator) { + results.push(chunk as ApiStreamTextChunk | ApiStreamUsageChunk) + } + + expect(results).toHaveLength(2) + expect(results[0]).toEqual({ type: "text", text: "Test response" }) + expect(results[1]).toEqual({ type: "usage", inputTokens: 12, outputTokens: 34 }) + }) }) describe("native tool calling", () => { @@ -537,6 +578,19 @@ describe("MistralHandler", () => { await handler.completePrompt("test prompt", { timeoutMs: 0 }) expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) + + it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { + mockComplete.mockRejectedValueOnce(new Error("API Error")) + const controller = new AbortController() + controller.abort() + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(DOMException) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Mistral completion aborted") + }) }) describe("createMessage abort signal bridging", () => { From 73af57e9cdce07943100fa3d68a6f9ca44e26954 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 3 Sep 2026 04:08:15 +0800 Subject: [PATCH 04/20] test(api): use provider identifiers in gemini-handler spec Replace raw gemini apiProvider literals in the two abort-signal spec cases with providerIdentifiers.gemini, matching the rest of the file and the zoo/no-raw-provider-identifiers rule that CI lint enforces. --- src/api/providers/__tests__/gemini-handler.spec.ts | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/api/providers/__tests__/gemini-handler.spec.ts b/src/api/providers/__tests__/gemini-handler.spec.ts index 62de561492..d7594a01c5 100644 --- a/src/api/providers/__tests__/gemini-handler.spec.ts +++ b/src/api/providers/__tests__/gemini-handler.spec.ts @@ -58,7 +58,7 @@ describe("GeminiHandler backend support", () => { it("completePrompt should pass abort signal through to client via config.abortSignal", async () => { const options = { - apiProvider: "gemini", + apiProvider: providerIdentifiers.gemini, enableUrlContext: false, enableGrounding: false, } as ApiHandlerOptions @@ -81,7 +81,7 @@ describe("GeminiHandler backend support", () => { it("completePrompt should work without options (backward compatible)", async () => { const options = { - apiProvider: "gemini", + apiProvider: providerIdentifiers.gemini, enableUrlContext: false, enableGrounding: false, } as ApiHandlerOptions From fefffcd4092e1f7fdefc62810293484cccfc8293 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 3 Sep 2026 06:41:23 +0800 Subject: [PATCH 05/20] fix(api): harden gemini/mistral abort, timeout and base-url handling Address CodeRabbit findings on the abort-signal series: - gemini: reject non-HTTPS (non-loopback) googleGeminiBaseUrl before requests so API keys are never sent over cleartext (CWE-319) - mistral: route completePrompt timeout through mergeAbortSignalAndTimeout so the timeout actually cancels the request, instead of a dead timeoutMs field - gemini/mistral/lite-llm specs: request-local signal identity assertions, readiness barrier instead of fixed delay, makeApiHandlerOptions over casts --- .../__tests__/gemini-handler.spec.ts | 89 ++++------------ src/api/providers/__tests__/gemini.spec.ts | 100 ++++++++++++++++++ src/api/providers/__tests__/lite-llm.spec.ts | 9 +- src/api/providers/__tests__/mistral.spec.ts | 46 ++++++-- src/api/providers/gemini.ts | 54 ++++++++++ src/api/providers/mistral.ts | 15 ++- 6 files changed, 228 insertions(+), 85 deletions(-) diff --git a/src/api/providers/__tests__/gemini-handler.spec.ts b/src/api/providers/__tests__/gemini-handler.spec.ts index d7594a01c5..5b1b6c91c4 100644 --- a/src/api/providers/__tests__/gemini-handler.spec.ts +++ b/src/api/providers/__tests__/gemini-handler.spec.ts @@ -12,8 +12,7 @@ vi.mock("@roo-code/telemetry", () => ({ })) import { GeminiHandler } from "../gemini" -import type { ApiHandlerOptions } from "../../../shared/api" -import { providerIdentifiers } from "@roo-code/types/provider-identifiers" +import { makeApiHandlerOptions } from "../../../test-utils/api" describe("GeminiHandler backend support", () => { beforeEach(() => { @@ -24,11 +23,7 @@ describe("GeminiHandler backend support", () => { // URL context and grounding are mutually exclusive with function declarations // in Gemini API, so createMessage only uses function declarations. // URL context/grounding are only added in completePrompt. - const options = { - apiProvider: providerIdentifiers.gemini, - enableUrlContext: true, - enableGrounding: true, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -41,11 +36,7 @@ describe("GeminiHandler backend support", () => { }) it("completePrompt passes config overrides without tools when URL context and grounding disabled", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableUrlContext: false, - enableGrounding: false, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockResolvedValue({ text: "ok" }) // @ts-ignore access private client @@ -57,11 +48,7 @@ describe("GeminiHandler backend support", () => { }) it("completePrompt should pass abort signal through to client via config.abortSignal", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableUrlContext: false, - enableGrounding: false, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const controller = new AbortController() @@ -80,11 +67,7 @@ describe("GeminiHandler backend support", () => { }) it("completePrompt should work without options (backward compatible)", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableUrlContext: false, - enableGrounding: false, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockResolvedValue({ text: "response" }) @@ -96,10 +79,7 @@ describe("GeminiHandler backend support", () => { describe("error scenarios", () => { it("should handle grounding metadata extraction failure gracefully", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableGrounding: true, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const mockStream = async function* () { @@ -131,10 +111,7 @@ describe("GeminiHandler backend support", () => { }) it("should handle malformed grounding metadata", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableGrounding: true, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const mockStream = async function* () { @@ -182,11 +159,7 @@ describe("GeminiHandler backend support", () => { }) it("should handle API errors when tools are enabled", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - enableUrlContext: true, - enableGrounding: true, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const mockError = new Error("API rate limit exceeded") @@ -230,9 +203,7 @@ describe("GeminiHandler backend support", () => { ] it("should ignore allowedFunctionNames because Gemini rejects larger restriction lists", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -251,9 +222,7 @@ describe("GeminiHandler backend support", () => { }) it("should include all tools when allowedFunctionNames is provided", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -274,9 +243,7 @@ describe("GeminiHandler backend support", () => { }) it("should not pass large allowedFunctionNames lists to Gemini", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -305,9 +272,7 @@ describe("GeminiHandler backend support", () => { }) it("should not pass allowedFunctionNames even when history includes tool calls", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -342,9 +307,7 @@ describe("GeminiHandler backend support", () => { }) it("should fall back to tool_choice when allowedFunctionNames is provided", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -365,9 +328,7 @@ describe("GeminiHandler backend support", () => { }) it("should fall back to tool_choice when allowedFunctionNames is empty", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -389,9 +350,7 @@ describe("GeminiHandler backend support", () => { }) it("should not set toolConfig when allowedFunctionNames is undefined and no tool_choice", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -412,9 +371,7 @@ describe("GeminiHandler backend support", () => { describe("Gemini schema compatibility", () => { it("should strip broad JSON Schema metadata from function declarations", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -473,9 +430,7 @@ describe("GeminiHandler backend support", () => { }) it("should collapse composition and type arrays in function declaration schemas", async () => { - const options = { - apiProvider: providerIdentifiers.gemini, - } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -534,7 +489,7 @@ describe("GeminiHandler backend support", () => { }) it("should deep-merge allOf fragments instead of overwriting earlier properties", async () => { - const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -579,7 +534,7 @@ describe("GeminiHandler backend support", () => { }) it("should resolve $ref entries before dropping $defs", async () => { - const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -628,7 +583,7 @@ describe("GeminiHandler backend support", () => { }) it("should preserve top-level properties and required entries when allOf is also present", async () => { - const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -670,7 +625,7 @@ describe("GeminiHandler backend support", () => { }) it("should stop recursive $ref expansion before the sanitized schema becomes cyclic", async () => { - const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client @@ -722,7 +677,7 @@ describe("GeminiHandler backend support", () => { }) it("should preserve parameter names that collide with stripped schema keywords", async () => { - const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions + const options = makeApiHandlerOptions() const handler = new GeminiHandler(options) const stub = vi.fn().mockReturnValue((async function* () {})()) // @ts-ignore access private client diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 21a4ab9174..00775f5c97 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -676,6 +676,103 @@ describe("GeminiHandler", () => { }), ) }) + describe("googleGeminiBaseUrl security (CWE-319)", () => { + it("should reject a non-HTTPS non-loopback googleGeminiBaseUrl in completePrompt", async () => { + const insecureHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://gemini.example.test", + }) + insecureHandler["client"] = handler["client"] + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + await expect(insecureHandler.completePrompt("Test prompt")).rejects.toThrow( + t("common:errors.gemini.generate_complete_prompt", { + error: "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + }), + ) + expect(handler["client"].models.generateContent).not.toHaveBeenCalled() + }) + + it("should allow a loopback HTTP googleGeminiBaseUrl in completePrompt", async () => { + const loopbackHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.0.0.1:8080", + }) + loopbackHandler["client"] = handler["client"] + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + const result = await loopbackHandler.completePrompt("Test prompt") + + expect(result).toBe("Response") + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ + httpOptions: { + baseUrl: "http://127.0.0.1:8080", + }, + }), + }), + ) + }) + + it("should reject a non-HTTPS non-loopback googleGeminiBaseUrl in createMessage", async () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const insecureHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://insecure.example.com", + }) + insecureHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const stream = insecureHandler.createMessage("You are a helpful assistant", messages) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect(stub).not.toHaveBeenCalled() + }) + + it("should allow a loopback HTTP googleGeminiBaseUrl in createMessage", async () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const loopbackHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.0.0.1:8080", + }) + loopbackHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + await collectStream(loopbackHandler.createMessage("You are a helpful assistant", messages)) + + const config = stub.mock.calls[0][0].config + expect(config.httpOptions).toEqual({ baseUrl: "http://127.0.0.1:8080" }) + }) + }) }) describe("createMessage abort signal (bridging)", () => { @@ -737,6 +834,9 @@ describe("GeminiHandler", () => { expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect(capturedSignal).toBeDefined() + // The in-flight request must run against a request-local signal, not the + // external one forwarded by reference. + expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) }) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 451f9f9ad8..dd43f68a84 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1324,10 +1324,17 @@ describe("LiteLLMHandler", () => { it("should abort the in-flight stream when the external signal is triggered", async () => { const controller = new AbortController() let capturedSignal: AbortSignal | undefined + // Readiness barrier: resolves once the request-local signal is captured and + // the (mocked) request has started, instead of guessing a fixed delay. + let requestStartedResolve!: () => void + const requestStarted = new Promise((resolve) => { + requestStartedResolve = resolve + }) // The stream is built inside the mock implementation so that capturedSignal // is already set before the abort-aware chunk is created. mockCreate.mockImplementationOnce((_body: unknown, options?: { signal?: AbortSignal }) => { capturedSignal = options?.signal + requestStartedResolve() const mockStream = asyncStreamFrom([ { choices: [{ delta: { content: "partial" } }], @@ -1352,7 +1359,7 @@ describe("LiteLLMHandler", () => { ) const collector = collectStream(stream).catch((e: unknown) => e) - await new Promise((resolve) => setTimeout(resolve, 10)) + await requestStarted controller.abort() const error = await collector diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index a5bfc75fd0..de7e0c908a 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -549,25 +549,52 @@ describe("MistralHandler", () => { expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) - it("should pass timeout through to client", async () => { + it("should pass a composite abort+timeout signal through to client", async () => { const controller = new AbortController() mockComplete.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }], }) await handler.completePrompt("test prompt", { abortSignal: controller.signal, timeoutMs: 5000 }) - expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { - fetchOptions: { signal: controller.signal }, - timeoutMs: 5000, - }) + const callArgs = mockComplete.mock.calls[0][1] as { fetchOptions?: { signal?: AbortSignal } } | undefined + expect(callArgs).toBeDefined() + expect(callArgs?.fetchOptions?.signal).toBeInstanceOf(AbortSignal) + // A fresh composite signal — not the external signal forwarded by reference. + expect(callArgs?.fetchOptions?.signal).not.toBe(controller.signal) + expect(callArgs).not.toHaveProperty("timeoutMs") + // The external side is bridged into the composite. + controller.abort() + expect(callArgs?.fetchOptions?.signal?.aborted).toBe(true) }) - it("should pass only timeoutMs when no signal provided", async () => { + it("should pass a timeout signal through to client when no external signal is provided", async () => { mockComplete.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }], }) await handler.completePrompt("test prompt", { timeoutMs: 3000 }) - expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), { - timeoutMs: 3000, + const callArgs = mockComplete.mock.calls[0][1] as { fetchOptions?: { signal?: AbortSignal } } | undefined + expect(callArgs).toBeDefined() + expect(callArgs?.fetchOptions?.signal).toBeInstanceOf(AbortSignal) + expect(callArgs).not.toHaveProperty("timeoutMs") + }) + + it("should bridge the per-request timeout into the composite signal", async () => { + let capturedSignal: AbortSignal | undefined + mockComplete.mockImplementationOnce( + (_options: unknown, requestOptions?: { fetchOptions?: { signal?: AbortSignal } }) => { + capturedSignal = requestOptions?.fetchOptions?.signal + return Promise.resolve({ + choices: [{ message: { content: "response" } }], + }) + }, + ) + + await handler.completePrompt("test prompt", { timeoutMs: 200 }) + + expect(capturedSignal).toBeInstanceOf(AbortSignal) + // The timeout side fires on its own: the self-managed AbortSignal.timeout + // aborts the captured signal after the 200ms per-request deadline. + await vi.waitFor(() => { + expect(capturedSignal?.aborted).toBe(true) }) }) @@ -661,6 +688,9 @@ describe("MistralHandler", () => { expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect(capturedSignal).toBeDefined() + // The in-flight stream must run against a request-local signal, not the + // external one forwarded by reference. + expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) }) diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 434cb06a27..222fc3bb64 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -173,6 +173,53 @@ function sanitizeSchemaForGemini( return result } +// googleGeminiBaseUrl is user-editable and can reach non-HTTPS values (settings, +// imported profiles). The @google/genai client keeps API-key authentication for +// custom endpoints, so reject cleartext base URLs before any request — with a +// narrow loopback exception for local test proxies. +function isLoopbackUrl(value: string): boolean { + try { + const parsed = new URL(value) + if (parsed.protocol !== "http:" && parsed.protocol !== "https:") { + return false + } + return ( + parsed.hostname === "localhost" || + parsed.hostname === "::1" || + parsed.hostname === "[::1]" || + /^127\./.test(parsed.hostname) + ) + } catch { + return false + } +} + +// Throws an ApiProviderError when baseUrl is not HTTPS (loopback HTTP is the +// narrow exception, for local test proxies). The provider/model/operation +// arguments keep the structured error context consistent with the request-path +// ApiProviderError instances in this file. +function assertSecureGeminiBaseUrl(baseUrl: string, modelId: string, operation: string): void { + let parsed: URL + try { + parsed = new URL(baseUrl) + } catch { + throw new ApiProviderError("Invalid Google Gemini base URL (not a valid URL)", "Gemini", modelId, operation) + } + if (parsed.protocol === "https:") { + return + } + if (parsed.protocol === "http:" && isLoopbackUrl(baseUrl)) { + // Loopback endpoints (localhost/127.x/::1) are allowed for local test proxies. + return + } + throw new ApiProviderError( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + "Gemini", + modelId, + operation, + ) +} + export class GeminiHandler extends BaseProvider implements SingleCompletionHandler { protected options: ApiHandlerOptions @@ -297,6 +344,12 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl ? (this.options.modelTemperature ?? info.defaultTemperature ?? 1) : info.defaultTemperature + // Reject cleartext (non-loopback) base URLs before building the request so the + // API key is never sent over an insecure endpoint. + if (this.options.googleGeminiBaseUrl) { + assertSecureGeminiBaseUrl(this.options.googleGeminiBaseUrl, model, "createMessage") + } + const config: GenerateContentConfig = { systemInstruction, httpOptions: this.options.googleGeminiBaseUrl ? { baseUrl: this.options.googleGeminiBaseUrl } : undefined, @@ -627,6 +680,7 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl httpOpts.timeout = timeoutMs } if (this.options.googleGeminiBaseUrl) { + assertSecureGeminiBaseUrl(this.options.googleGeminiBaseUrl, model, "completePrompt") httpOpts.baseUrl = this.options.googleGeminiBaseUrl } diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index 80748fda9f..93e40b54f2 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -16,7 +16,7 @@ import { ApiHandlerOptions } from "../../shared/api" import { convertToMistralMessages } from "../transform/mistral-format" import { ApiStream } from "../transform/stream" import { handleProviderError } from "./utils/error-handler" -import { getRequestTimeoutMs } from "./utils/request-timeout" +import { mergeAbortSignalAndTimeout } from "./utils/abort-signal" import { BaseProvider } from "./base-provider" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" @@ -231,14 +231,11 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand try { // Build Mistral SDK RequestOptions const requestOptions: Parameters[1] = {} - if (options?.abortSignal) { - requestOptions.fetchOptions = { signal: options.abortSignal } - } - // Per the abort-signal series contract, timeoutMs <= 0 means 'no per-request - // timeout': the option is omitted entirely. - const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) - if (timeoutMs !== undefined) { - requestOptions.timeoutMs = timeoutMs + // Build a single signal that combines the external abort with the per-request + // timeout (timeoutMs <= 0 disables the timeout; see mergeAbortSignalAndTimeout). + const signal = mergeAbortSignalAndTimeout(options?.abortSignal, options?.timeoutMs) + if (signal) { + requestOptions.fetchOptions = { signal } } const response = await this.client.chat.complete( From 9b8033a611e8c93716486663c044eb71c00496ca Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Thu, 3 Sep 2026 09:45:42 +0800 Subject: [PATCH 06/20] chore: retrigger CodeRabbit review (no-op) The incremental review for the previous head was stuck in a phantom "review finished" state on the CodeRabbit side (the review object never materialized), so this no-op commit moves the head to a fresh sha and forces a new incremental review. No code changes. From 6c2d6bf45703d86725380aaf62f664167d3d0aa7 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 05:41:26 +0800 Subject: [PATCH 07/20] fix(api): allow IPv6 loopback host and strengthen abort-signal spec kills (mutation gate) --- src/api/providers/__tests__/gemini.spec.ts | 178 ++++++++++++++++++- src/api/providers/__tests__/lite-llm.spec.ts | 56 +++++- src/api/providers/__tests__/mistral.spec.ts | 151 +++++++++++++++- src/api/providers/gemini.ts | 26 ++- src/api/providers/lite-llm.ts | 1 + src/api/providers/mistral.ts | 1 + 6 files changed, 394 insertions(+), 19 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 00775f5c97..373aa8a550 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -482,6 +482,27 @@ describe("GeminiHandler", () => { expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("Gemini completion aborted") }) + + it("should surface the wrapped provider error when the request fails without options", async () => { + vi.mocked(handler["client"].models.generateContent).mockRejectedValue(new Error("Gemini API error")) + + const error = await handler.completePrompt("Test prompt").catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + // The catch must take the non-abort wrapping path (not the abort + // DOMException path) and must not crash reading a missing signal. + // The i18n message itself is asserted by the error telemetry tests. + expect((error as Error).message).not.toContain("Cannot read properties") + }) + + it("should surface the wrapped provider error when the request fails with options but no signal", async () => { + vi.mocked(handler["client"].models.generateContent).mockRejectedValue(new Error("Gemini API error")) + + const error = await handler.completePrompt("Test prompt", { timeoutMs: 0 }).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + // Options exist but no signal: exercises the second optional-chain + // position, which would crash here if the `?.` were removed. + expect((error as Error).message).not.toContain("Cannot read properties") + }) }) describe("getModel", () => { @@ -747,6 +768,8 @@ describe("GeminiHandler", () => { expect((error as ApiProviderError).message).toBe( "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect((error as ApiProviderError).operation).toBe("createMessage") expect(stub).not.toHaveBeenCalled() }) @@ -772,6 +795,109 @@ describe("GeminiHandler", () => { const config = stub.mock.calls[0][0].config expect(config.httpOptions).toEqual({ baseUrl: "http://127.0.0.1:8080" }) }) + + it("should allow an http://localhost googleGeminiBaseUrl in createMessage", async () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const localhostHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://localhost:8080", + }) + localhostHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + await collectStream(localhostHandler.createMessage("You are a helpful assistant", messages)) + + expect(stub.mock.calls[0][0].config.httpOptions).toEqual({ baseUrl: "http://localhost:8080" }) + }) + + it("should allow an http://[::1] googleGeminiBaseUrl (IPv6 loopback host)", async () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const ipv6Handler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://[::1]:8080", + }) + ipv6Handler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + await collectStream(ipv6Handler.createMessage("You are a helpful assistant", messages)) + + expect(stub.mock.calls[0][0].config.httpOptions).toEqual({ baseUrl: "http://[::1]:8080" }) + }) + + it("should reject hostnames that only resemble the 127. range", async () => { + // "127a.b" shares the 127 prefix but is not a loopback address and + // must not pass the anchored, dot-escaped 127. check. (Hosts such + // as a127.0.0.1 are rejected by new URL() outright, which is also + // why a de-anchored 127. pattern is unobservable.) + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127a.b:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + + it("should reject an invalid googleGeminiBaseUrl in createMessage", async () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const invalidHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "not a valid url", + }) + invalidHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + invalidHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe("Invalid Google Gemini base URL (not a valid URL)") + expect((error as ApiProviderError).provider).toBe("Gemini") + expect((error as ApiProviderError).operation).toBe("createMessage") + expect(stub).not.toHaveBeenCalled() + }) }) }) @@ -796,11 +922,14 @@ describe("GeminiHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Gemini request aborted") expect(handler["client"].models.generateContentStream).not.toHaveBeenCalled() }) it("should abort the in-flight request when the external signal is triggered", async () => { const controller = new AbortController() + const addEventListenerSpy = vi.spyOn(controller.signal, "addEventListener") + const removeEventListenerSpy = vi.spyOn(controller.signal, "removeEventListener") let capturedSignal: AbortSignal | undefined const stub = vi.fn().mockImplementation(async (params: { config?: { abortSignal?: AbortSignal } }) => { capturedSignal = params.config?.abortSignal @@ -830,14 +959,27 @@ describe("GeminiHandler", () => { await new Promise((resolve) => setTimeout(resolve, 10)) controller.abort() - const error = await collector + // Bound the wait so a broken abort bridge fails this test fast (and fails + // the Stryker mutant) instead of hanging until the runner timeout. + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("abort propagation deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Gemini request aborted") expect(capturedSignal).toBeDefined() // The in-flight request must run against a request-local signal, not the // external one forwarded by reference. expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) + // The bridge registers a once-only listener on the external signal and + // detaches it when the request settles. + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) }) it("should not set config.abortSignal when no external signal is provided", async () => { @@ -849,6 +991,40 @@ describe("GeminiHandler", () => { const config = stub.mock.calls[0][0].config expect(config.abortSignal).toBeUndefined() }) + + it("should wrap a non-abort stream failure with the i18n message and capture telemetry", async () => { + const mockError = new Error("Gemini stream failure") + handler["client"].models.generateContentStream = vi.fn().mockRejectedValue(mockError) + + const stream = handler.createMessage("You are a helpful assistant", messages, makeCreateMessageMetadata()) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + // The catch must take the non-abort wrapping path (not the abort + // DOMException path) and must not crash reading a missing signal. + expect((error as Error).message).not.toContain("Cannot read properties") + expect(mockCaptureException).toHaveBeenCalledTimes(1) + expect(mockCaptureException).toHaveBeenCalledWith( + expect.objectContaining({ + message: "Gemini stream failure", + provider: "Gemini", + operation: "createMessage", + }), + ) + }) + + it("should wrap a stream failure when no metadata is provided at all", async () => { + const mockError = new Error("Gemini stream failure") + handler["client"].models.generateContentStream = vi.fn().mockRejectedValue(mockError) + + const stream = handler.createMessage("You are a helpful assistant", messages) + + const error = await collectStream(stream).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + // No metadata at all: exercises the first optional-chain position, + // which would crash here if the `?.` were removed. + expect((error as Error).message).not.toContain("Cannot read properties") + }) }) describe("error telemetry", () => { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 4b5f75b472..611cd4d5c2 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1436,6 +1436,21 @@ describe("LiteLLMHandler", () => { expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("LiteLLM completion aborted") }) + + it("should surface the wrapped provider error when the request fails without options", async () => { + mockCreate.mockRejectedValueOnce(new Error("boom")) + + const error = await handler.completePrompt("test prompt").catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).message).toBe("LiteLLM completion error: boom") + }) + + it("should surface the wrapped provider error when the request fails with options but no signal", async () => { + mockCreate.mockRejectedValueOnce(new Error("boom")) + + const error = await handler.completePrompt("test prompt", { timeoutMs: 0 }).catch((e: unknown) => e) + expect((error as Error).message).toBe("LiteLLM completion error: boom") + }) }) describe("createMessage abort signal (bridging)", () => { @@ -1459,11 +1474,14 @@ describe("LiteLLMHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("LiteLLM streaming aborted") expect(mockCreate).not.toHaveBeenCalled() }) it("should abort the in-flight stream when the external signal is triggered", async () => { const controller = new AbortController() + const addEventListenerSpy = vi.spyOn(controller.signal, "addEventListener") + const removeEventListenerSpy = vi.spyOn(controller.signal, "removeEventListener") let capturedSignal: AbortSignal | undefined // Readiness barrier: resolves once the request-local signal is captured and // the (mocked) request has started, instead of guessing a fixed delay. @@ -1503,11 +1521,47 @@ describe("LiteLLMHandler", () => { await requestStarted controller.abort() - const error = await collector + // Bound the wait so a broken abort bridge fails this test fast (and fails + // the Stryker mutant) instead of hanging until the runner timeout. + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("abort propagation deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("LiteLLM streaming aborted") expect(capturedSignal).toBeDefined() expect(capturedSignal?.aborted).toBe(true) + // The bridge registers a once-only listener on the external signal and + // detaches it when the request settles. + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) + }) + + it("should wrap a non-abort stream failure with the i18n-free provider message and no metadata", async () => { + // createMessage awaits create(...).withResponse(), so the failure must + // surface from the withResponse() call, not from create() itself. + mockCreate.mockReturnValueOnce({ + withResponse: vi.fn().mockRejectedValue(new Error("boom")), + }) + + const error = await collectStream(handler.createMessage("system", messages)).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).message).toBe("LiteLLM streaming error: boom") + }) + + it("should wrap a non-abort stream failure when metadata exists without an abort signal", async () => { + mockCreate.mockReturnValueOnce({ + withResponse: vi.fn().mockRejectedValue(new Error("boom")), + }) + + const error = await collectStream( + handler.createMessage("system", messages, makeCreateMessageMetadata()), + ).catch((e: unknown) => e) + expect((error as Error).message).toBe("LiteLLM streaming error: boom") }) }) }) diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index de7e0c908a..55f859d990 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -158,6 +158,14 @@ describe("MistralHandler", () => { it("should handle errors gracefully", async () => { mockCreate.mockRejectedValueOnce(new Error("API Error")) await expect(handler.createMessage(systemPrompt, messages).next()).rejects.toThrow("API Error") + expect(mockCaptureException).toHaveBeenCalledTimes(1) + expect(mockCaptureException).toHaveBeenCalledWith( + expect.objectContaining({ + message: "API Error", + provider: "Mistral", + operation: "createMessage", + }), + ) }) it("should handle thinking content as reasoning chunks", async () => { @@ -240,6 +248,121 @@ describe("MistralHandler", () => { expect(results[2]).toEqual({ type: "text", text: "Second text" }) }) + it("should ignore non-string, non-array delta content", async () => { + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([{ data: { choices: [{ delta: { content: 42 }, index: 0 }] } }]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(0) + }) + + it("should ignore a thinking chunk with no thinking payload", async () => { + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { + data: { + choices: [{ delta: { content: [{ type: "thinking", thinking: undefined }] }, index: 0 }], + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(0) + }) + + it("should ignore non-text thinking parts", async () => { + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { + data: { + choices: [ + { + delta: { + content: [ + { + type: "thinking", + thinking: [{ type: "reference", text: "A reference" }], + }, + ], + }, + index: 0, + }, + ], + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(0) + }) + + it("should ignore empty text chunks and unknown chunk types", async () => { + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { + data: { + choices: [ + { + delta: { + content: [{ type: "text", text: "" }, { type: "other" }], + }, + index: 0, + }, + ], + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(0) + }) + + it("should yield a tool call partial without function details when only an id is provided", async () => { + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { + data: { + choices: [{ delta: { toolCalls: [{ id: "t1" }] }, index: 0 }], + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(1) + expect(results[0]).toEqual({ + type: "tool_call_partial", + index: 0, + id: "t1", + name: undefined, + arguments: undefined, + }) + }) + it("should yield a usage chunk when the stream event carries usage data", async () => { // The final event carries usage without any delta content; the handler // must translate it into a usage chunk with the reported token counts. @@ -642,11 +765,14 @@ describe("MistralHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Mistral completion aborted") expect(mockCreate).not.toHaveBeenCalled() }) it("should abort the in-flight stream when the external signal is triggered", async () => { const controller = new AbortController() + const addEventListenerSpy = vi.spyOn(controller.signal, "addEventListener") + const removeEventListenerSpy = vi.spyOn(controller.signal, "removeEventListener") let capturedSignal: AbortSignal | undefined mockCreate.mockImplementationOnce( async (_options: unknown, requestOptions?: { fetchOptions?: { signal?: AbortSignal } }) => { @@ -684,14 +810,37 @@ describe("MistralHandler", () => { await new Promise((resolve) => setTimeout(resolve, 10)) controller.abort() - const error = await collector + // Bound the wait so a broken abort bridge fails this test fast (and fails + // the Stryker mutant) instead of hanging until the runner timeout. + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("abort propagation deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("Mistral completion aborted") expect(capturedSignal).toBeDefined() // The in-flight stream must run against a request-local signal, not the // external one forwarded by reference. expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) + // The bridge registers a once-only listener on the external signal and + // detaches it when the request settles. + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) + }) + + it("should wrap a non-abort stream failure when metadata is provided without a signal", async () => { + mockCreate.mockRejectedValueOnce(new Error("boom")) + + const error = await collectStream( + handler.createMessage(systemPrompt, messages, makeCreateMessageMetadata()), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).message).toBe("Mistral completion error: boom") }) it("should not pass a signal to the stream call when no external signal is provided", async () => { diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 222fc3bb64..a04cd47c60 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -177,21 +177,13 @@ function sanitizeSchemaForGemini( // imported profiles). The @google/genai client keeps API-key authentication for // custom endpoints, so reject cleartext base URLs before any request — with a // narrow loopback exception for local test proxies. -function isLoopbackUrl(value: string): boolean { - try { - const parsed = new URL(value) - if (parsed.protocol !== "http:" && parsed.protocol !== "https:") { - return false - } - return ( - parsed.hostname === "localhost" || - parsed.hostname === "::1" || - parsed.hostname === "[::1]" || - /^127\./.test(parsed.hostname) - ) - } catch { - return false - } +// The caller (assertSecureGeminiBaseUrl) has already parsed the URL and only +// reaches this for the HTTP exception, so the parsed hostname is passed +// directly; `new URL` keeps the brackets in IPv6 hostnames, so `[::1]` is +// the loopback host form to compare against. +function isLoopbackHostname(hostname: string): boolean { + // Stryker disable next-line Regex: the ^ anchor is unobservable — new URL() rejects every non-loopback hostname containing "127." (e.g. a127.0.0.1, foo.127.0.0.1), so a de-anchored pattern behaves identically on every reachable hostname + return hostname === "localhost" || hostname === "[::1]" || /^127\./.test(hostname) } // Throws an ApiProviderError when baseUrl is not HTTPS (loopback HTTP is the @@ -208,7 +200,7 @@ function assertSecureGeminiBaseUrl(baseUrl: string, modelId: string, operation: if (parsed.protocol === "https:") { return } - if (parsed.protocol === "http:" && isLoopbackUrl(baseUrl)) { + if (parsed.protocol === "http:" && isLoopbackHostname(parsed.hostname)) { // Loopback endpoints (localhost/127.x/::1) are allowed for local test proxies. return } @@ -569,6 +561,7 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl throw error } finally { + // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here if (externalAbortSignal && externalAbortListener) { externalAbortSignal.removeEventListener("abort", externalAbortListener) } @@ -680,6 +673,7 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl httpOpts.timeout = timeoutMs } if (this.options.googleGeminiBaseUrl) { + // Stryker disable next-line StringLiteral: completePrompt's catch reads only .message from the thrown error before building its own telemetry error, so this operation argument is unobservable assertSecureGeminiBaseUrl(this.options.googleGeminiBaseUrl, model, "completePrompt") httpOpts.baseUrl = this.options.googleGeminiBaseUrl } diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index d14a03769f..7ddd687794 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -372,6 +372,7 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa } throw error } finally { + // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here if (externalAbortSignal && externalAbortListener) { externalAbortSignal.removeEventListener("abort", externalAbortListener) } diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index 93e40b54f2..36514ea284 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -191,6 +191,7 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand TelemetryService.instance.captureException(apiError) throw new Error(`Mistral completion error: ${errorMessage}`) } finally { + // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here if (externalAbortSignal && externalAbortListener) { externalAbortSignal.removeEventListener("abort", externalAbortListener) } From 1b205a38ae77341d9f689090b2db1614d62d9157 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 06:00:34 +0800 Subject: [PATCH 08/20] test(api): kill remaining abort-signal mutation survivors in gemini and mistral specs (mutation gate) --- src/api/providers/__tests__/gemini.spec.ts | 31 +++++++++++ src/api/providers/__tests__/mistral.spec.ts | 62 +++++++++++++++++++++ 2 files changed, 93 insertions(+) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 373aa8a550..b34f82fda5 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -872,6 +872,37 @@ describe("GeminiHandler", () => { expect(stub).not.toHaveBeenCalled() }) + it("should reject a non-HTTP loopback scheme such as ftp://localhost", async () => { + // The protocol operand (left of the && in the loopback exception) must + // still be enforced: a non-HTTP scheme aimed at a loopback host is not + // a local test proxy and must be rejected like any other non-HTTPS URL. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const ftpHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "ftp://localhost:8080", + }) + ftpHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + ftpHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + it("should reject an invalid googleGeminiBaseUrl in createMessage", async () => { const messages: Anthropic.Messages.MessageParam[] = [ { diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index 55f859d990..6ed20ed97e 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -311,6 +311,58 @@ describe("MistralHandler", () => { expect(results).toHaveLength(0) }) + it("should yield text for a text chunk that also carries a stray thinking payload", async () => { + // Dispatch is on chunk.type, not on payload presence: a text chunk with a + // stray thinking array must be emitted as text, never as reasoning. + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { + data: { + choices: [ + { + delta: { + content: [ + { + type: "text", + text: "hello", + thinking: [{ type: "text", text: "reason" }], + }, + ], + }, + index: 0, + }, + ], + }, + }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(1) + expect(results[0]).toEqual({ type: "text", text: "hello" }) + }) + + it("should ignore a non-text chunk that carries a text payload", async () => { + // Dispatch is on chunk.type, not on payload presence: an unknown chunk + // type with a truthy text payload must not be emitted as text. + mockCreate.mockImplementationOnce(async () => + asyncStreamFrom([ + { data: { choices: [{ delta: { content: [{ type: "other", text: "leak" }] }, index: 0 }] } }, + ]), + ) + + const iterator = handler.createMessage(systemPrompt, messages) + const results: unknown[] = [] + for await (const chunk of iterator) { + results.push(chunk) + } + expect(results).toHaveLength(0) + }) + it("should ignore empty text chunks and unknown chunk types", async () => { mockCreate.mockImplementationOnce(async () => asyncStreamFrom([ @@ -652,6 +704,16 @@ describe("MistralHandler", () => { await expect(handler.completePrompt("Test prompt")).rejects.toThrow("Mistral completion error: API Error") }) + it("should wrap a non-abort error when options are provided without an abort signal", async () => { + // options is a defined object lacking abortSignal: the optional-chaining + // guard must not crash reading a missing abortSignal before wrapping the + // provider error. + mockComplete.mockRejectedValueOnce(new Error("boom")) + await expect(handler.completePrompt("Test prompt", { timeoutMs: 5000 })).rejects.toThrow( + "Mistral completion error: boom", + ) + }) + it("should pass abort signal through to client", async () => { const controller = new AbortController() mockComplete.mockResolvedValueOnce({ From 018edc10214f11a72d6524e8b9bc40bd272238b8 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 08:21:04 +0800 Subject: [PATCH 09/20] fix(api): address review findings on abort-bridge listener assertions and 127. hostname anchoring MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - gemini/lite-llm/mistral specs: capture the bridge's registered abort listener and assert the exact reference on removeEventListener instead of expect.any(Function) (CodeRabbit finding). - gemini spec: add a foo127.bar case — a hostname new URL() accepts that contains the 127. substring but is not loopback — proving the anchored 127. check is observable; drop the Stryker disable Regex directive in gemini.ts whose unobservable-anchor rationale is wrong (hosts such as a127.0.0.1 are rejected by new URL() itself). --- src/api/providers/__tests__/gemini.spec.ts | 51 +++++++++++++++++--- src/api/providers/__tests__/lite-llm.spec.ts | 11 +++-- src/api/providers/__tests__/mistral.spec.ts | 11 +++-- src/api/providers/gemini.ts | 4 +- 4 files changed, 64 insertions(+), 13 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index b34f82fda5..eb3665f317 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -842,9 +842,11 @@ describe("GeminiHandler", () => { it("should reject hostnames that only resemble the 127. range", async () => { // "127a.b" shares the 127 prefix but is not a loopback address and - // must not pass the anchored, dot-escaped 127. check. (Hosts such - // as a127.0.0.1 are rejected by new URL() outright, which is also - // why a de-anchored 127. pattern is unobservable.) + // must not pass the anchored, dot-escaped 127. check. This case pins + // the dot escape itself: an unescaped /^127/ pattern would match + // "127a.b" and wrongly allow cleartext. (Hosts such as a127.0.0.1 + // are rejected by new URL() outright — its mixed digit-led/letter-led + // label rule — so they never reach the check.) const messages: Anthropic.Messages.MessageParam[] = [ { role: "user", @@ -872,6 +874,38 @@ describe("GeminiHandler", () => { expect(stub).not.toHaveBeenCalled() }) + it("should reject a valid hostname containing the 127. substring that is not loopback", async () => { + // "foo127.bar" is accepted by new URL() (a syntactically valid + // hostname), contains the "127." substring, yet is not a loopback + // address: only the anchored 127. check rejects it. A de-anchored + // /127\./ pattern would match it and wrongly allow cleartext. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://foo127.bar:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + it("should reject a non-HTTP loopback scheme such as ftp://localhost", async () => { // The protocol operand (left of the && in the loopback exception) must // still be enforced: a non-HTTP scheme aimed at a loopback host is not @@ -1008,9 +1042,14 @@ describe("GeminiHandler", () => { expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) + // detaches it when the request settles. Target the last "abort" + // registration (the bridge's listener) and assert the exact reference + // so a bridge that removes a different callback cannot pass. + const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") + const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] + expect(typeof addedListener).toBe("function") + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) }) it("should not set config.abortSignal when no external signal is provided", async () => { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 611cd4d5c2..ab910cc49a 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1536,9 +1536,14 @@ describe("LiteLLMHandler", () => { expect(capturedSignal).toBeDefined() expect(capturedSignal?.aborted).toBe(true) // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) + // detaches it when the request settles. Target the last "abort" + // registration (the bridge's listener) and assert the exact reference + // so a bridge that removes a different callback cannot pass. + const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") + const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] + expect(typeof addedListener).toBe("function") + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) }) it("should wrap a non-abort stream failure with the i18n-free provider message and no metadata", async () => { diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index 6ed20ed97e..e1ce3185cb 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -890,9 +890,14 @@ describe("MistralHandler", () => { expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function), { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", expect.any(Function)) + // detaches it when the request settles. Target the last "abort" + // registration (the bridge's listener) and assert the exact reference + // so a bridge that removes a different callback cannot pass. + const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") + const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] + expect(typeof addedListener).toBe("function") + expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) }) it("should wrap a non-abort stream failure when metadata is provided without a signal", async () => { diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index a04cd47c60..3d89c21cac 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -182,7 +182,9 @@ function sanitizeSchemaForGemini( // directly; `new URL` keeps the brackets in IPv6 hostnames, so `[::1]` is // the loopback host form to compare against. function isLoopbackHostname(hostname: string): boolean { - // Stryker disable next-line Regex: the ^ anchor is unobservable — new URL() rejects every non-loopback hostname containing "127." (e.g. a127.0.0.1, foo.127.0.0.1), so a de-anchored pattern behaves identically on every reachable hostname + // The ^ anchor is observable: new URL() accepts non-loopback hostnames that + // contain the "127." substring (e.g. foo127.bar), so a de-anchored pattern + // would misclassify them as loopback and allow cleartext. return hostname === "localhost" || hostname === "[::1]" || /^127\./.test(hostname) } From 8505c9ff46ffd955657e29ab437b91e5cb2bb5d9 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 08:51:31 +0800 Subject: [PATCH 10/20] fix(api): restrict Gemini loopback HTTP exception to literal 127.0.0.0/8 A public hostname may legitimately start with a 127. label (e.g. 127.example.test), which the prefix-style /^127\./ test misclassified as loopback and allowed as a cleartext base URL with API-key auth. isLoopbackHostname now requires a literal IPv4 loopback address: four dot-separated parts, first part exactly 127, remaining parts decimal octets 0-255. localhost and [::1] handling is unchanged. Adds regression tests rejecting http://127.example.test and http://10.0.0.1 and accepting the inclusive boundary 127.255.255.255. --- src/api/providers/__tests__/gemini.spec.ts | 101 +++++++++++++++++++-- src/api/providers/gemini.ts | 18 +++- 2 files changed, 109 insertions(+), 10 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index eb3665f317..5671519a42 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -841,12 +841,11 @@ describe("GeminiHandler", () => { }) it("should reject hostnames that only resemble the 127. range", async () => { - // "127a.b" shares the 127 prefix but is not a loopback address and - // must not pass the anchored, dot-escaped 127. check. This case pins - // the dot escape itself: an unescaped /^127/ pattern would match - // "127a.b" and wrongly allow cleartext. (Hosts such as a127.0.0.1 - // are rejected by new URL() outright — its mixed digit-led/letter-led - // label rule — so they never reach the check.) + // "127a.b" is not a four-part 127.0.0.0/8 literal, so it must not + // pass the loopback check. A prefix-style "starts with 127." test + // would have matched it and wrongly allowed cleartext. (Hosts such + // as a127.0.0.1 are rejected by new URL() outright — its mixed + // digit-led/letter-led label rule — so they never reach the check.) const messages: Anthropic.Messages.MessageParam[] = [ { role: "user", @@ -906,6 +905,96 @@ describe("GeminiHandler", () => { expect(stub).not.toHaveBeenCalled() }) + it("should reject a public hostname whose first label is 127", async () => { + // "127.example.test" is a syntactically valid hostname (a digit-led + // first label is legal DNS) that new URL() accepts, but it is a + // public domain — not a literal 127.0.0.0/8 address — so cleartext + // must be rejected. A "starts with 127." prefix check would wrongly + // allow it and leak the API key in cleartext. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.example.test:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + + it("should allow the full 127.0.0.0/8 range including 127.255.255.255", async () => { + // The octet boundary is inclusive: 255 is the largest valid decimal + // octet and must be accepted as loopback. A < instead of <= (or a + // lowered boundary) would wrongly reject valid loopback endpoints. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const loopbackHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.255.255.255:8080", + }) + loopbackHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + await collectStream(loopbackHandler.createMessage("You are a helpful assistant", messages)) + + const config = stub.mock.calls[0][0].config + expect(config.httpOptions).toEqual({ baseUrl: "http://127.255.255.255:8080" }) + }) + + it("should reject a four-part IPv4 host outside the 127.0.0.0/8 range", async () => { + // "10.0.0.1" is a valid IPv4 literal that new URL() accepts, but its + // first octet is not 127, so it is not loopback and cleartext must + // be rejected. This pins the 127 label comparison itself. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://10.0.0.1:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + it("should reject a non-HTTP loopback scheme such as ftp://localhost", async () => { // The protocol operand (left of the && in the loopback exception) must // still be enforced: a non-HTTP scheme aimed at a loopback host is not diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 3d89c21cac..55bdfb777a 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -182,10 +182,20 @@ function sanitizeSchemaForGemini( // directly; `new URL` keeps the brackets in IPv6 hostnames, so `[::1]` is // the loopback host form to compare against. function isLoopbackHostname(hostname: string): boolean { - // The ^ anchor is observable: new URL() accepts non-loopback hostnames that - // contain the "127." substring (e.g. foo127.bar), so a de-anchored pattern - // would misclassify them as loopback and allow cleartext. - return hostname === "localhost" || hostname === "[::1]" || /^127\./.test(hostname) + if (hostname === "localhost" || hostname === "[::1]") { + return true + } + // Only literal IPv4 loopback (127.0.0.0/8) qualifies: public hostnames may + // start with a "127." label (e.g. 127.example.test), which a prefix test + // would misclassify as loopback and allow cleartext. + const parts = hostname.split(".") + if (parts.length !== 4 || parts[0] !== "127") { + return false + } + // The remaining parts must be decimal octets 0-255. Non-numeric parts + // (e.g. "example" in 127.example.test) fail the digit check, and Number() + // of a non-numeric string is NaN, which also fails the <= 255 check. + return parts.slice(1).every((octet) => /^\d{1,3}$/.test(octet) && Number(octet) <= 255) } // Throws an ApiProviderError when baseUrl is not HTTPS (loopback HTTP is the From 2f8b22881c2e31e59158d2b1f0e4e8e582405360 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 09:28:46 +0800 Subject: [PATCH 11/20] test(api): restructure Gemini loopback octet check for mutation observability Extract the octet predicate into isOctet() so the single-line 127.0.0.0/8 check has no every/slice method-chain mutants, and pin the remaining unobservable ConditionalExpression variants with true-rationale Stryker directives. Add reject tests for 127.0.0.a (non-numeric last octet), 127.0.a.b (non-numeric middle octet) and 127.0.0.1.a (five-part host). --- src/api/providers/__tests__/gemini.spec.ts | 101 ++++++++++++++++++++- src/api/providers/gemini.ts | 30 +++--- 2 files changed, 114 insertions(+), 17 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 5671519a42..b72cde15b7 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -875,9 +875,9 @@ describe("GeminiHandler", () => { it("should reject a valid hostname containing the 127. substring that is not loopback", async () => { // "foo127.bar" is accepted by new URL() (a syntactically valid - // hostname), contains the "127." substring, yet is not a loopback - // address: only the anchored 127. check rejects it. A de-anchored - // /127\./ pattern would match it and wrongly allow cleartext. + // hostname) and contains the "127." substring, yet is not a + // loopback address: its first label is "foo127", not the literal + // 127 octet, so the four-part 127.0.0.0/8 check rejects it. const messages: Anthropic.Messages.MessageParam[] = [ { role: "user", @@ -938,6 +938,101 @@ describe("GeminiHandler", () => { expect(stub).not.toHaveBeenCalled() }) + it("should reject a four-part host with a non-numeric last octet", async () => { + // "127.0.0.a" is accepted by new URL() (a non-digit last label + // skips IPv4 validation), but the last label is not a decimal + // octet, so the host is not loopback and cleartext must be + // rejected. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.0.0.a:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + + it("should reject a four-part host with a non-numeric middle octet", async () => { + // "127.0.a.b" is accepted by new URL(), but its third label is + // not a decimal octet, so the host is not loopback and cleartext + // must be rejected even though the first and fourth labels are + // numeric. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.0.a.b:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + + it("should reject a five-part host that starts with a loopback address", async () => { + // "127.0.0.1.a" is a hostname (not an IPv4 literal) that new + // URL() accepts: its five labels mean it must not pass the + // four-part 127.0.0.0/8 check, so cleartext must be rejected. + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + const stub = vi.fn().mockReturnValue((async function* () {})()) + const restrictedHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + googleGeminiBaseUrl: "http://127.0.0.1.a:8080", + }) + restrictedHandler["client"] = handler["client"] + handler["client"].models.generateContentStream = stub + + const error = await collectStream( + restrictedHandler.createMessage("You are a helpful assistant", messages), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(ApiProviderError) + expect((error as ApiProviderError).message).toBe( + "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", + ) + expect((error as ApiProviderError).provider).toBe("Gemini") + expect(stub).not.toHaveBeenCalled() + }) + it("should allow the full 127.0.0.0/8 range including 127.255.255.255", async () => { // The octet boundary is inclusive: 255 is the largest valid decimal // octet and must be accepted as loopback. A < instead of <= (or a diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 55bdfb777a..86b983fdf9 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -177,25 +177,27 @@ function sanitizeSchemaForGemini( // imported profiles). The @google/genai client keeps API-key authentication for // custom endpoints, so reject cleartext base URLs before any request — with a // narrow loopback exception for local test proxies. -// The caller (assertSecureGeminiBaseUrl) has already parsed the URL and only -// reaches this for the HTTP exception, so the parsed hostname is passed -// directly; `new URL` keeps the brackets in IPv6 hostnames, so `[::1]` is -// the loopback host form to compare against. +// A decimal octet is a 1-3 digit string; a non-numeric part also fails the +// range check, because Number() of a non-numeric string is NaN, which is not +// <= 255. +function isOctet(octet: string): boolean { + // Stryker disable next-line Regex,LogicalOperator,ConditionalExpression: a malformed octet newly accepted by a mutated regex or operand is non-numeric (Number() is NaN, failing <= 255), out of range (e.g. 256), or 4+ digits, all of which are unreachable because new URL() rejects the all-digit host; the ->false and stricter-regex variants are killed by the 127.0.0.1 and 127.255.255.255 accept tests + return /^\d{1,3}$/.test(octet) && Number(octet) <= 255 +} + +// Only literal IPv4 loopback (127.0.0.0/8) qualifies: public hostnames may +// start with a "127." label (e.g. 127.example.test), which a prefix test would +// misclassify as loopback and allow cleartext. assertSecureGeminiBaseUrl has +// already parsed the URL and only reaches this for the HTTP exception, so the +// parsed hostname is passed directly; `new URL` keeps the brackets in IPv6 +// hostnames, so `[::1]` is the loopback host form to compare against. function isLoopbackHostname(hostname: string): boolean { if (hostname === "localhost" || hostname === "[::1]") { return true } - // Only literal IPv4 loopback (127.0.0.0/8) qualifies: public hostnames may - // start with a "127." label (e.g. 127.example.test), which a prefix test - // would misclassify as loopback and allow cleartext. const parts = hostname.split(".") - if (parts.length !== 4 || parts[0] !== "127") { - return false - } - // The remaining parts must be decimal octets 0-255. Non-numeric parts - // (e.g. "example" in 127.example.test) fail the digit check, and Number() - // of a non-numeric string is NaN, which also fails the <= 255 check. - return parts.slice(1).every((octet) => /^\d{1,3}$/.test(octet) && Number(octet) <= 255) + // Stryker disable next-line ConditionalExpression: the ->true variants of isOctet(parts[1]) and isOctet(parts[2]) are unobservable because any four-part host with a non-numeric or out-of-range middle octet is rejected by new URL() before reaching this check; the remaining variants are killed by the 127.0.0.1, 127.255.255.255, 10.0.0.1, 127.0.0.a and 127.0.0.1.a tests + return parts.length === 4 && parts[0] === "127" && isOctet(parts[1]) && isOctet(parts[2]) && isOctet(parts[3]) } // Throws an ApiProviderError when baseUrl is not HTTPS (loopback HTTP is the From 4e9a602367cccf6d9265b5925e469d97686a068f Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Mon, 7 Sep 2026 10:02:26 +0800 Subject: [PATCH 12/20] test(api): collapse duplicated Gemini base-URL rejection tests into a table-driven case --- src/api/providers/__tests__/gemini.spec.ts | 130 +++------------------ 1 file changed, 13 insertions(+), 117 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index b72cde15b7..09b000f8f5 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -905,107 +905,34 @@ describe("GeminiHandler", () => { expect(stub).not.toHaveBeenCalled() }) - it("should reject a public hostname whose first label is 127", async () => { + const INSECURE_BASE_URL_CASES = [ // "127.example.test" is a syntactically valid hostname (a digit-led // first label is legal DNS) that new URL() accepts, but it is a // public domain — not a literal 127.0.0.0/8 address — so cleartext // must be rejected. A "starts with 127." prefix check would wrongly // allow it and leak the API key in cleartext. - const messages: Anthropic.Messages.MessageParam[] = [ - { - role: "user", - content: "Hello", - }, - ] - const stub = vi.fn().mockReturnValue((async function* () {})()) - const restrictedHandler = new GeminiHandler({ - apiKey: "test-key", - apiModelId: GEMINI_MODEL_NAME, - geminiApiKey: "test-key", - googleGeminiBaseUrl: "http://127.example.test:8080", - }) - restrictedHandler["client"] = handler["client"] - handler["client"].models.generateContentStream = stub - - const error = await collectStream( - restrictedHandler.createMessage("You are a helpful assistant", messages), - ).catch((e: unknown) => e) - expect(error).toBeInstanceOf(ApiProviderError) - expect((error as ApiProviderError).message).toBe( - "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", - ) - expect((error as ApiProviderError).provider).toBe("Gemini") - expect(stub).not.toHaveBeenCalled() - }) - - it("should reject a four-part host with a non-numeric last octet", async () => { + ["a public hostname whose first label is 127", "http://127.example.test:8080"], // "127.0.0.a" is accepted by new URL() (a non-digit last label // skips IPv4 validation), but the last label is not a decimal // octet, so the host is not loopback and cleartext must be // rejected. - const messages: Anthropic.Messages.MessageParam[] = [ - { - role: "user", - content: "Hello", - }, - ] - const stub = vi.fn().mockReturnValue((async function* () {})()) - const restrictedHandler = new GeminiHandler({ - apiKey: "test-key", - apiModelId: GEMINI_MODEL_NAME, - geminiApiKey: "test-key", - googleGeminiBaseUrl: "http://127.0.0.a:8080", - }) - restrictedHandler["client"] = handler["client"] - handler["client"].models.generateContentStream = stub - - const error = await collectStream( - restrictedHandler.createMessage("You are a helpful assistant", messages), - ).catch((e: unknown) => e) - expect(error).toBeInstanceOf(ApiProviderError) - expect((error as ApiProviderError).message).toBe( - "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", - ) - expect((error as ApiProviderError).provider).toBe("Gemini") - expect(stub).not.toHaveBeenCalled() - }) - - it("should reject a four-part host with a non-numeric middle octet", async () => { + ["a four-part host with a non-numeric last octet", "http://127.0.0.a:8080"], // "127.0.a.b" is accepted by new URL(), but its third label is // not a decimal octet, so the host is not loopback and cleartext // must be rejected even though the first and fourth labels are // numeric. - const messages: Anthropic.Messages.MessageParam[] = [ - { - role: "user", - content: "Hello", - }, - ] - const stub = vi.fn().mockReturnValue((async function* () {})()) - const restrictedHandler = new GeminiHandler({ - apiKey: "test-key", - apiModelId: GEMINI_MODEL_NAME, - geminiApiKey: "test-key", - googleGeminiBaseUrl: "http://127.0.a.b:8080", - }) - restrictedHandler["client"] = handler["client"] - handler["client"].models.generateContentStream = stub - - const error = await collectStream( - restrictedHandler.createMessage("You are a helpful assistant", messages), - ).catch((e: unknown) => e) - expect(error).toBeInstanceOf(ApiProviderError) - expect((error as ApiProviderError).message).toBe( - "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", - ) - expect((error as ApiProviderError).provider).toBe("Gemini") - expect(stub).not.toHaveBeenCalled() - }) - - it("should reject a five-part host that starts with a loopback address", async () => { + ["a four-part host with a non-numeric middle octet", "http://127.0.a.b:8080"], // "127.0.0.1.a" is a hostname (not an IPv4 literal) that new // URL() accepts: its five labels mean it must not pass the // four-part 127.0.0.0/8 check, so cleartext must be rejected. + ["a five-part host that starts with a loopback address", "http://127.0.0.1.a:8080"], + // "10.0.0.1" is a valid IPv4 literal that new URL() accepts, but its + // first octet is not 127, so it is not loopback and cleartext must + // be rejected. This pins the 127 label comparison itself. + ["a four-part IPv4 host outside the 127.0.0.0/8 range", "http://10.0.0.1:8080"], + ] as const + + it.each(INSECURE_BASE_URL_CASES)("should reject %s", async (_name, googleGeminiBaseUrl) => { const messages: Anthropic.Messages.MessageParam[] = [ { role: "user", @@ -1017,7 +944,7 @@ describe("GeminiHandler", () => { apiKey: "test-key", apiModelId: GEMINI_MODEL_NAME, geminiApiKey: "test-key", - googleGeminiBaseUrl: "http://127.0.0.1.a:8080", + googleGeminiBaseUrl, }) restrictedHandler["client"] = handler["client"] handler["client"].models.generateContentStream = stub @@ -1059,37 +986,6 @@ describe("GeminiHandler", () => { expect(config.httpOptions).toEqual({ baseUrl: "http://127.255.255.255:8080" }) }) - it("should reject a four-part IPv4 host outside the 127.0.0.0/8 range", async () => { - // "10.0.0.1" is a valid IPv4 literal that new URL() accepts, but its - // first octet is not 127, so it is not loopback and cleartext must - // be rejected. This pins the 127 label comparison itself. - const messages: Anthropic.Messages.MessageParam[] = [ - { - role: "user", - content: "Hello", - }, - ] - const stub = vi.fn().mockReturnValue((async function* () {})()) - const restrictedHandler = new GeminiHandler({ - apiKey: "test-key", - apiModelId: GEMINI_MODEL_NAME, - geminiApiKey: "test-key", - googleGeminiBaseUrl: "http://10.0.0.1:8080", - }) - restrictedHandler["client"] = handler["client"] - handler["client"].models.generateContentStream = stub - - const error = await collectStream( - restrictedHandler.createMessage("You are a helpful assistant", messages), - ).catch((e: unknown) => e) - expect(error).toBeInstanceOf(ApiProviderError) - expect((error as ApiProviderError).message).toBe( - "Google Gemini base URL must use HTTPS (or a loopback HTTP endpoint for local test proxies)", - ) - expect((error as ApiProviderError).provider).toBe("Gemini") - expect(stub).not.toHaveBeenCalled() - }) - it("should reject a non-HTTP loopback scheme such as ftp://localhost", async () => { // The protocol operand (left of the && in the loopback exception) must // still be enforced: a non-HTTP scheme aimed at a loopback host is not From b428b81328c4ba7b0535c33e94dfd5cf80256fae Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Fri, 11 Sep 2026 10:13:34 +0800 Subject: [PATCH 13/20] refactor(api): leverage abort-signal utils and RequestConfigBuilder for gemini/mistral/lite-llm Address review feedback (edelauna, #1303): - catch blocks: isRequestAborted(error, signal) catches SDK-native abort errors that surface before the signal flag propagates (error-name and exact-message branches), and createAbortError normalizes the thrown error to the series abort contract - createMessage bridge: RequestConfigBuilder.addMergedSignal (AbortSignal.any) replaces the manual AbortController + addEventListener/removeEventListener plumbing; the finally cleanup blocks are gone - pre-abort fast-fail uses the throwIfAborted helper - specs: message assertions updated to the helper messages, listener-identity assertions inverted to assert no manual listener management, and new regression tests pin the SDK-native-abort branch per provider --- src/api/providers/__tests__/gemini.spec.ts | 48 ++++++++++------ src/api/providers/__tests__/lite-llm.spec.ts | 35 +++++++----- src/api/providers/__tests__/mistral.spec.ts | 50 ++++++++++------ src/api/providers/gemini.ts | 53 +++++++---------- src/api/providers/lite-llm.ts | 56 +++++++----------- src/api/providers/mistral.ts | 60 ++++++++------------ 6 files changed, 151 insertions(+), 151 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index 09b000f8f5..b036e07eaf 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -478,9 +478,25 @@ describe("GeminiHandler", () => { const error = await handler .completePrompt("Test prompt", { abortSignal: controller.signal }) .catch((e: unknown) => e) - expect(error).toBeInstanceOf(DOMException) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + }) + + it("should surface a standard AbortError when the SDK throws an abort error before the signal flag propagates", async () => { + // The signal flag has not propagated yet, but the SDK rejects with its + // own abort error: isRequestAborted must catch the error-name branch. + const controller = new AbortController() + vi.mocked(handler["client"].models.generateContent).mockRejectedValue( + Object.assign(new Error("Request was aborted."), { name: "AbortError" }), + ) + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Gemini completion aborted") + expect((error as Error).message).toBe("The Gemini request was aborted") }) it("should surface the wrapped provider error when the request fails without options", async () => { @@ -1067,7 +1083,7 @@ describe("GeminiHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Gemini request aborted") + expect((error as Error).message).toBe("This operation was aborted") expect(handler["client"].models.generateContentStream).not.toHaveBeenCalled() }) @@ -1115,31 +1131,29 @@ describe("GeminiHandler", () => { }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Gemini request aborted") + expect((error as Error).message).toBe("The Gemini request was aborted") expect(capturedSignal).toBeDefined() - // The in-flight request must run against a request-local signal, not the - // external one forwarded by reference. + // The in-flight request must run against a request-local merged signal, not + // the external one forwarded by reference. expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) - // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. Target the last "abort" - // registration (the bridge's listener) and assert the exact reference - // so a bridge that removes a different callback cannot pass. - const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") - const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] - expect(typeof addedListener).toBe("function") - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) + // The bridge (RequestConfigBuilder.addMergedSignal) uses AbortSignal.any, + // so the external signal must never be managed with manual listeners. + expect(addEventListenerSpy).not.toHaveBeenCalled() + expect(removeEventListenerSpy).not.toHaveBeenCalled() }) - it("should not set config.abortSignal when no external signal is provided", async () => { + it("should set a live request-local config.abortSignal when no external signal is provided", async () => { const stub = vi.fn().mockReturnValue((async function* () {})()) handler["client"].models.generateContentStream = stub await collectStream(handler.createMessage("You are a helpful assistant", messages)) const config = stub.mock.calls[0][0].config - expect(config.abortSignal).toBeUndefined() + // The request-local controller signal is always present so the provider + // keeps its own abort handle; with no external signal it never aborts. + expect(config.abortSignal).toBeInstanceOf(AbortSignal) + expect(config.abortSignal?.aborted).toBe(false) }) it("should wrap a non-abort stream failure with the i18n message and capture telemetry", async () => { diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index ab910cc49a..e34dc428ee 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1432,9 +1432,23 @@ describe("LiteLLMHandler", () => { const error = await handler .completePrompt("test prompt", { abortSignal: controller.signal }) .catch((e: unknown) => e) - expect(error).toBeInstanceOf(DOMException) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("should surface a standard AbortError when the SDK throws an abort error before the signal flag propagates", async () => { + // The signal flag has not propagated yet, but the SDK rejects with its + // own abort error: isRequestAborted must catch the error-name branch. + mockCreate.mockRejectedValueOnce(Object.assign(new Error("Request was aborted."), { name: "AbortError" })) + const controller = new AbortController() + + const error = await handler + .completePrompt("test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("LiteLLM completion aborted") + expect((error as Error).message).toBe("The LiteLLM request was aborted") }) it("should surface the wrapped provider error when the request fails without options", async () => { @@ -1474,7 +1488,7 @@ describe("LiteLLMHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("LiteLLM streaming aborted") + expect((error as Error).message).toBe("This operation was aborted") expect(mockCreate).not.toHaveBeenCalled() }) @@ -1532,18 +1546,13 @@ describe("LiteLLMHandler", () => { }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("LiteLLM streaming aborted") + expect((error as Error).message).toBe("The LiteLLM request was aborted") expect(capturedSignal).toBeDefined() expect(capturedSignal?.aborted).toBe(true) - // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. Target the last "abort" - // registration (the bridge's listener) and assert the exact reference - // so a bridge that removes a different callback cannot pass. - const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") - const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] - expect(typeof addedListener).toBe("function") - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) + // The bridge (RequestConfigBuilder.addMergedSignal) uses AbortSignal.any, + // so the external signal must never be managed with manual listeners. + expect(addEventListenerSpy).not.toHaveBeenCalled() + expect(removeEventListenerSpy).not.toHaveBeenCalled() }) it("should wrap a non-abort stream failure with the i18n-free provider message and no metadata", async () => { diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index e1ce3185cb..aa96934b4a 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -135,6 +135,7 @@ describe("MistralHandler", () => { tools: expect.any(Array), toolChoice: "any", }), + expect.objectContaining({ fetchOptions: { signal: expect.any(AbortSignal) } }), ) expect(result.value).toBeDefined() @@ -501,6 +502,7 @@ describe("MistralHandler", () => { ]), toolChoice: "any", }), + expect.objectContaining({ fetchOptions: { signal: expect.any(AbortSignal) } }), ) }) @@ -518,6 +520,7 @@ describe("MistralHandler", () => { tools: expect.any(Array), toolChoice: "any", }), + expect.objectContaining({ fetchOptions: { signal: expect.any(AbortSignal) } }), ) }) @@ -655,6 +658,7 @@ describe("MistralHandler", () => { expect.objectContaining({ toolChoice: "any", }), + expect.objectContaining({ fetchOptions: { signal: expect.any(AbortSignal) } }), ) }) }) @@ -799,9 +803,23 @@ describe("MistralHandler", () => { const error = await handler .completePrompt("Test prompt", { abortSignal: controller.signal }) .catch((e: unknown) => e) - expect(error).toBeInstanceOf(DOMException) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + }) + + it("should surface a standard AbortError when the SDK throws an abort error before the signal flag propagates", async () => { + // The signal flag has not propagated yet, but the SDK rejects with its + // own abort error: isRequestAborted must catch the error-name branch. + mockComplete.mockRejectedValueOnce(Object.assign(new Error("Request was aborted."), { name: "AbortError" })) + const controller = new AbortController() + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Mistral completion aborted") + expect((error as Error).message).toBe("The Mistral request was aborted") }) }) @@ -827,7 +845,7 @@ describe("MistralHandler", () => { const error = await collectStream(stream).catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Mistral completion aborted") + expect((error as Error).message).toBe("This operation was aborted") expect(mockCreate).not.toHaveBeenCalled() }) @@ -883,21 +901,16 @@ describe("MistralHandler", () => { }) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("Mistral completion aborted") + expect((error as Error).message).toBe("The Mistral request was aborted") expect(capturedSignal).toBeDefined() - // The in-flight stream must run against a request-local signal, not the - // external one forwarded by reference. + // The in-flight stream must run against a request-local merged signal, not + // the external one forwarded by reference. expect(capturedSignal).not.toBe(controller.signal) expect(capturedSignal?.aborted).toBe(true) - // The bridge registers a once-only listener on the external signal and - // detaches it when the request settles. Target the last "abort" - // registration (the bridge's listener) and assert the exact reference - // so a bridge that removes a different callback cannot pass. - const abortAddCalls = addEventListenerSpy.mock.calls.filter(([event]) => event === "abort") - const addedListener = abortAddCalls[abortAddCalls.length - 1]?.[1] - expect(typeof addedListener).toBe("function") - expect(addEventListenerSpy).toHaveBeenCalledWith("abort", addedListener, { once: true }) - expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", addedListener) + // The bridge (RequestConfigBuilder.addMergedSignal) uses AbortSignal.any, + // so the external signal must never be managed with manual listeners. + expect(addEventListenerSpy).not.toHaveBeenCalled() + expect(removeEventListenerSpy).not.toHaveBeenCalled() }) it("should wrap a non-abort stream failure when metadata is provided without a signal", async () => { @@ -910,11 +923,14 @@ describe("MistralHandler", () => { expect((error as Error).message).toBe("Mistral completion error: boom") }) - it("should not pass a signal to the stream call when no external signal is provided", async () => { + it("should pass a live request-local signal to the stream call when no external signal is provided", async () => { const stream = handler.createMessage(systemPrompt, messages) await collectStream(stream) const streamOptions = mockCreate.mock.calls[0][1] - expect(streamOptions).toBeUndefined() + // The request-local controller signal is always present so the provider + // keeps its own abort handle; with no external signal it never aborts. + expect(streamOptions?.fetchOptions?.signal).toBeInstanceOf(AbortSignal) + expect(streamOptions?.fetchOptions?.signal?.aborted).toBe(false) }) }) }) diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 86b983fdf9..71209e1c52 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -27,6 +27,8 @@ import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, Complete import { BaseProvider } from "./base-provider" import { NOT_PROVIDED } from "./constants" import { parseVertexJsonCredentials } from "./utils/vertex-credentials" +import { createAbortError, isRequestAborted, throwIfAborted } from "./utils/abort-signal" +import { RequestConfigBuilder } from "./config-builder/request-config-builder" import { getRequestTimeoutMs } from "./utils/request-timeout" type GeminiHandlerOptions = ApiHandlerOptions & { @@ -404,29 +406,21 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } } - // Bridge the external abort signal from Task (metadata.abortSignal) into a - // request-local controller so the in-flight generateContentStream request - // can be cancelled. The @google/genai SDK merges this signal with its own - // timeout handling, which is preserved rather than replaced. - // A pre-aborted signal rejects immediately with an AbortError. - const externalAbortSignal = metadata?.abortSignal - let requestAbortController: AbortController | undefined - let externalAbortListener: (() => void) | undefined - if (externalAbortSignal) { - if (externalAbortSignal.aborted) { - throw new DOMException("Gemini request aborted", "AbortError") - } - const controller = new AbortController() - requestAbortController = controller - const onExternalAbort = () => controller.abort() - externalAbortListener = onExternalAbort - externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) - } + // Fast-fail if the request was already aborted before building. + throwIfAborted(metadata?.abortSignal) + + // The request-local controller is the provider-owned abort handle; the + // signal the SDK receives merges it with the external Task signal + // (AbortSignal.any inside RequestConfigBuilder), so external aborts + // cancel the in-flight request without manual listener management. + const requestBuilder = new RequestConfigBuilder<{ signal?: AbortSignal }>() + requestBuilder.addMergedSignal(new AbortController(), metadata) + const requestSignal = requestBuilder.getOption("signal") const params: GenerateContentParameters = { model, contents, - config: requestAbortController ? { ...config, abortSignal: requestAbortController.signal } : config, + config: requestSignal ? { ...config, abortSignal: requestSignal } : config, } try { @@ -560,10 +554,10 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } } } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (metadata?.abortSignal?.aborted) { - throw new DOMException("Gemini request aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, metadata?.abortSignal)) { + throw createAbortError("Gemini") } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "createMessage") @@ -574,11 +568,6 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } throw error - } finally { - // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here - if (externalAbortSignal && externalAbortListener) { - externalAbortSignal.removeEventListener("abort", externalAbortListener) - } } } @@ -723,10 +712,10 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl return text } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (options?.abortSignal?.aborted) { - throw new DOMException("Gemini completion aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, options?.abortSignal)) { + throw createAbortError("Gemini") } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "completePrompt") diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index 7ddd687794..fd9482362f 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -22,6 +22,8 @@ import { sanitizeOpenAiCallId } from "../../utils/tool-id" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" import { RouterProvider } from "./router-provider" import { extractReasoningFromDelta } from "./utils/extract-reasoning" +import { createAbortError, isRequestAborted, throwIfAborted } from "./utils/abort-signal" +import { RequestConfigBuilder } from "./config-builder/request-config-builder" import { getRequestTimeoutMs } from "./utils/request-timeout" /** @@ -268,31 +270,20 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa requestHeaders["X-Zoo-Session-ID"] = metadata.taskId } - // Bridge the external abort signal from Task (metadata.abortSignal) into a - // request-local controller so the in-flight streaming request can be - // cancelled. A pre-aborted signal rejects immediately with an AbortError. - const externalAbortSignal = metadata?.abortSignal - let requestAbortController: AbortController | undefined - let externalAbortListener: (() => void) | undefined - if (externalAbortSignal) { - if (externalAbortSignal.aborted) { - throw new DOMException("LiteLLM streaming aborted", "AbortError") - } - const controller = new AbortController() - requestAbortController = controller - const onExternalAbort = () => controller.abort() - externalAbortListener = onExternalAbort - externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) - } + // Fast-fail if the request was already aborted before building. + throwIfAborted(metadata?.abortSignal) + + // The request-local controller is the provider-owned abort handle; the + // request signal merges it with the external Task signal (AbortSignal.any + // inside RequestConfigBuilder), so external aborts cancel the in-flight + // request without manual listener management. + const requestBuilder = new RequestConfigBuilder<{ signal?: AbortSignal }>() + requestBuilder.addMergedSignal(new AbortController(), metadata) + const requestSignal = requestBuilder.getOption("signal") try { const { data: completion } = await this.client.chat.completions - .create( - requestOptions, - requestAbortController - ? { headers: requestHeaders, signal: requestAbortController.signal } - : { headers: requestHeaders }, - ) + .create(requestOptions, { headers: requestHeaders, signal: requestSignal }) .withResponse() let lastUsage @@ -362,20 +353,15 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa yield usageData } } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (metadata?.abortSignal?.aborted) { - throw new DOMException("LiteLLM streaming aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, metadata?.abortSignal)) { + throw createAbortError("LiteLLM") } if (error instanceof Error) { throw new Error(`LiteLLM streaming error: ${error.message}`) } throw error - } finally { - // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here - if (externalAbortSignal && externalAbortListener) { - externalAbortSignal.removeEventListener("abort", externalAbortListener) - } } } @@ -425,10 +411,10 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa ) return response.choices[0]?.message.content || "" } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (options?.abortSignal?.aborted) { - throw new DOMException("LiteLLM completion aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, options?.abortSignal)) { + throw createAbortError("LiteLLM") } if (error instanceof Error) { throw new Error(`LiteLLM completion error: ${error.message}`) diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index 36514ea284..260df0326c 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -16,7 +16,8 @@ import { ApiHandlerOptions } from "../../shared/api" import { convertToMistralMessages } from "../transform/mistral-format" import { ApiStream } from "../transform/stream" import { handleProviderError } from "./utils/error-handler" -import { mergeAbortSignalAndTimeout } from "./utils/abort-signal" +import { createAbortError, isRequestAborted, mergeAbortSignalAndTimeout, throwIfAborted } from "./utils/abort-signal" +import { RequestConfigBuilder } from "./config-builder/request-config-builder" import { BaseProvider } from "./base-provider" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" @@ -102,32 +103,22 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand // Temporary debug log for QA // console.log("[MISTRAL DEBUG] Raw API request body:", requestOptions) - // Bridge the external abort signal from Task (metadata.abortSignal) into a - // request-local controller so the in-flight streaming request can be - // cancelled. A pre-aborted signal rejects immediately with an AbortError. - const externalAbortSignal = metadata?.abortSignal - let requestAbortController: AbortController | undefined - let externalAbortListener: (() => void) | undefined - if (externalAbortSignal) { - if (externalAbortSignal.aborted) { - throw new DOMException("Mistral completion aborted", "AbortError") - } - const controller = new AbortController() - requestAbortController = controller - const onExternalAbort = () => controller.abort() - externalAbortListener = onExternalAbort - externalAbortSignal.addEventListener("abort", onExternalAbort, { once: true }) - } + // Fast-fail if the request was already aborted before building. + throwIfAborted(metadata?.abortSignal) + + // The request-local controller is the provider-owned abort handle; the + // fetch signal merges it with the external Task signal (AbortSignal.any + // inside RequestConfigBuilder), so external aborts cancel the in-flight + // request without manual listener management. + const requestBuilder = new RequestConfigBuilder<{ signal?: AbortSignal }>() + requestBuilder.addMergedSignal(new AbortController(), metadata) + const requestSignal = requestBuilder.getOption("signal") let response try { - if (requestAbortController) { - response = await this.client.chat.stream(requestOptions, { - fetchOptions: { signal: requestAbortController.signal }, - }) - } else { - response = await this.client.chat.stream(requestOptions) - } + response = await this.client.chat.stream(requestOptions, { + fetchOptions: { signal: requestSignal }, + }) for await (const event of response) { const delta = event.data.choices[0]?.delta @@ -181,20 +172,15 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand } } } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (metadata?.abortSignal?.aborted) { - throw new DOMException("Mistral completion aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, metadata?.abortSignal)) { + throw createAbortError("Mistral") } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "createMessage") TelemetryService.instance.captureException(apiError) throw new Error(`Mistral completion error: ${errorMessage}`) - } finally { - // Stryker disable next-line LogicalOperator: externalAbortListener is only assigned when externalAbortSignal is truthy, so && and || evaluate identically here - if (externalAbortSignal && externalAbortListener) { - externalAbortSignal.removeEventListener("abort", externalAbortListener) - } } } @@ -260,10 +246,10 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand return content || "" } catch (error) { - // User-initiated abort: surface a standard AbortError rather than the - // wrapped provider error so callers can distinguish cancellation. - if (options?.abortSignal?.aborted) { - throw new DOMException("Mistral completion aborted", "AbortError") + // Aborted request: covers both the external signal and SDK-native + // abort errors that may surface before the signal flag propagates. + if (isRequestAborted(error, options?.abortSignal)) { + throw createAbortError("Mistral") } const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model, "completePrompt") From ffedfdbdb593a4ce519b71844c9d24f9e4f63a2c Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Fri, 11 Sep 2026 13:27:34 +0800 Subject: [PATCH 14/20] fix(api): run LiteLLM pre-abort fast-fail before model discovery Address CodeRabbit finding on #1303 (review 5174654329): LiteLLMHandler.createMessage awaited fetchModel() before throwIfAborted, so an already-aborted request could reach provider model discovery (getModels/refreshModels) on a cold cache, and a discovery failure would escape as a model-fetch error instead of AbortError. completePrompt had the same shape. Both methods now fast-fail with the canonical AbortError before model discovery begins. Tests: the completePrompt aborted-signal case is rewritten as a true in-flight abort (still covering catch-block classification via signal.aborted, without leaking a once-mock), a new completePrompt pre-abort test added, and the createMessage pre-abort test now pins that fetchModel is never called for an already-aborted request. --- src/api/providers/__tests__/lite-llm.spec.ts | 52 +++++++++++++++++--- src/api/providers/lite-llm.ts | 13 +++-- 2 files changed, 56 insertions(+), 9 deletions(-) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index e34dc428ee..e68f8ecf59 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1393,6 +1393,26 @@ describe("LiteLLMHandler", () => { ) }) + it("should reject immediately with AbortError when the signal is pre-aborted, before model discovery", async () => { + const fetchModelSpy = vi.spyOn(handler, "fetchModel").mockResolvedValue({ + id: litellmDefaultModelId, + info: litellmDefaultModelInfo, + }) + const controller = new AbortController() + controller.abort() + + const error = await handler + .completePrompt("test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("This operation was aborted") + // The pre-abort fast-fail must run before provider model discovery: + // an already-aborted request must not trigger getModels/refreshModels. + expect(fetchModelSpy).not.toHaveBeenCalled() + expect(mockCreate).not.toHaveBeenCalled() + }) + it("should pass timeout through to client", async () => { mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }] }) await handler.completePrompt("test prompt", { timeoutMs: 5000 }) @@ -1424,14 +1444,27 @@ describe("LiteLLMHandler", () => { expect(mockCreate).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) - it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { - mockCreate.mockRejectedValueOnce(new Error("LiteLLM API error")) + it("should surface a standard AbortError when the signal is aborted while the request is in flight", async () => { const controller = new AbortController() + // The request stays pending until the external signal aborts it; the + // mock rejects with a plain (non-abort) error once the signal has + // aborted, so the catch block must classify it via signal.aborted. + mockCreate.mockImplementationOnce((_body: unknown, options?: { signal?: AbortSignal }) => { + const signal = options?.signal + return new Promise((_resolve, reject) => { + const fail = () => reject(new Error("LiteLLM API error")) + if (signal?.aborted) { + fail() + return + } + signal?.addEventListener("abort", fail, { once: true }) + }) + }) + const promise = handler.completePrompt("test prompt", { abortSignal: controller.signal }) + // Abort while the (pending) request is in flight — after entry, so the + // pre-abort fast-fail does not apply and the catch block handles it. controller.abort() - - const error = await handler - .completePrompt("test prompt", { abortSignal: controller.signal }) - .catch((e: unknown) => e) + const error = await promise.catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("The LiteLLM request was aborted") @@ -1476,6 +1509,10 @@ describe("LiteLLMHandler", () => { ] it("should reject immediately with AbortError when the external signal is pre-aborted", async () => { + const fetchModelSpy = vi.spyOn(handler, "fetchModel").mockResolvedValue({ + id: litellmDefaultModelId, + info: litellmDefaultModelInfo, + }) const controller = new AbortController() controller.abort() @@ -1489,6 +1526,9 @@ describe("LiteLLMHandler", () => { expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("This operation was aborted") + // The pre-abort fast-fail must run before provider model discovery: + // an already-aborted request must not trigger getModels/refreshModels. + expect(fetchModelSpy).not.toHaveBeenCalled() expect(mockCreate).not.toHaveBeenCalled() }) diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index fd9482362f..2acad4d9c9 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -136,6 +136,11 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa messages: Anthropic.Messages.MessageParam[], metadata?: ApiHandlerCreateMessageMetadata, ): ApiStream { + // Fast-fail if the request was already aborted before building, so an + // already-aborted request fails with AbortError before provider model + // discovery (getModels/refreshModels) begins. + throwIfAborted(metadata?.abortSignal) + const { id: modelId, info } = await this.fetchModel() // Models that require reasoning_content to be echoed back during tool-call @@ -270,9 +275,6 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa requestHeaders["X-Zoo-Session-ID"] = metadata.taskId } - // Fast-fail if the request was already aborted before building. - throwIfAborted(metadata?.abortSignal) - // The request-local controller is the provider-owned abort handle; the // request signal merges it with the external Task signal (AbortSignal.any // inside RequestConfigBuilder), so external aborts cancel the in-flight @@ -366,6 +368,11 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa } async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { + // Fast-fail if the request was already aborted before building, so an + // already-aborted request fails with AbortError before provider model + // discovery (getModels/refreshModels) begins. + throwIfAborted(options?.abortSignal) + const { id: modelId, info } = await this.fetchModel() // Check if this is a GPT-5 model that requires max_completion_tokens instead of max_tokens From d3e510009a39bebd709feab4889302c78782b286 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Fri, 11 Sep 2026 14:16:17 +0800 Subject: [PATCH 15/20] fix(api): stop provider streams after abort instead of leaking buffered content Class h of the abort-signal contract (streaming yield granularity): a for-await loop keeps pulling and yielding buffered content after the request is aborted, and the SDK swallows the mid-stream AbortError, so an aborted request could finish as a normal stream completion. createMessage in gemini/mistral/lite-llm now: - breaks at the top of the loop once the request-local signal is aborted (for-await pulls the already-buffered element before the check runs, so the break prevents processing, not pulling); - re-checks before every yield site reachable after a suspension point (a guard before the first yield of an iteration is dead code - the top check runs without a suspension point before it - and is intentionally absent); - throws the canonical AbortError after the loop, so a break or a swallowed mid-stream abort surfaces as "The request was aborted" instead of a normal completion. completePrompt in gemini/mistral additionally fast-fails before building the request (lite-llm received that in the previous commit). Tests: new "streaming loop abort defense" suites per provider (per-yield-site mid-chunk aborts plus a structural break/post-loop case with a pull counter, using the unguarded first-yield shape so only the break can prevent a leak), completePrompt pre-abort fast-fail tests, and the pre-aborted catch-path cases rewritten as true in-flight aborts (the mock stays pending until the external signal aborts - no once-mock leaks). The lite-llm in-flight mock attaches a no-op rejection handler to its pending element, which the top-of-loop break correctly never reads. Extra kill tests for lines swept into the diff hunks: responseId capture, non-Error throw rethrow/wrap branches (createMessage and completePrompt), supportsTemperature true/false temperature config, and a usage-only empty-choices chunk. Dead write-only hasContent/hasReasoning flags in gemini createMessage removed (they generated equivalent BooleanLiteral mutants inside the diff hunks). Mutation gate: three equivalent OptionalChaining mutants on the top-of-loop checks carry mutator-specific Stryker directives (requestSignal is always set by the addMergedSignal call two lines above - a request-local controller signal exists even without an external signal). Verified: tsc clean; 209/209 across the six delta suites; eslint --prune-suppressions with flat counts; zdt mutation preflight 204 valid / 204 killed / 0 survived (exit 0); zdt align check exit 0 (0 ERROR; 3 WARN = statically unguarded first-yield sites subsumed by the loop-top check; 2 DIVERGENCE = recorded per-method mechanism divergence). --- src/api/providers/__tests__/gemini.spec.ts | 321 ++++++++++++++++++- src/api/providers/__tests__/lite-llm.spec.ts | 133 +++++++- src/api/providers/__tests__/mistral.spec.ts | 168 +++++++++- src/api/providers/gemini.ts | 30 +- src/api/providers/lite-llm.ts | 17 + src/api/providers/mistral.ts | 23 +- 6 files changed, 665 insertions(+), 27 deletions(-) diff --git a/src/api/providers/__tests__/gemini.spec.ts b/src/api/providers/__tests__/gemini.spec.ts index b036e07eaf..28a9c8c056 100644 --- a/src/api/providers/__tests__/gemini.spec.ts +++ b/src/api/providers/__tests__/gemini.spec.ts @@ -407,6 +407,20 @@ describe("GeminiHandler", () => { }) }) + it("should reject immediately with AbortError when the external signal is pre-aborted, before the request is built", async () => { + const controller = new AbortController() + controller.abort() + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("This operation was aborted") + // The fast-fail must run before the client is touched. + expect(handler["client"].models.generateContent).not.toHaveBeenCalled() + }) + it("should work without options (backward compatible)", async () => { vi.mocked(handler["client"].models.generateContent).mockResolvedValue( stubGenerateContentResponse("response"), @@ -470,19 +484,55 @@ describe("GeminiHandler", () => { }) }) - it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { + it("should surface a standard AbortError when the signal is aborted while the request is in flight", async () => { const controller = new AbortController() + // The request stays pending until the external signal aborts it; the + // mock rejects with a plain (non-abort) error once the signal has + // aborted, so the catch block must classify it via signal.aborted. + vi.mocked(handler["client"].models.generateContent).mockImplementationOnce(() => { + return new Promise((_resolve, reject) => { + const fail = () => reject(new Error("Gemini API error")) + if (controller.signal.aborted) { + fail() + return + } + controller.signal.addEventListener("abort", fail, { once: true }) + }) + }) + const promise = handler.completePrompt("Test prompt", { abortSignal: controller.signal }) + // Abort after entry — the pre-abort fast-fail must not apply, and the + // catch block handles the failure. controller.abort() - vi.mocked(handler["client"].models.generateContent).mockRejectedValue(new Error("Gemini API error")) - - const error = await handler - .completePrompt("Test prompt", { abortSignal: controller.signal }) - .catch((e: unknown) => e) + const error = await promise.catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("The Gemini request was aborted") }) + it("rethrows non-Error completion failures as-is after capturing telemetry", async () => { + // The instanceof branch only covers Error instances; a non-Error rejection + // must be rethrown untouched (wrapping it would lose the original value). + vi.mocked(handler["client"].models.generateContent).mockRejectedValueOnce("Gemini: not an error object") + + const error = await handler.completePrompt("Test prompt").catch((e: unknown) => e) + expect(error).toBe("Gemini: not an error object") + expect(error).not.toBeInstanceOf(Error) + }) + + it("wraps completion failures in a new error before rethrowing", async () => { + // The consumer-facing error must be the provider's wrap (a NEW Error + // built from the localized template), not the raw SDK error instance: + // this is what the instanceof branch is for (and what makes the branch + // observable to the mutation gate). Unit tests cannot assert the i18n + // string itself (i18next is not initialized with resources here). + const mockError = new Error("Gemini completion failure") + vi.mocked(handler["client"].models.generateContent).mockRejectedValueOnce(mockError) + + const error = await handler.completePrompt("Test prompt").catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect(error).not.toBe(mockError) + }) + it("should surface a standard AbortError when the SDK throws an abort error before the signal flag propagates", async () => { // The signal flag has not propagated yet, but the SDK rejects with its // own abort error: isRequestAborted must catch the error-name branch. @@ -653,6 +703,70 @@ describe("GeminiHandler", () => { }) describe("completePrompt request options", () => { + it("applies the explicit model temperature when the model supports temperature", async () => { + const temperatureHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + modelTemperature: 0.5, + }) + temperatureHandler["client"] = handler["client"] + const baseInfo = temperatureHandler.getModel() + // Explicit supportsTemperature: true makes the `!== false` check observable — + // with the default model info (undefined) the literal cannot be mutated. + vi.spyOn(temperatureHandler, "getModel").mockReturnValue({ + ...baseInfo, + info: { + ...baseInfo.info, + supportsTemperature: true, + defaultTemperature: 1, + }, + }) + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + await temperatureHandler.completePrompt("Test prompt") + + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ temperature: 0.5 }), + }), + ) + }) + + it("ignores the explicit model temperature when the model does not support temperature", async () => { + const noTemperatureHandler = new GeminiHandler({ + apiKey: "test-key", + apiModelId: GEMINI_MODEL_NAME, + geminiApiKey: "test-key", + modelTemperature: 0.5, + }) + noTemperatureHandler["client"] = handler["client"] + const baseInfo = noTemperatureHandler.getModel() + vi.spyOn(noTemperatureHandler, "getModel").mockReturnValue({ + ...baseInfo, + info: { + ...baseInfo.info, + supportsTemperature: false, + defaultTemperature: 0.3, + }, + }) + vi.mocked(handler["client"].models.generateContent).mockResolvedValue( + stubGenerateContentResponse("Response"), + ) + + await noTemperatureHandler.completePrompt("Test prompt") + + // supportsTemperature: false forces the fallback to the model default — + // the user's explicit modelTemperature must not be applied. + expect(handler["client"].models.generateContent).toHaveBeenCalledWith( + expect.objectContaining({ + config: expect.objectContaining({ temperature: 0.3 }), + }), + ) + }) + it("should pass timeout and baseUrl through httpOptions", async () => { const handlerWithBaseUrl = new GeminiHandler({ apiKey: "test-key", @@ -1225,6 +1339,46 @@ describe("GeminiHandler", () => { expect(capturedError).toBeInstanceOf(ApiProviderError) }) + it("rethrows non-Error stream failures as-is after capturing telemetry", async () => { + // The instanceof branch only covers Error instances; a non-Error throw + // must be rethrown untouched (wrapping it would lose the original value). + handler["client"].models.generateContentStream = vi.fn().mockReturnValue( + (async function* () { + yield { candidates: [] } + throw "Gemini stream: not an error object" + })(), + ) + + const error = await collectStream(handler.createMessage(systemPrompt, mockMessages)).catch( + (e: unknown) => e, + ) + expect(error).toBe("Gemini stream: not an error object") + expect(error).not.toBeInstanceOf(Error) + // Telemetry still captures the stringified failure. + expect(mockCaptureException).toHaveBeenCalledWith( + expect.objectContaining({ + message: "Gemini stream: not an error object", + operation: "createMessage", + }), + ) + }) + + it("wraps stream failures in a new error before rethrowing", async () => { + // The consumer-facing error must be the provider's wrap (a NEW Error + // built from the localized template), not the raw SDK error instance: + // this is what the instanceof branch is for (and what makes the branch + // observable to the mutation gate). Unit tests cannot assert the i18n + // string itself (i18next is not initialized with resources here). + const mockError = new Error("Gemini stream failure") + handler["client"].models.generateContentStream = vi.fn().mockRejectedValue(mockError) + + const error = await collectStream(handler.createMessage(systemPrompt, mockMessages)).catch( + (e: unknown) => e, + ) + expect(error).toBeInstanceOf(Error) + expect(error).not.toBe(mockError) + }) + it("should capture telemetry on completePrompt error", async () => { const mockError = new Error("Gemini completion error") ;(handler["client"].models.generateContent as any).mockRejectedValue(mockError) @@ -1260,4 +1414,159 @@ describe("GeminiHandler", () => { expect(mockCaptureException).toHaveBeenCalled() }) }) + + describe("createMessage responseId capture", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + it("captures the final response's responseId so api history can store it", async () => { + handler["client"].models.generateContentStream = vi.fn().mockReturnValue( + asyncStreamFrom([ + { + candidates: [ + { + content: { parts: [{ text: "final" }] }, + finishReason: "STOP", + }, + ], + responseId: "resp-123", + }, + ]), + ) + + await collectStream(handler.createMessage("You are a helpful assistant", messages)) + + expect(handler["lastResponseId"]).toBe("resp-123") + }) + }) + + describe("createMessage streaming loop abort defense", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + function makePartsChunk(parts: Array>) { + return { + candidates: [ + { + content: { parts }, + }, + ], + } + } + + function startStream(signal: AbortSignal, generator: AsyncGenerator) { + handler["client"].models.generateContentStream = vi.fn().mockReturnValue(generator) + return handler.createMessage( + "You are a helpful assistant", + messages, + makeCreateMessageMetadata({ abortSignal: signal }), + ) + } + + it("rejects with AbortError instead of yielding text content after a mid-chunk abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([makePartsChunk([{ thought: true, text: "reasoning" }, { text: "content" }])]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("reasoning") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + }) + + it("rejects with AbortError instead of yielding reasoning after a mid-chunk abort (text-first chunk)", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([makePartsChunk([{ text: "content" }, { thought: true, text: "reasoning" }])]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + }) + + it("rejects with AbortError instead of yielding a tool-call name after a mid-chunk abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([makePartsChunk([{ text: "content" }, { functionCall: { name: "f", args: {} } }])]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + }) + + it("rejects with AbortError instead of yielding tool-call arguments after a mid-chunk abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([ + makePartsChunk([ + { functionCall: { name: "f", args: {} } }, + { functionCall: { name: "g", args: {} } }, + ]), + ]), + ) + const name = await stream.next() + expect(name.value?.type).toBe("tool_call_partial") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + }) + + it("does not process or pull chunks beyond the aborted point and surfaces AbortError instead of completing the stream", async () => { + const controller = new AbortController() + let pulls = 0 + const generator = (async function* () { + pulls++ + yield makePartsChunk([{ text: "first chunk" }]) + pulls++ + // Fallback shape (no candidates): its text yield is the first yield + // of the iteration and carries no pre-yield guard, so only the + // top-of-loop break prevents it from leaking. + yield { text: "buffered chunk" } + pulls++ + yield makePartsChunk([{ text: "third chunk" }]) + })() + const stream = startStream(controller.signal, generator) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Gemini request was aborted") + // The for-await mechanism pulls the already-buffered chunk before the + // top-of-loop check runs; the break must ensure it is not processed and + // that no chunk beyond it is pulled. + expect(pulls).toBe(2) + }) + }) }) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index e68f8ecf59..cc16ded6c0 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1548,19 +1548,24 @@ describe("LiteLLMHandler", () => { mockCreate.mockImplementationOnce((_body: unknown, options?: { signal?: AbortSignal }) => { capturedSignal = options?.signal requestStartedResolve() + const pendingAbort = new Promise((_resolve, reject) => { + const onAbort = () => reject(new DOMException("aborted", "AbortError")) + if (capturedSignal?.aborted) { + onAbort() + return + } + capturedSignal?.addEventListener("abort", onAbort, { once: true }) + }) + // A correctly aborted request stops pulling at the top-of-loop break + // and never reads this element, so its rejection must be handled + // here to keep it from being reported as an unhandled rejection. + void pendingAbort.catch(() => undefined) const mockStream = asyncStreamFrom([ { choices: [{ delta: { content: "partial" } }], usage: undefined, }, - new Promise((_resolve, reject) => { - const onAbort = () => reject(new DOMException("aborted", "AbortError")) - if (capturedSignal?.aborted) { - onAbort() - return - } - capturedSignal?.addEventListener("abort", onAbort, { once: true }) - }), + pendingAbort, ]) return { withResponse: vi.fn().mockResolvedValue({ data: mockStream }) } }) @@ -1618,4 +1623,116 @@ describe("LiteLLMHandler", () => { expect((error as Error).message).toBe("LiteLLM streaming error: boom") }) }) + + describe("createMessage streaming loop abort defense", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + function chunkOf(delta: Record) { + return { choices: [{ delta }] } + } + + function startStream(signal: AbortSignal, generator: AsyncGenerator) { + mockCreate.mockReturnValue({ + withResponse: vi.fn().mockResolvedValue({ data: generator }), + }) + return handler.createMessage("system", messages, makeCreateMessageMetadata({ abortSignal: signal })) + } + + it("rejects with AbortError instead of yielding text content after a mid-chunk abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([chunkOf({ reasoning_content: "reasoning", content: "content" })]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("reasoning") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("rejects with AbortError instead of yielding a tool call after a mid-chunk abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([ + chunkOf({ + content: "content", + tool_calls: [ + { + index: 0, + id: "1", + type: "function", + function: { name: "f", arguments: "{}" }, + }, + ], + }), + ]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("does not process or pull chunks beyond the aborted point and surfaces AbortError instead of completing the stream", async () => { + const controller = new AbortController() + let pulls = 0 + const generator = (async function* () { + pulls++ + yield chunkOf({ content: "first" }) + pulls++ + // Reasoning shape: its reasoning yield is the first yield of the + // iteration and carries no pre-yield guard, so only the top-of-loop + // break prevents it from leaking. + yield chunkOf({ reasoning_content: "buffered" }) + pulls++ + yield chunkOf({ content: "third" }) + })() + const stream = startStream(controller.signal, generator) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + // The for-await mechanism pulls the already-buffered chunk before the + // top-of-loop check runs; the break must ensure it is not processed and + // that no chunk beyond it is pulled. + expect(pulls).toBe(2) + }) + + it("completes normally for a usage-only chunk with empty choices (delta is undefined)", async () => { + const controller = new AbortController() + // Real OpenAI-shaped streams end with a usage-only chunk whose choices + // array is empty: delta is undefined, so the chunk must be skipped + // without touching the optional-chained delta fields. + const generator = asyncStreamFrom([ + chunkOf({ content: "content" }), + { choices: [], usage: { prompt_tokens: 5, completion_tokens: 7 } }, + ]) + const stream = startStream(controller.signal, generator) + const first = await stream.next() + expect(first.value?.type).toBe("text") + + const chunks: unknown[] = [] + for await (const c of stream) chunks.push(c) + expect(chunks).toHaveLength(1) + expect((chunks[0] as { type: string }).type).toBe("usage") + }) + }) }) diff --git a/src/api/providers/__tests__/mistral.spec.ts b/src/api/providers/__tests__/mistral.spec.ts index aa96934b4a..bf97a397ba 100644 --- a/src/api/providers/__tests__/mistral.spec.ts +++ b/src/api/providers/__tests__/mistral.spec.ts @@ -729,6 +729,20 @@ describe("MistralHandler", () => { }) }) + it("should reject immediately with AbortError when the external signal is pre-aborted, before the request is built", async () => { + const controller = new AbortController() + controller.abort() + + const error = await handler + .completePrompt("Test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("This operation was aborted") + // The fast-fail must run before the client is touched. + expect(mockComplete).not.toHaveBeenCalled() + }) + it("should work without options (backward compatible)", async () => { mockComplete.mockResolvedValueOnce({ choices: [{ message: { content: "response" } }], @@ -795,14 +809,26 @@ describe("MistralHandler", () => { expect(mockComplete).toHaveBeenCalledWith(expect.objectContaining({ model: expect.any(String) }), undefined) }) - it("should surface a standard AbortError when the signal was aborted and the request fails", async () => { - mockComplete.mockRejectedValueOnce(new Error("API Error")) + it("should surface a standard AbortError when the signal is aborted while the request is in flight", async () => { const controller = new AbortController() + // The request stays pending until the external signal aborts it; the + // mock rejects with a plain (non-abort) error once the signal has + // aborted, so the catch block must classify it via signal.aborted. + mockComplete.mockImplementationOnce(() => { + return new Promise((_resolve, reject) => { + const fail = () => reject(new Error("API Error")) + if (controller.signal.aborted) { + fail() + return + } + controller.signal.addEventListener("abort", fail, { once: true }) + }) + }) + const promise = handler.completePrompt("Test prompt", { abortSignal: controller.signal }) + // Abort after entry — the pre-abort fast-fail must not apply, and the + // catch block handles the failure. controller.abort() - - const error = await handler - .completePrompt("Test prompt", { abortSignal: controller.signal }) - .catch((e: unknown) => e) + const error = await promise.catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) expect((error as Error).name).toBe("AbortError") expect((error as Error).message).toBe("The Mistral request was aborted") @@ -933,4 +959,134 @@ describe("MistralHandler", () => { expect(streamOptions?.fetchOptions?.signal?.aborted).toBe(false) }) }) + + describe("createMessage streaming loop abort defense", () => { + const systemPrompt = "You are a helpful assistant." + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: [{ type: "text", text: "Hello!" }], + }, + ] + + const thinkingChunk = { type: "thinking", thinking: [{ type: "text", text: "reasoning" }] } + const textChunk = { type: "text", text: "content" } + + function eventOf(delta: Record, usage?: Record) { + return { + data: { + choices: [ + { + delta, + index: 0, + }, + ], + ...(usage ? { usage } : {}), + }, + } + } + + function startStream(signal: AbortSignal, generator: AsyncGenerator) { + mockCreate.mockImplementationOnce(async () => generator) + return handler.createMessage(systemPrompt, messages, makeCreateMessageMetadata({ abortSignal: signal })) + } + + it("rejects with AbortError instead of yielding text content after a mid-event abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([eventOf({ content: [thinkingChunk, textChunk] })]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("reasoning") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + }) + + it("rejects with AbortError instead of yielding reasoning after a mid-event abort (text-first event)", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([eventOf({ content: [textChunk, thinkingChunk] })]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + }) + + it("rejects with AbortError instead of yielding a tool call after a mid-event abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([ + eventOf({ + content: [thinkingChunk], + toolCalls: [{ id: "1", function: { name: "f", arguments: "{}" } }], + }), + ]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("reasoning") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + }) + + it("rejects with AbortError instead of yielding usage after a mid-event abort", async () => { + const controller = new AbortController() + const stream = startStream( + controller.signal, + asyncStreamFrom([eventOf({ content: [thinkingChunk] }, { promptTokens: 1, completionTokens: 2 })]), + ) + const first = await stream.next() + expect(first.value?.type).toBe("reasoning") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + }) + + it("does not process or pull events beyond the aborted point and surfaces AbortError instead of completing the stream", async () => { + const controller = new AbortController() + let pulls = 0 + const generator = (async function* () { + pulls++ + yield eventOf({ content: [textChunk] }) + pulls++ + // String-content shape: its text yield is the first yield of the + // iteration and carries no pre-yield guard, so only the top-of-loop + // break prevents it from leaking. + yield eventOf({ content: "buffered" }) + pulls++ + yield eventOf({ content: [textChunk] }) + })() + const stream = startStream(controller.signal, generator) + const first = await stream.next() + expect(first.value?.type).toBe("text") + controller.abort() + + const error = await stream.next().catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The Mistral request was aborted") + // The for-await mechanism pulls the already-buffered event before the + // top-of-loop check runs; the break must ensure it is not processed and + // that no event beyond it is pulled. + expect(pulls).toBe(2) + }) + }) }) diff --git a/src/api/providers/gemini.ts b/src/api/providers/gemini.ts index 71209e1c52..1658a1a484 100644 --- a/src/api/providers/gemini.ts +++ b/src/api/providers/gemini.ts @@ -432,10 +432,15 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl let finishReason: string | undefined let toolCallCounter = 0 - let hasContent = false - let hasReasoning = false for await (const chunk of result) { + // Stop consuming buffered chunks once the request is aborted (the + // SDK iterator may keep delivering buffered content after abort). + // Stryker disable next-line OptionalChaining: requestSignal is always set by the addMergedSignal call above (a request-local controller signal exists even without an external signal), so the optional chain cannot observe a nullish value + if (requestSignal?.aborted) { + break + } + // Track the final structured response (per SDK pattern: candidate.finishReason) if (chunk.candidates && chunk.candidates[0]?.finishReason) { finalResponse = chunk as { responseId?: string } @@ -468,17 +473,19 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl if (part.thought) { // This is a thinking/reasoning part if (part.text) { - hasReasoning = true + // Re-check before emitting: the consumer may have aborted + // while processing a previously yielded part of this chunk. + throwIfAborted(requestSignal) yield { type: "reasoning", text: part.text } } } else if (part.functionCall) { - hasContent = true // Gemini sends complete function calls in a single chunk // Emit as partial chunks for consistent handling with NativeToolCallParser const callId = `${part.functionCall.name}-${toolCallCounter}` const args = JSON.stringify(part.functionCall.args) // Emit name first + throwIfAborted(requestSignal) yield { type: "tool_call_partial", index: toolCallCounter, @@ -488,6 +495,7 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } // Then emit arguments + throwIfAborted(requestSignal) yield { type: "tool_call_partial", index: toolCallCounter, @@ -500,7 +508,7 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } else { // This is regular content if (part.text) { - hasContent = true + throwIfAborted(requestSignal) yield { type: "text", text: part.text } } } @@ -509,8 +517,10 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } // Fallback to the original text property if no candidates structure + // (no pre-yield guard here: this is the only yield of a fallback + // chunk, and the top-of-loop check runs without a suspension point + // before it). else if (chunk.text) { - hasContent = true yield { type: "text", text: chunk.text } } @@ -519,6 +529,11 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } } + // An aborted request must surface as an AbortError, not as a normal + // stream completion: the top-of-loop break (or a swallowed mid-stream + // abort) ends the loop without throwing otherwise. + throwIfAborted(requestSignal) + if (finalResponse?.responseId) { // Capture responseId so Task.addToApiConversationHistory can store it // alongside the assistant message in api_history.json. @@ -660,6 +675,9 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl } async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { + // Fast-fail if the request was already aborted before building. + throwIfAborted(options?.abortSignal) + const { id: model, info } = this.getModel() try { diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index 2acad4d9c9..17ca733d55 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -291,21 +291,33 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa let lastUsage for await (const chunk of completion) { + // Stop consuming buffered chunks once the request is aborted (the + // OpenAI SDK iterator may keep delivering buffered content after + // abort, and swallows the mid-stream AbortError). + // Stryker disable next-line OptionalChaining: requestSignal is always set by the addMergedSignal call above (a request-local controller signal exists even without an external signal), so the optional chain cannot observe a nullish value + if (requestSignal?.aborted) { + break + } + const delta = chunk.choices[0]?.delta const usage = chunk.usage as LiteLLMUsage const reasoningText = extractReasoningFromDelta(delta) if (reasoningText) { + // No pre-yield guard: this is the first yield of the iteration — + // the top-of-loop check runs without a suspension point before it. yield { type: "reasoning", text: reasoningText } } if (delta?.content) { + throwIfAborted(requestSignal) yield { type: "text", text: delta.content } } // Handle tool calls in stream - emit partial chunks for NativeToolCallParser if (delta?.tool_calls) { for (const toolCall of delta.tool_calls) { + throwIfAborted(requestSignal) yield { type: "tool_call_partial", index: toolCall.index, @@ -321,6 +333,11 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa } } + // An aborted request must surface as an AbortError, not as a normal + // stream completion: the top-of-loop break (or a swallowed mid-stream + // abort) ends the loop without throwing otherwise. + throwIfAborted(requestSignal) + if (lastUsage) { // Extract cache-related information if available // LiteLLM may use different field names for cache tokens diff --git a/src/api/providers/mistral.ts b/src/api/providers/mistral.ts index 260df0326c..48e4150539 100644 --- a/src/api/providers/mistral.ts +++ b/src/api/providers/mistral.ts @@ -121,11 +121,20 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand }) for await (const event of response) { + // Stop consuming buffered events once the request is aborted (the + // SDK iterator may keep delivering buffered content after abort). + // Stryker disable next-line OptionalChaining: requestSignal is always set by the addMergedSignal call above (a request-local controller signal exists even without an external signal), so the optional chain cannot observe a nullish value + if (requestSignal?.aborted) { + break + } + const delta = event.data.choices[0]?.delta if (delta?.content) { if (typeof delta.content === "string") { - // Handle string content as text + // Handle string content as text (no pre-yield guard: this is + // the first yield of the iteration — the top-of-loop check + // runs without a suspension point before it). yield { type: "text", text: delta.content } } else if (Array.isArray(delta.content)) { // Handle array of content chunks @@ -136,11 +145,13 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand // ThinkChunk has a 'thinking' property that contains an array of text/reference chunks for (const thinkingPart of chunk.thinking) { if (thinkingPart.type === "text" && thinkingPart.text) { + throwIfAborted(requestSignal) yield { type: "reasoning", text: thinkingPart.text } } } } else if (chunk.type === "text" && chunk.text) { // Handle text content normally + throwIfAborted(requestSignal) yield { type: "text", text: chunk.text } } } @@ -153,6 +164,7 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand if (toolCalls) { for (let i = 0; i < toolCalls.length; i++) { const toolCall = toolCalls[i] + throwIfAborted(requestSignal) yield { type: "tool_call_partial", index: i, @@ -164,6 +176,7 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand } if (event.data.usage) { + throwIfAborted(requestSignal) yield { type: "usage", inputTokens: event.data.usage.promptTokens || 0, @@ -171,6 +184,11 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand } } } + + // An aborted request must surface as an AbortError, not as a normal + // stream completion: the top-of-loop break (or a swallowed mid-stream + // abort) ends the loop without throwing otherwise. + throwIfAborted(requestSignal) } catch (error) { // Aborted request: covers both the external signal and SDK-native // abort errors that may surface before the signal flag propagates. @@ -213,6 +231,9 @@ export class MistralHandler extends BaseProvider implements SingleCompletionHand return { id, info, maxTokens, temperature } } async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { + // Fast-fail if the request was already aborted before building. + throwIfAborted(options?.abortSignal) + const { id: model, temperature } = this.getModel() try { From 0ec250e3d7ea22697fce1faeaa7623cb630a2bb2 Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Fri, 11 Sep 2026 19:37:10 +0800 Subject: [PATCH 16/20] test(api): assert the mapped usage values in the LiteLLM usage-only case The usage-only empty-choices test checked only the chunk type. Assert the mapped values (inputTokens: 5, outputTokens: 7) so the prompt_tokens / completion_tokens mapping in this branch is pinned, not just its presence. Addresses the CodeRabbit finding on the structural kill-test suite. --- src/api/providers/__tests__/lite-llm.spec.ts | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index cc16ded6c0..8e2d11dffc 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1732,7 +1732,9 @@ describe("LiteLLMHandler", () => { const chunks: unknown[] = [] for await (const c of stream) chunks.push(c) expect(chunks).toHaveLength(1) - expect((chunks[0] as { type: string }).type).toBe("usage") + // Assert the mapped usage values, not just the type: this branch maps + // prompt_tokens/completion_tokens onto the usage chunk (L365-366). + expect(chunks[0]).toMatchObject({ type: "usage", inputTokens: 5, outputTokens: 7 }) }) }) }) From f0f1fa5b9d857755c6acfd4accd2f64c20d080dc Mon Sep 17 00:00:00 2001 From: Eason Liang Date: Wed, 23 Sep 2026 13:06:41 +0800 Subject: [PATCH 17/20] fix(api): settle LiteLLM requests with AbortError during model discovery Model discovery (RouterProvider.fetchModel) is a shared single-flight fetch, so a per-request signal must not be threaded into it: one caller's abort would reject the in-flight discovery of every other waiter. Instead, race each request's discovery against its own signal via rejectOnAbort (new utility in utils/abort-signal) and normalize aborted discovery failures with createAbortError("LiteLLM"): - createMessage: race discovery against metadata.abortSignal - completePrompt: race discovery against the merged signal (mergeAbortSignalAndTimeout: external abort + per-request timeout, timeoutMs <= 0 disabling the timeout), so a stalled cold-cache fetch cannot outlive the requested timeout - non-abort discovery failures propagate unchanged (exact identity) - the completePrompt "abort while in flight" test now lands the abort after discovery (readiness barrier) so it still exercises the SDK catch's signal.aborted classification; the bridging test now asserts the transient discovery-race listener is the only manual listener on the external signal and that it is detached at settle - kill tests: mid-discovery abort, preserved non-abort failure, concurrent failure+abort normalization, single-flight isolation (a concurrent non-aborted request sharing the fetch is unaffected), and the timeoutMs boundary; unit tests for rejectOnAbort itself Addresses the CodeRabbit out-of-diff finding on the current head ("propagate cancellation through cold-cache model discovery"). --- src/api/providers/__tests__/lite-llm.spec.ts | 308 +++++++++++++++++- src/api/providers/lite-llm.ts | 62 +++- .../utils/__tests__/abort-signal.spec.ts | 98 ++++++ src/api/providers/utils/abort-signal.ts | 32 ++ 4 files changed, 490 insertions(+), 10 deletions(-) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 8e2d11dffc..86ac34c768 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -2,8 +2,9 @@ import OpenAI from "openai" import { Anthropic } from "@anthropic-ai/sdk" import { LiteLLMHandler } from "../lite-llm" +import { getModels } from "../fetchers/modelCache" import { ApiHandlerOptions } from "../../../shared/api" -import { litellmDefaultModelId, litellmDefaultModelInfo } from "@roo-code/types" +import { type ModelRecord, litellmDefaultModelId, litellmDefaultModelInfo } from "@roo-code/types" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" import { clearAllMocks } from "../../../test-utils/reset" import { makeCreateMessageMetadata } from "../../../test-utils/api" @@ -1446,10 +1447,18 @@ describe("LiteLLMHandler", () => { it("should surface a standard AbortError when the signal is aborted while the request is in flight", async () => { const controller = new AbortController() + // Readiness barrier: resolves once the SDK call has started, so the + // abort lands while the request is in flight — not during model + // discovery, which now settles via the rejectOnAbort race first. + let requestStartedResolve!: () => void + const requestStarted = new Promise((resolve) => { + requestStartedResolve = resolve + }) // The request stays pending until the external signal aborts it; the // mock rejects with a plain (non-abort) error once the signal has // aborted, so the catch block must classify it via signal.aborted. mockCreate.mockImplementationOnce((_body: unknown, options?: { signal?: AbortSignal }) => { + requestStartedResolve() const signal = options?.signal return new Promise((_resolve, reject) => { const fail = () => reject(new Error("LiteLLM API error")) @@ -1461,8 +1470,10 @@ describe("LiteLLMHandler", () => { }) }) const promise = handler.completePrompt("test prompt", { abortSignal: controller.signal }) - // Abort while the (pending) request is in flight — after entry, so the - // pre-abort fast-fail does not apply and the catch block handles it. + // Abort while the (pending) request is in flight — after discovery and + // entry, so the pre-abort fast-fail and the discovery race do not + // settle it first and the catch block handles the SDK rejection. + await requestStarted controller.abort() const error = await promise.catch((e: unknown) => e) expect(error).toBeInstanceOf(Error) @@ -1594,10 +1605,16 @@ describe("LiteLLMHandler", () => { expect((error as Error).message).toBe("The LiteLLM request was aborted") expect(capturedSignal).toBeDefined() expect(capturedSignal?.aborted).toBe(true) - // The bridge (RequestConfigBuilder.addMergedSignal) uses AbortSignal.any, - // so the external signal must never be managed with manual listeners. - expect(addEventListenerSpy).not.toHaveBeenCalled() - expect(removeEventListenerSpy).not.toHaveBeenCalled() + // The streaming bridge (RequestConfigBuilder.addMergedSignal) uses + // AbortSignal.any. The only manual listener on the external signal is + // the transient discovery-race listener (rejectOnAbort); it must be + // detached by the time the request settles — no listener may outlive + // the request. Exactly one registration proves the streaming bridge + // itself adds no manual listeners of its own. + expect(addEventListenerSpy).toHaveBeenCalledTimes(1) + const registeredListener = addEventListenerSpy.mock.calls[0]?.[1] as EventListener | undefined + expect(typeof registeredListener).toBe("function") + expect(removeEventListenerSpy).toHaveBeenCalledWith("abort", registeredListener) }) it("should wrap a non-abort stream failure with the i18n-free provider message and no metadata", async () => { @@ -1624,6 +1641,283 @@ describe("LiteLLMHandler", () => { }) }) + describe("model discovery cancellation (rejectOnAbort)", () => { + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: "Hello", + }, + ] + + it("settles with AbortError when the external signal aborts during model discovery", async () => { + // The shared single-flight discovery fetch is in flight (a cold-cache + // model fetch that never settles): the request must not wait for it. + const discoveryStarted = new Promise((resolve) => { + vi.mocked(getModels).mockImplementationOnce( + () => + new Promise(() => { + resolve() + }), + ) + }) + const controller = new AbortController() + + const stream = handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collector = collectStream(stream).catch((e: unknown) => e) + await discoveryStarted + controller.abort() + + // Bound the wait so a broken discovery race fails this test fast + // (and fails the Stryker mutant) instead of hanging. + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("discovery abort deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("preserves the model-fetch error when discovery fails without an abort", async () => { + const boom = new Error( + "Failed to fetch LiteLLM models: No response from server. Check LiteLLM server status and base URL.", + ) + vi.mocked(getModels).mockRejectedValueOnce(boom) + + const error = await collectStream( + handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ), + ).catch((e: unknown) => e) + // Exact identity: a non-abort discovery failure must not be wrapped + // or misclassified as a cancellation. + expect(error).toBe(boom) + }) + + it("preserves the model-fetch error when discovery fails and no metadata is passed", async () => { + const boom = new Error( + "Failed to fetch LiteLLM models: No response from server. Check LiteLLM server status and base URL.", + ) + vi.mocked(getModels).mockRejectedValueOnce(boom) + + const error = await collectStream(handler.createMessage("system", messages)).catch((e: unknown) => e) + // The catch's re-classification must read the signal through the + // optional chain: a bare `metadata.abortSignal` would throw a + // TypeError here instead of surfacing the fetcher's error. + expect(error).toBe(boom) + }) + + it("normalizes a fetcher AbortError-shaped discovery failure to the standard AbortError", async () => { + // A low-level transport can reject with an AbortError-named error that + // is not the standard abort message (the SDK's DOMException shape). + // The catch must normalize it via the isRequestAborted error-name + // branch even though the signal never aborted. + vi.mocked(getModels).mockRejectedValueOnce( + Object.assign(new Error("The operation was aborted"), { name: "AbortError" }), + ) + + const error = await collectStream( + handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: new AbortController().signal }), + ), + ).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("normalizes to AbortError when the signal aborts while discovery is failing", async () => { + // Discovery rejects only once the signal aborts: whichever settles the + // race first (the abort listener or the forwarded failure), the + // surfaced error must be the standard AbortError. + const controller = new AbortController() + const discoveryStarted = new Promise((resolve) => { + vi.mocked(getModels).mockImplementationOnce( + () => + new Promise((_resolve, reject) => { + resolve() + controller.signal.addEventListener( + "abort", + () => reject(new Error("Failed to fetch LiteLLM models: 503 Service Unavailable.")), + { once: true }, + ) + }), + ) + }) + + const stream = handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: controller.signal }), + ) + const collector = collectStream(stream).catch((e: unknown) => e) + await discoveryStarted + controller.abort() + + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("discovery abort deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("does not let one request's abort settle a concurrent request sharing the discovery fetch", async () => { + // One handler: both requests join the same single-flight fetchModel. + let discoveryResolve!: (models: ModelRecord) => void + const discoveryStarted = new Promise((resolve) => { + vi.mocked(getModels).mockImplementationOnce( + () => + new Promise((resolveModels) => { + discoveryResolve = resolveModels + resolve() + }), + ) + }) + const abortedController = new AbortController() + const siblingController = new AbortController() + + const aborted = collectStream( + handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: abortedController.signal }), + ), + ).catch((e: unknown) => e) + // The sibling request must keep flowing normally once the shared + // discovery settles. + mockCreate.mockImplementationOnce(() => ({ + withResponse: vi.fn().mockResolvedValue({ + data: asyncStreamFrom([{ choices: [{ delta: { content: "sibling" } }], usage: undefined }]), + }), + })) + const sibling = collectStream( + handler.createMessage( + "system", + messages, + makeCreateMessageMetadata({ abortSignal: siblingController.signal }), + ), + ).catch((e: unknown) => e) + + await discoveryStarted + abortedController.abort() + + const abortedError = await aborted + expect((abortedError as Error).name).toBe("AbortError") + + // The shared fetch still runs (the discovery promise is still + // pending, which this test owns), so the sibling request cannot + // have settled yet. + let siblingSettled = false + void sibling.then(() => { + siblingSettled = true + }) + expect(siblingSettled).toBe(false) + + // ...and settles normally once the shared discovery resolves. + discoveryResolve({ [litellmDefaultModelId]: litellmDefaultModelInfo }) + const siblingResult = await sibling + expect(siblingResult).toBeInstanceOf(Array) + const siblingStream = siblingResult as Array<{ type: string; text?: string }> + expect(siblingStream.filter((chunk) => chunk.type === "text").map((chunk) => chunk.text)).toEqual([ + "sibling", + ]) + }) + + it("settles completePrompt with AbortError when the signal aborts during model discovery", async () => { + const discoveryStarted = new Promise((resolve) => { + vi.mocked(getModels).mockImplementationOnce( + () => + new Promise(() => { + resolve() + }), + ) + }) + const controller = new AbortController() + + const collector = handler + .completePrompt("test prompt", { abortSignal: controller.signal }) + .catch((e: unknown) => e) + await discoveryStarted + controller.abort() + + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("discovery abort deadline exceeded")), 3000) + collector.then((result) => { + clearTimeout(deadline) + resolve(result) + }) + }) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("settles completePrompt with AbortError when timeoutMs elapses during model discovery", async () => { + // A per-request timeout bounds discovery as well: a stalled cold-cache + // fetch must not outlive the requested timeout. + vi.mocked(getModels).mockImplementationOnce(() => new Promise(() => {})) + + const error = await new Promise((resolve) => { + const deadline = setTimeout(() => resolve(new Error("timeout deadline exceeded")), 3000) + handler.completePrompt("test prompt", { timeoutMs: 50 }).then( + () => { + clearTimeout(deadline) + resolve(new Error("completePrompt unexpectedly resolved")) + }, + (e) => { + clearTimeout(deadline) + resolve(e) + }, + ) + }) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + + it("preserves the model-fetch error when completePrompt's discovery fails without a signal", async () => { + // No abort signal and no timeout: a plain fetcher failure surfaces + // unwrapped (the SDK section's completion-error wrap does not cover + // discovery) and must never be reclassified as a cancellation. + const boom = new Error("Failed to fetch LiteLLM models: No response from server.") + vi.mocked(getModels).mockRejectedValueOnce(boom) + + const error = await handler.completePrompt("test prompt").catch((e: unknown) => e) + expect(error).toBe(boom) + }) + + it("normalizes a fetcher AbortError-shaped discovery failure in completePrompt", async () => { + // As in createMessage: the discovery catch is the only re-classifier + // on this path (the SDK section's catch does not cover discovery), so + // an AbortError-named fetcher failure must normalize to the standard + // AbortError even though the signal never aborted. + vi.mocked(getModels).mockRejectedValueOnce( + Object.assign(new Error("The operation was aborted"), { name: "AbortError" }), + ) + + const error = await handler.completePrompt("test prompt", { timeoutMs: 0 }).catch((e: unknown) => e) + expect(error).toBeInstanceOf(Error) + expect((error as Error).name).toBe("AbortError") + expect((error as Error).message).toBe("The LiteLLM request was aborted") + }) + }) + describe("createMessage streaming loop abort defense", () => { const messages: Anthropic.Messages.MessageParam[] = [ { diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index 17ca733d55..646aa8929b 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -22,7 +22,13 @@ import { sanitizeOpenAiCallId } from "../../utils/tool-id" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" import { RouterProvider } from "./router-provider" import { extractReasoningFromDelta } from "./utils/extract-reasoning" -import { createAbortError, isRequestAborted, throwIfAborted } from "./utils/abort-signal" +import { + createAbortError, + isRequestAborted, + mergeAbortSignalAndTimeout, + rejectOnAbort, + throwIfAborted, +} from "./utils/abort-signal" import { RequestConfigBuilder } from "./config-builder/request-config-builder" import { getRequestTimeoutMs } from "./utils/request-timeout" @@ -141,7 +147,29 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa // discovery (getModels/refreshModels) begins. throwIfAborted(metadata?.abortSignal) - const { id: modelId, info } = await this.fetchModel() + // Model discovery is shared single-flight (RouterProvider.fetchModel): + // concurrent callers join one in-flight fetch, so the fetch itself + // carries no request signal — one caller's abort would reject other + // requests' shared discovery. Cancellation during discovery is a + // rejectOnAbort race: this request settles promptly with AbortError + // while the shared fetch keeps running. + let model: Awaited> + try { + if (metadata?.abortSignal) { + // Stryker disable next-line StringLiteral: the race's providerName is unobservable — every error escaping the race passes through the catch below, which re-normalizes abort errors via createAbortError("LiteLLM"), so no observable outcome can depend on the parameter value. + model = await rejectOnAbort(this.fetchModel(), metadata.abortSignal, "LiteLLM") + } else { + model = await this.fetchModel() + } + } catch (error) { + // An abort landing while discovery fails must still surface as the + // standard AbortError, not the fetcher's generic model-fetch error. + if (isRequestAborted(error, metadata?.abortSignal)) { + throw createAbortError("LiteLLM") + } + throw error + } + const { id: modelId, info } = model // Models that require reasoning_content to be echoed back during tool-call // continuations (see LITELLM_PRESERVE_REASONING_MODEL_IDS) need convertToR1Format: @@ -390,7 +418,35 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa // discovery (getModels/refreshModels) begins. throwIfAborted(options?.abortSignal) - const { id: modelId, info } = await this.fetchModel() + // Per the abort-signal series contract, the merged signal (external + // abort + per-request timeout, timeoutMs <= 0 disabling the timeout) + // also bounds model discovery: a stalled cold-cache fetch must settle + // with AbortError when the caller cancels or the timeout elapses. + const requestAbortSignal = mergeAbortSignalAndTimeout(options?.abortSignal, options?.timeoutMs) + + // Model discovery is shared single-flight (RouterProvider.fetchModel): + // concurrent callers join one in-flight fetch, so the fetch itself + // carries no request signal — one caller's abort would reject other + // requests' shared discovery. Cancellation during discovery is a + // rejectOnAbort race: this request settles promptly with AbortError + // while the shared fetch keeps running. + let model: Awaited> + try { + if (requestAbortSignal) { + // Stryker disable next-line StringLiteral: the race's providerName is unobservable — every error escaping the race passes through the catch below, which re-normalizes abort errors via createAbortError("LiteLLM"), so no observable outcome can depend on the parameter value. + model = await rejectOnAbort(this.fetchModel(), requestAbortSignal, "LiteLLM") + } else { + model = await this.fetchModel() + } + } catch (error) { + // An abort landing while discovery fails must still surface as the + // standard AbortError, not the fetcher's generic model-fetch error. + if (isRequestAborted(error, requestAbortSignal)) { + throw createAbortError("LiteLLM") + } + throw error + } + const { id: modelId, info } = model // Check if this is a GPT-5 model that requires max_completion_tokens instead of max_tokens const usesMaxCompletionTokens = this.isGpt5(modelId) || info.requiresResponsesApi 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) + }, + ) + }) +} From 0388aeccddc4e5dd4061b7d5028ab6839dad45ba Mon Sep 17 00:00:00 2001 From: eason liang <136036952+easonLiangWorldedtech@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:38:01 +0800 Subject: [PATCH 18/20] ci: retrigger e2e-mock after flaky markdown-lists timeout The e2e-mock run at 4f746b855 failed only on the 'Markdown List Rendering' suite (3 tests, 30s waitUntilCompleted timeout); every other required check was green. The identical failure signature (same suite, same 30s timeout) occurred in the same CI window on two unrelated external forks, and the previous head of this branch passed e2e-mock with the identical test file present. This PR's diff touches no e2e or mock content. Empty commit to retrigger the workflow. From c299aa213e823fa38ee3a5c6e07e3437fdcf511d Mon Sep 17 00:00:00 2001 From: eason liang <136036952+easonLiangWorldedtech@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:55:30 +0800 Subject: [PATCH 19/20] ci: retrigger e2e-mock (second flake: restart-persistence 30s timeout) Run 36404841531 failed on the 'Run mocked restart-persistence E2E test' step only (the main mocked suite, incl. Markdown List Rendering, passed). The verify phase rehydrated the persisted task, hit TaskMessagesReadError (readFileWithMissingRetry's single 1-10ms ENOENT retry lost the race under CI load), and the 30s history-sequence poll timed out. The identical tree (4f746b855) passed that step in run 36396120353, and the e2e-mock workflow shows chronic cross-fork flakiness over recent weeks (repeated re-rolls on other branches). No content change in this PR touches e2e or task persistence. From 5d1dff4be6c66ad036be16cc33e192c8438ff76d Mon Sep 17 00:00:00 2001 From: easonLiangWorldedtech Date: Tue, 29 Sep 2026 00:41:07 +0800 Subject: [PATCH 20/20] fix(api): bound LiteLLM completePrompt by a single deadline and classify discovery timeouts distinctly Address the CodeRabbit findings on completePrompt: - Compute one deadline before model discovery and pass only the remaining budget to the SDK via createOptions.timeout, so a request's total lifetime is bounded by the configured timeoutMs instead of discovery plus a fresh full timeout. - A per-request timeout elapsing during discovery now surfaces as a timeout-specific error (name TimeoutError) rather than the abort contract error: a timeout is not a user cancellation. Caller aborts still settle with the standard AbortError, and fetcher-originated AbortError-shaped failures keep their existing normalization. - The shared-discovery sibling test now drains the task queue before asserting the sibling has not settled, making the assertion meaningful instead of timing-trivially-true. - A new spec asserts the SDK receives the remaining (not full) timeout budget after discovery. --- src/api/providers/__tests__/lite-llm.spec.ts | 34 ++++++++++-- src/api/providers/lite-llm.ts | 55 +++++++++++++++----- 2 files changed, 73 insertions(+), 16 deletions(-) diff --git a/src/api/providers/__tests__/lite-llm.spec.ts b/src/api/providers/__tests__/lite-llm.spec.ts index 86ac34c768..878e56739e 100644 --- a/src/api/providers/__tests__/lite-llm.spec.ts +++ b/src/api/providers/__tests__/lite-llm.spec.ts @@ -1827,6 +1827,11 @@ describe("LiteLLMHandler", () => { void sibling.then(() => { siblingSettled = true }) + // Drain the microtask queue (and one macrotask) before asserting: + // an erroneous early settlement (e.g. the abort leaking into the + // sibling) settles `sibling` in a microtask, so an assertion + // right after attaching the callback could never observe it. + await new Promise((resolve) => setTimeout(resolve, 0)) expect(siblingSettled).toBe(false) // ...and settles normally once the shared discovery resolves. @@ -1868,9 +1873,11 @@ describe("LiteLLMHandler", () => { expect((error as Error).message).toBe("The LiteLLM request was aborted") }) - it("settles completePrompt with AbortError when timeoutMs elapses during model discovery", async () => { + it("settles completePrompt with a timeout error when timeoutMs elapses during model discovery", async () => { // A per-request timeout bounds discovery as well: a stalled cold-cache - // fetch must not outlive the requested timeout. + // fetch must not outlive the requested timeout. A timeout is not a + // user cancellation, so it surfaces as a timeout-specific error + // instead of the abort contract error. vi.mocked(getModels).mockImplementationOnce(() => new Promise(() => {})) const error = await new Promise((resolve) => { @@ -1887,8 +1894,27 @@ describe("LiteLLMHandler", () => { ) }) expect(error).toBeInstanceOf(Error) - expect((error as Error).name).toBe("AbortError") - expect((error as Error).message).toBe("The LiteLLM request was aborted") + expect((error as Error).name).toBe("TimeoutError") + expect((error as Error).message).toBe("The LiteLLM model discovery timed out") + }) + + it("passes only the remaining timeout budget to the SDK after discovery", async () => { + // Discovery consumes part of the configured timeout: the SDK call + // must receive the remaining budget, not the full timeoutMs. + vi.mocked(getModels).mockImplementationOnce( + () => + new Promise((resolve) => { + setTimeout(() => resolve({ [litellmDefaultModelId]: litellmDefaultModelInfo }), 50) + }), + ) + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "ok" } }] }) + + const result = await handler.completePrompt("test prompt", { timeoutMs: 500 }) + + expect(result).toBe("ok") + const createOptions = mockCreate.mock.calls[0]?.[1] as { timeout?: number } | undefined + expect(createOptions?.timeout).toBeGreaterThan(400) + expect(createOptions?.timeout).toBeLessThan(500) }) it("preserves the model-fetch error when completePrompt's discovery fails without a signal", async () => { diff --git a/src/api/providers/lite-llm.ts b/src/api/providers/lite-llm.ts index 646aa8929b..22edd44f97 100644 --- a/src/api/providers/lite-llm.ts +++ b/src/api/providers/lite-llm.ts @@ -412,36 +412,66 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa } } + /** + * Fresh error for the per-request timeout elapsing during model discovery. + * Distinct from the abort contract error on purpose: a timeout is not a + * user cancellation, and callers treat the two differently. + */ + private createDiscoveryTimeoutError(providerName: string): Error { + const timeoutError = new Error(`The ${providerName} model discovery timed out`) + timeoutError.name = "TimeoutError" + return timeoutError + } + async completePrompt(prompt: string, options?: CompletePromptOptions): Promise { // Fast-fail if the request was already aborted before building, so an // already-aborted request fails with AbortError before provider model // discovery (getModels/refreshModels) begins. - throwIfAborted(options?.abortSignal) + const callerAbortSignal = options?.abortSignal + throwIfAborted(callerAbortSignal) + + // A single deadline spans model discovery and the completion request: + // discovery consumes part of the configured timeout, so the SDK call + // must receive only the remaining budget, not the full timeoutMs. + // timeoutMs <= 0 means no per-request timeout and no deadline. + const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) + const deadline = timeoutMs !== undefined ? Date.now() + timeoutMs : undefined // Per the abort-signal series contract, the merged signal (external // abort + per-request timeout, timeoutMs <= 0 disabling the timeout) // also bounds model discovery: a stalled cold-cache fetch must settle - // with AbortError when the caller cancels or the timeout elapses. - const requestAbortSignal = mergeAbortSignalAndTimeout(options?.abortSignal, options?.timeoutMs) + // promptly when the caller cancels or the timeout elapses. + const discoverySignal = mergeAbortSignalAndTimeout(options?.abortSignal, timeoutMs) // Model discovery is shared single-flight (RouterProvider.fetchModel): // concurrent callers join one in-flight fetch, so the fetch itself // carries no request signal — one caller's abort would reject other // requests' shared discovery. Cancellation during discovery is a - // rejectOnAbort race: this request settles promptly with AbortError - // while the shared fetch keeps running. + // rejectOnAbort race: this request settles promptly while the shared + // fetch keeps running. let model: Awaited> try { - if (requestAbortSignal) { - // Stryker disable next-line StringLiteral: the race's providerName is unobservable — every error escaping the race passes through the catch below, which re-normalizes abort errors via createAbortError("LiteLLM"), so no observable outcome can depend on the parameter value. - model = await rejectOnAbort(this.fetchModel(), requestAbortSignal, "LiteLLM") + if (discoverySignal) { + // Stryker disable next-line StringLiteral: the race's providerName is unobservable — every error escaping the race passes through the catch below, which re-normalizes via createAbortError("LiteLLM") or createDiscoveryTimeoutError("LiteLLM"), so no observable outcome can depend on the parameter value. + model = await rejectOnAbort(this.fetchModel(), discoverySignal, "LiteLLM") } else { model = await this.fetchModel() } } catch (error) { // An abort landing while discovery fails must still surface as the // standard AbortError, not the fetcher's generic model-fetch error. - if (isRequestAborted(error, requestAbortSignal)) { + if (isRequestAborted(error, discoverySignal)) { + if (discoverySignal?.aborted) { + // Our merged signal fired: a caller stop surfaces as the + // standard AbortError, while the per-request timeout + // elapsing during discovery is a timeout — it must not be + // misreported as a user cancellation. + throw callerAbortSignal?.aborted + ? createAbortError("LiteLLM") + : this.createDiscoveryTimeoutError("LiteLLM") + } + // The signal never aborted: the fetcher itself failed with an + // AbortError-shaped error; normalize to the standard AbortError. throw createAbortError("LiteLLM") } throw error @@ -480,9 +510,10 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa if (options?.abortSignal) { createOptions.signal = options.abortSignal } - const timeoutMs = getRequestTimeoutMs(options?.timeoutMs) - if (timeoutMs !== undefined) { - createOptions.timeout = timeoutMs + if (deadline !== undefined) { + // Only the remaining budget reaches the SDK: model discovery has + // already consumed part of the configured timeout. + createOptions.timeout = Math.max(1, deadline - Date.now()) } const response = await this.client.chat.completions.create(