From eb1a46baf5823db251703a128ba4a1864fbe92ac Mon Sep 17 00:00:00 2001 From: can1357 Date: Tue, 9 Jun 2026 04:09:49 +0200 Subject: [PATCH] feat: added injectable fetch transport across AI and coding network flows - Added optional FetchImpl fields to compaction, proxy, AI, coding-agent, and mnemopi options. - Threaded injected fetch implementations through OAuth, discovery, and search/LLM request flows. - Removed exported hookFetch utility and its package entrypoint from utils. - Replaced global-fetch test monkeypatching with per-test FetchImpl mocks across test suites. --- .omp/rules/ts-hook-fetch.md | 51 ---- packages/agent/CHANGELOG.md | 4 +- packages/agent/src/compaction/compaction.ts | 5 + packages/agent/src/compaction/openai.ts | 8 +- packages/agent/src/proxy.ts | 5 +- .../test/proxy-stream-disconnect.test.ts | 27 +- packages/agent/test/remote-compaction.test.ts | 14 +- packages/ai/CHANGELOG.md | 6 +- packages/ai/src/auth-storage.ts | 1 + packages/ai/src/provider-models/google.ts | 19 +- packages/ai/src/provider-models/ollama.ts | 10 +- .../ai/src/provider-models/openai-compat.ts | 97 ++++++-- packages/ai/src/registry/api-key-login.ts | 3 + .../ai/src/registry/api-key-validation.ts | 14 +- packages/ai/src/registry/kilo.ts | 5 +- .../oauth/__tests__/xai-oauth.test.ts | 11 +- packages/ai/src/registry/oauth/anthropic.ts | 52 ++-- .../ai/src/registry/oauth/github-copilot.ts | 103 ++++---- .../ai/src/registry/oauth/minimax-code.ts | 2 + packages/ai/src/registry/oauth/types.ts | 2 + packages/ai/src/registry/oauth/xai-oauth.ts | 24 +- packages/ai/src/registry/oauth/xiaomi.ts | 13 +- packages/ai/src/registry/types.ts | 4 +- packages/ai/src/usage.ts | 4 +- packages/ai/src/utils/discovery/gemini.ts | 4 +- .../src/utils/discovery/openai-compatible.ts | 4 +- packages/ai/test/anthropic-client.test.ts | 3 +- packages/ai/test/anthropic-oauth.test.ts | 27 +- .../ai/test/anthropic-stream-timeout.test.ts | 3 - .../azure-openai-responses-stream.test.ts | 23 +- packages/ai/test/claude-usage-retry.test.ts | 19 +- packages/ai/test/firepass.live.ts | 13 +- packages/ai/test/firepass.test.ts | 30 ++- .../github-copilot-anthropic-auth.test.ts | 13 +- packages/ai/test/github-copilot-login.test.ts | 16 +- .../test/github-copilot-model-limits.test.ts | 21 +- .../github-copilot-openai-base-url.test.ts | 49 ++-- .../ai/test/google-antigravity-usage.test.ts | 7 +- .../google-gemini-cli-3x-thinking.test.ts | 38 +-- .../test/google-gemini-cli-alignment.test.ts | 19 +- packages/ai/test/google-system-prompt.test.ts | 17 +- packages/ai/test/helpers/fetch-mock.ts | 8 + packages/ai/test/issue-1203-repro.test.ts | 21 +- packages/ai/test/issue-1399-repro.test.ts | 21 +- packages/ai/test/issue-1617-repro.test.ts | 81 +++--- packages/ai/test/issue-1776-repro.test.ts | 28 +-- packages/ai/test/issue-1838-repro.test.ts | 16 +- packages/ai/test/issue-1846-repro.test.ts | 32 +-- packages/ai/test/issue-2080-repro.test.ts | 43 ++-- packages/ai/test/issue-2105-repro.test.ts | 80 +++--- packages/ai/test/issue-2113-repro.test.ts | 29 +-- packages/ai/test/issue-772-repro.test.ts | 19 +- packages/ai/test/issue-827-repro.test.ts | 8 +- packages/ai/test/issue-847-repro.test.ts | 16 +- packages/ai/test/issue-887-repro.test.ts | 15 +- packages/ai/test/issue-911-repro.test.ts | 31 ++- packages/ai/test/issue-912-repro.test.ts | 11 +- packages/ai/test/issue-945-repro.test.ts | 8 +- packages/ai/test/issue-955-repro.test.ts | 8 +- packages/ai/test/issue-957-repro.test.ts | 68 ++--- packages/ai/test/issue-959-repro.test.ts | 28 +-- packages/ai/test/issue-969-repro.test.ts | 15 +- packages/ai/test/kilo-login.test.ts | 26 +- packages/ai/test/minimax-code-login.test.ts | 12 +- packages/ai/test/nanogpt-login.test.ts | 18 +- packages/ai/test/nanogpt-model-limits.test.ts | 15 +- packages/ai/test/oauth-deepseek.test.ts | 34 +-- .../ai/test/ollama-cloud-provider.test.ts | 87 ++++--- packages/ai/test/ollama-provider.test.ts | 28 +-- packages/ai/test/openai-codex-stream.test.ts | 154 +++++++----- packages/ai/test/openai-codex-usage.test.ts | 3 +- .../ai/test/openai-completions-compat.test.ts | 93 ++++--- ...enai-completions-disable-reasoning.test.ts | 15 +- .../openai-completions-progress-chunk.test.ts | 14 +- ...enai-completions-upstream-provider.test.ts | 28 +-- .../test/openai-first-event-timeout.test.ts | 110 +++++---- .../openai-responses-cache-affinity.test.ts | 9 +- ...i-responses-omit-max-output-tokens.test.ts | 16 +- .../openai-responses-system-prompt.test.ts | 10 +- .../ai/test/openai-tool-strict-mode.test.ts | 40 +-- packages/ai/test/openrouter-login.test.ts | 13 +- .../ai/test/provider-fetch-override.test.ts | 131 +++++----- packages/ai/test/provider-response.test.ts | 15 +- packages/ai/test/raw-sse-sdk-capture.test.ts | 30 ++- packages/ai/test/request-debug.test.ts | 7 +- .../ai/test/stream-markup-healing.test.ts | 88 +++---- packages/ai/test/synthetic-login.test.ts | 14 +- packages/ai/test/wafer.live.ts | 14 +- packages/ai/test/wafer.test.ts | 37 ++- packages/ai/test/xiaomi-oauth.test.ts | 22 +- .../test/xiaomi-tp-login-integration.test.ts | 64 ++--- packages/ai/test/zenmux-login.test.ts | 19 +- packages/ai/test/zenmux-provider.test.ts | 9 +- packages/ai/test/zhipu-compat.test.ts | 27 +- packages/coding-agent/CHANGELOG.md | 4 +- .../coding-agent/src/config/model-registry.ts | 20 +- .../src/extensibility/custom-tools/types.ts | 4 +- .../coding-agent/src/mcp/oauth-discovery.ts | 11 +- packages/coding-agent/src/mcp/oauth-flow.ts | 17 +- packages/coding-agent/src/tools/fetch.ts | 12 +- packages/coding-agent/src/tools/image-gen.ts | 44 +++- .../src/tools/report-tool-issue.ts | 8 +- packages/coding-agent/src/web/kagi.ts | 7 +- packages/coding-agent/src/web/parallel.ts | 6 +- .../src/web/search/providers/anthropic.ts | 6 +- .../src/web/search/providers/base.ts | 3 +- .../src/web/search/providers/brave.ts | 7 +- .../src/web/search/providers/codex.ts | 8 +- .../src/web/search/providers/exa.ts | 99 +++++++- .../src/web/search/providers/gemini.ts | 6 + .../src/web/search/providers/jina.ts | 20 +- .../src/web/search/providers/kagi.ts | 11 +- .../src/web/search/providers/kimi.ts | 22 +- .../src/web/search/providers/parallel.ts | 8 +- .../src/web/search/providers/perplexity.ts | 11 +- .../src/web/search/providers/searxng.ts | 8 +- .../src/web/search/providers/synthetic.ts | 14 +- .../src/web/search/providers/tavily.ts | 6 +- .../src/web/search/providers/zai.ts | 19 +- .../test/auth-storage-minimax-login.test.ts | 9 +- packages/coding-agent/test/compaction.test.ts | 69 +++--- .../coding-agent/test/helpers/fetch-mock.ts | 18 ++ ...-1528-discovery-default-max-tokens.test.ts | 27 +- ...ssue-970-custom-provider-discovery.test.ts | 15 +- .../coding-agent/test/lm-studio-fix.test.ts | 31 ++- .../model-registry-runtime-provider.test.ts | 15 +- .../coding-agent/test/model-registry.test.ts | 156 ++++++------ .../coding-agent/test/model-resolver.test.ts | 3 +- .../coding-agent/test/oauth-discovery.test.ts | 35 ++- packages/coding-agent/test/oauth-flow.test.ts | 129 ++++------ .../test/tools/fetch-jina-stall.test.ts | 43 ++-- .../test/tools/fetch-kagi-toggle.test.ts | 35 +-- .../coding-agent/test/tools/image-gen.test.ts | 8 +- .../test/tools/report-tool-issue.test.ts | 61 +++-- .../web-scrapers/youtube-parallel.test.ts | 45 ++-- .../test/tools/web-search-codex.test.ts | 139 ++++++----- .../test/tools/web-search-exa.test.ts | 232 ++++++++++-------- .../test/tools/web-search-gemini.test.ts | 31 +-- .../test/tools/web-search-kagi.test.ts | 185 +++++++------- .../test/tools/web-search-parallel.test.ts | 32 +-- .../test/tools/web-search-searxng.test.ts | 110 +++++---- .../test/tools/web-search-tavily.test.ts | 20 +- .../test/web/search/abort-and-timeout.test.ts | 26 +- .../test/web/search/codex-broker.test.ts | 8 +- .../test/web/search/perplexity.test.ts | 62 +++-- .../test/web/search/tavily.test.ts | 23 +- packages/mnemopi/CHANGELOG.md | 7 +- packages/mnemopi/src/core/extraction.ts | 5 +- .../mnemopi/src/core/extraction/client.ts | 15 +- packages/mnemopi/src/core/llm-backends.ts | 3 + packages/mnemopi/src/core/local-llm.ts | 43 +++- .../test/extraction-integration.test.ts | 19 +- packages/mnemopi/test/local-llm.test.ts | 33 +-- packages/utils/CHANGELOG.md | 6 +- packages/utils/src/hook-fetch.ts | 30 --- packages/utils/src/index.ts | 1 - 156 files changed, 2491 insertions(+), 2197 deletions(-) delete mode 100644 .omp/rules/ts-hook-fetch.md create mode 100644 packages/ai/test/helpers/fetch-mock.ts create mode 100644 packages/coding-agent/test/helpers/fetch-mock.ts delete mode 100644 packages/utils/src/hook-fetch.ts diff --git a/.omp/rules/ts-hook-fetch.md b/.omp/rules/ts-hook-fetch.md deleted file mode 100644 index 09a91a39d..000000000 --- a/.omp/rules/ts-hook-fetch.md +++ /dev/null @@ -1,51 +0,0 @@ ---- -description: Use hookFetch instead of assigning or spying on globalThis.fetch in tests -condition: "globalThis\\.fetch\\s*=|spyOn\\(globalThis.*fetch" -scope: "tool:edit(**/*.test.{ts,tsx,js,jsx}), tool:write(**/*.test.{ts,tsx,js,jsx})" ---- - -**Do not assign `globalThis.fetch` or use `vi.spyOn(globalThis, "fetch")` in tests.** - -## Why it's wrong - -- Forgetting restoration leaks state across tests -- `vi.spyOn` ties fetch mocking to vitest lifecycle instead of explicit scoping -- Makes test mocking inconsistent across the codebase - -## What to use instead - -Use `hookFetch` from `@oh-my-pi/pi-utils`. It returns a `Disposable` — use `using` for automatic cleanup: - -```ts -import { hookFetch } from "@oh-my-pi/pi-utils"; - -using _hook = hookFetch((input, init, next) => { - // return a mocked Response, or delegate with next(input, init) -}); -``` - -## Examples - -```ts -// WRONG -globalThis.fetch = async () => new Response("ok"); - -// WRONG -vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("ok")); - -// RIGHT — fixed response -using _hook = hookFetch(() => new Response("ok")); - -// RIGHT — conditional mock with passthrough -using _hook = hookFetch((input, init, next) => { - if (String(input).includes("127.0.0.1")) { - return new Response(JSON.stringify({ data: [] })); - } - return next(input, init); -}); - -// RIGHT — when you need vi.fn() for mock assertions -const fetchSpy = vi.fn(() => new Response("ok")); -using _hook = hookFetch(fetchSpy); -// later: expect(fetchSpy.mock.calls[0]).toEqual(...) -``` \ No newline at end of file diff --git a/packages/agent/CHANGELOG.md b/packages/agent/CHANGELOG.md index 7793f55c8..dc50913f6 100644 --- a/packages/agent/CHANGELOG.md +++ b/packages/agent/CHANGELOG.md @@ -1,9 +1,11 @@ # Changelog ## [Unreleased] - ### Added +- Added optional `fetch` overrides to `SummaryOptions` and `compact`/`generateSummary` so remote compaction can use custom HTTP clients +- Added optional `fetch` option to `ProxyStreamOptions` to control the HTTP request used by `streamProxy` +- Added optional `fetch` overrides to `requestOpenAiRemoteCompaction` and `requestRemoteCompaction` for injectable HTTP transport - Added the upstream provider that served a request (`AssistantMessage.upstreamProvider`, e.g. OpenRouter's routed provider) as a `pi.gen_ai.response.upstream_provider` chat-span telemetry attribute, alongside the existing response id and time-to-first-chunk. ## [15.10.5] - 2026-06-08 diff --git a/packages/agent/src/compaction/compaction.ts b/packages/agent/src/compaction/compaction.ts index 47c9a58ac..9bb2f0106 100644 --- a/packages/agent/src/compaction/compaction.ts +++ b/packages/agent/src/compaction/compaction.ts @@ -9,6 +9,7 @@ import { type AssistantMessage, clampThinkingLevelForModel, Effort, + type FetchImpl, type Message, type MessageAttribution, type Model, @@ -595,6 +596,8 @@ export interface SummaryOptions { * `resolveCompactionEffort` for the conversion contract. */ thinkingLevel?: ThinkingLevel; + /** Optional fetch implementation threaded into remote compaction calls. */ + fetch?: FetchImpl; } export async function generateSummary( @@ -647,6 +650,7 @@ export async function generateSummary( prompt: promptText, }, signal, + { fetch: options.fetch }, ); return remote.summary; } @@ -1015,6 +1019,7 @@ export async function compact( remoteHistory, summaryOptions.remoteInstructions ?? SUMMARIZATION_SYSTEM_PROMPT, signal, + { fetch: summaryOptions.fetch }, ); preserveData = withOpenAiRemoteCompactionPreserveData(previousPreserveData, remote); } catch (err) { diff --git a/packages/agent/src/compaction/openai.ts b/packages/agent/src/compaction/openai.ts index 1e0dea3af..0b6ad71e8 100644 --- a/packages/agent/src/compaction/openai.ts +++ b/packages/agent/src/compaction/openai.ts @@ -20,7 +20,7 @@ import { } from "@oh-my-pi/pi-ai/providers/openai-codex/constants"; import { parseTextSignature } from "@oh-my-pi/pi-ai/providers/openai-responses-shared"; import { transformMessages } from "@oh-my-pi/pi-ai/providers/transform-messages"; -import type { AssistantMessage, Message, Model } from "@oh-my-pi/pi-ai/types"; +import type { AssistantMessage, FetchImpl, Message, Model } from "@oh-my-pi/pi-ai/types"; import { getOpenAIResponsesHistoryItems, getOpenAIResponsesHistoryPayload, @@ -428,6 +428,7 @@ export async function requestOpenAiRemoteCompaction( compactInput: Array>, instructions: string, signal?: AbortSignal, + opts?: { fetch?: FetchImpl }, ): Promise { const endpoint = resolveOpenAiCompactEndpoint(model); const request: OpenAiRemoteCompactionRequest = { @@ -451,7 +452,7 @@ export async function requestOpenAiRemoteCompaction( headers[OPENAI_HEADERS.ORIGINATOR] = OPENAI_HEADER_VALUES.ORIGINATOR_CODEX; } - const response = await fetch(endpoint, { + const response = await (opts?.fetch ?? fetch)(endpoint, { method: "POST", headers, body: JSON.stringify(request), @@ -501,8 +502,9 @@ export async function requestRemoteCompaction( endpoint: string, request: RemoteCompactionRequest, signal?: AbortSignal, + opts?: { fetch?: FetchImpl }, ): Promise { - const response = await fetch(endpoint, { + const response = await (opts?.fetch ?? fetch)(endpoint, { method: "POST", headers: { "content-type": "application/json" }, body: JSON.stringify(request), diff --git a/packages/agent/src/proxy.ts b/packages/agent/src/proxy.ts index 5a4c424a9..5bb82ef81 100644 --- a/packages/agent/src/proxy.ts +++ b/packages/agent/src/proxy.ts @@ -7,6 +7,7 @@ import { type AssistantMessageEvent, type Context, EventStream, + type FetchImpl, type Model, type SimpleStreamOptions, type StopReason, @@ -61,6 +62,8 @@ export interface ProxyStreamOptions extends SimpleStreamOptions { authToken: string; /** Proxy server URL (e.g., "https://genai.example.com") */ proxyUrl: string; + /** Optional fetch implementation; defaults to global fetch. */ + fetch?: FetchImpl; } /** @@ -117,7 +120,7 @@ export function streamProxy(model: Model, context: Context, options: ProxyStream } try { - response = await fetch(`${options.proxyUrl}/api/stream`, { + response = await (options.fetch ?? fetch)(`${options.proxyUrl}/api/stream`, { method: "POST", headers: { Authorization: `Bearer ${options.authToken}`, diff --git a/packages/agent/test/proxy-stream-disconnect.test.ts b/packages/agent/test/proxy-stream-disconnect.test.ts index 81cf1264b..325fe5ca2 100644 --- a/packages/agent/test/proxy-stream-disconnect.test.ts +++ b/packages/agent/test/proxy-stream-disconnect.test.ts @@ -9,8 +9,7 @@ import { describe, expect, it } from "bun:test"; import type { ProxyAssistantMessageEvent } from "@oh-my-pi/pi-agent-core/proxy"; import { type ProxyMessageEventStream, streamProxy } from "@oh-my-pi/pi-agent-core/proxy"; -import type { AssistantMessageEvent, Context, Model } from "@oh-my-pi/pi-ai"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { AssistantMessageEvent, Context, FetchImpl, Model } from "@oh-my-pi/pi-ai"; const mockModel: Model = { id: "test-model", @@ -76,12 +75,12 @@ describe("streamProxy — server disconnect without terminal event", () => { it("emits an error event when server disconnects after start with no terminal event", async () => { const events: ProxyAssistantMessageEvent[] = [{ type: "start" }]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", + fetch: fetchMock, }); const collected = await collectEvents(stream); const errorEvent = collected.find(e => e.type === "error"); @@ -98,12 +97,12 @@ describe("streamProxy — server disconnect without terminal event", () => { { type: "text_delta", contentIndex: 0, delta: "Hel" }, ]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", + fetch: fetchMock, }); // Consume iterator so the internal async function runs @@ -123,13 +122,13 @@ describe("streamProxy — server disconnect without terminal event", () => { const events: ProxyAssistantMessageEvent[] = [{ type: "start" }]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", signal: abortController.signal, + fetch: fetchMock, }); const collected = await collectEvents(stream); @@ -150,13 +149,13 @@ describe("streamProxy — server disconnect without terminal event", () => { const events: ProxyAssistantMessageEvent[] = [{ type: "start" }]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", signal: abortController.signal, + fetch: fetchMock, }); await collectEvents(stream); @@ -180,12 +179,12 @@ describe("streamProxy — server disconnect without terminal event", () => { }, ]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", + fetch: fetchMock, }); const collected = await collectEvents(stream); @@ -209,12 +208,12 @@ describe("streamProxy — server disconnect without terminal event", () => { }, ]; const body = buildSseBody(events); - - using _hook = hookFetch(() => new Response(body, { status: 200 })); + const fetchMock: FetchImpl = () => Promise.resolve(new Response(body, { status: 200 })); const stream = streamProxy(mockModel, mockContext, { proxyUrl: "http://localhost:0", authToken: "test", + fetch: fetchMock, }); const collected = await collectEvents(stream); diff --git a/packages/agent/test/remote-compaction.test.ts b/packages/agent/test/remote-compaction.test.ts index 61d9c279f..693366cc2 100644 --- a/packages/agent/test/remote-compaction.test.ts +++ b/packages/agent/test/remote-compaction.test.ts @@ -1,7 +1,6 @@ import { describe, expect, test } from "bun:test"; import { buildOpenAiNativeHistory, requestOpenAiRemoteCompaction } from "@oh-my-pi/pi-agent-core/compaction/openai"; -import type { AssistantMessage, Model, ToolResultMessage } from "@oh-my-pi/pi-ai/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { AssistantMessage, FetchImpl, Model, ToolResultMessage } from "@oh-my-pi/pi-ai/types"; function makeOpenAiModel(overrides: Partial> = {}): Model<"openai-responses"> { return { @@ -198,13 +197,13 @@ describe("buildOpenAiNativeHistory call-id tracking", () => { describe("remote compaction input trimming", () => { test("trims custom tool outputs with their matching custom calls", async () => { let requestInput: Array> | undefined; - using _hook = hookFetch(async (_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { const body = JSON.parse(String(init?.body)) as { input: Array> }; requestInput = body.input; return Response.json({ output: [{ type: "compaction_summary", summary: "compact" }], }); - }); + }; await requestOpenAiRemoteCompaction( makeOpenAiModel({ contextWindow: 1 }), @@ -214,6 +213,8 @@ describe("remote compaction input trimming", () => { { type: "custom_tool_call_output", call_id: "call_apply_1", output: "patch applied".repeat(1_000) }, ], "compact", + undefined, + { fetch: fetchMock }, ); expect(requestInput?.some(item => item.type === "custom_tool_call")).toBe(false); @@ -224,7 +225,7 @@ describe("remote compaction input trimming", () => { describe("requestOpenAiRemoteCompaction abort", () => { test("rejects when the abort signal is aborted mid-fetch", async () => { const controller = new AbortController(); - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { // Honor the provided abort signal: hang until aborted, then reject. const signal = init?.signal as AbortSignal | undefined; const { promise, reject } = Promise.withResolvers(); @@ -236,7 +237,7 @@ describe("requestOpenAiRemoteCompaction abort", () => { reject(signal.reason instanceof Error ? signal.reason : new DOMException("Aborted", "AbortError")); }); return promise; - }); + }; const promise = requestOpenAiRemoteCompaction( makeOpenAiModel(), @@ -244,6 +245,7 @@ describe("requestOpenAiRemoteCompaction abort", () => { [{ type: "message", role: "user", content: [{ type: "input_text", text: "hi" }] }], "compact", controller.signal, + { fetch: fetchMock }, ); queueMicrotask(() => controller.abort()); diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index 752aa7375..069b6dfed 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,9 +1,11 @@ # Changelog ## [Unreleased] - ### Added +- Added optional `fetch` transport override (`fetch?: FetchImpl`) to Google, Ollama, and OpenAI-compatible model-manager options so dynamic model discovery and metadata lookups can use a caller-supplied HTTP client instead of only global `fetch` +- Added optional `fetch` on OAuth controller and API-key validation/login flows so token exchange, refresh, and device/PKCE login requests can be routed through a custom `fetch` implementation +- Added optional `fetch` support to usage polling context, allowing usage providers to execute usage checks using an injected HTTP client - Added `AssistantMessage.upstreamProvider`, capturing the upstream provider an aggregator routed the request to (OpenRouter reports it via a top-level `provider` field on every chunk, e.g. `"Anthropic"`). Surfaced from the OpenAI-completions stream alongside `responseId`. ### Fixed @@ -3098,4 +3100,4 @@ _Dedicated to Peter's shoulder ([@steipete](https://twitter.com/steipete))_ ## [0.9.4] - 2025-11-26 -Initial release with multi-provider LLM support. +Initial release with multi-provider LLM support. \ No newline at end of file diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index a99247959..2c1a96958 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -1571,6 +1571,7 @@ export class AuthStorage { onPrompt: ctrl.onPrompt, onManualCodeInput: ctrl.onManualCodeInput ?? manualCodeInput, signal: ctrl.signal, + fetch: ctrl.fetch, }); if (typeof result === "string") { // Some flows (e.g. ollama) return "" to signal that no key was entered. diff --git a/packages/ai/src/provider-models/google.ts b/packages/ai/src/provider-models/google.ts index 226a16047..00383b90f 100644 --- a/packages/ai/src/provider-models/google.ts +++ b/packages/ai/src/provider-models/google.ts @@ -5,6 +5,7 @@ import { fetchGeminiModels } from "../utils/discovery/gemini"; export interface GoogleModelManagerConfig { apiKey?: string; + fetch?: FetchImpl; } export interface GoogleVertexModelManagerConfig { @@ -18,22 +19,36 @@ export interface GoogleVertexModelManagerConfig { export interface GoogleAntigravityModelManagerConfig { oauthToken?: string; endpoint?: string; + fetch?: FetchImpl; } export interface GoogleGeminiCliModelManagerConfig { oauthToken?: string; endpoint?: string; + fetch?: FetchImpl; } const CLOUD_CODE_ASSIST_ENDPOINT = "https://cloudcode-pa.googleapis.com"; +function toDiscoveryFetch(fetchImpl: FetchImpl | undefined): typeof fetch | undefined { + if (!fetchImpl) { + return undefined; + } + return Object.assign( + (input: Parameters[0], init?: Parameters[1]) => fetchImpl(input, init), + { preconnect: fetchImpl.preconnect ?? fetch.preconnect }, + ); +} + export function googleModelManagerOptions( config?: GoogleModelManagerConfig, ): ModelManagerOptions<"google-generative-ai"> { const apiKey = config?.apiKey; return { providerId: "google", - ...(apiKey ? { fetchDynamicModels: () => fetchGeminiModels({ apiKey }) } : undefined), + ...(apiKey + ? { fetchDynamicModels: () => fetchGeminiModels({ apiKey, fetch: toDiscoveryFetch(config?.fetch) }) } + : undefined), }; } @@ -53,6 +68,7 @@ export function googleAntigravityModelManagerOptions( fetchAntigravityDiscoveryModels({ token, endpoint: config?.endpoint, + fetcher: toDiscoveryFetch(config?.fetch), }), } : undefined), @@ -72,6 +88,7 @@ export function googleGeminiCliModelManagerOptions( const models = await fetchAntigravityDiscoveryModels({ token, endpoint, + fetcher: toDiscoveryFetch(config?.fetch), }); if (models === null) { return null; diff --git a/packages/ai/src/provider-models/ollama.ts b/packages/ai/src/provider-models/ollama.ts index ed539d924..9dead83bd 100644 --- a/packages/ai/src/provider-models/ollama.ts +++ b/packages/ai/src/provider-models/ollama.ts @@ -1,12 +1,13 @@ import { fetchWithRetry } from "@oh-my-pi/pi-utils"; import { Effort } from "../effort"; import type { ModelManagerOptions } from "../model-manager"; -import type { ThinkingConfig } from "../types"; +import type { FetchImpl, ThinkingConfig } from "../types"; import { createBundledReferenceMap, createReferenceResolver } from "./bundled-references"; export interface OllamaCloudModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } type OllamaTagEntry = { @@ -65,13 +66,13 @@ function getThinkingConfig(capabilities: string[] | undefined): ThinkingConfig | maxLevel: Effort.High, }; } - async function fetchShowMetadata( baseUrl: string, apiKey: string, model: string, + fetchImpl: FetchImpl = fetch, ): Promise { - const response = await fetch(`${baseUrl}/api/show`, { + const response = await fetchImpl(`${baseUrl}/api/show`, { method: "POST", headers: { ...createCloudHeaders(apiKey), @@ -100,6 +101,7 @@ export function ollamaCloudModelManagerOptions( const response = await fetchWithRetry(`${baseUrl}/api/tags`, { method: "GET", headers: createCloudHeaders(apiKey), + fetch: config?.fetch, defaultDelayMs: OLLAMA_RETRY_DELAYS_MS, }); if (!response.ok) { @@ -116,7 +118,7 @@ export function ollamaCloudModelManagerOptions( const reference = resolveReference(id); let metadata: OllamaShowResponse | undefined; try { - metadata = await fetchShowMetadata(baseUrl, apiKey, id); + metadata = await fetchShowMetadata(baseUrl, apiKey, id, config?.fetch); } catch { metadata = undefined; } diff --git a/packages/ai/src/provider-models/openai-compat.ts b/packages/ai/src/provider-models/openai-compat.ts index 292b0dd8d..ad99ca1df 100644 --- a/packages/ai/src/provider-models/openai-compat.ts +++ b/packages/ai/src/provider-models/openai-compat.ts @@ -2,7 +2,7 @@ import { Effort } from "../effort"; import type { ModelManagerOptions } from "../model-manager"; import { getBundledModels } from "../models"; import { getGitHubCopilotBaseUrl, OPENCODE_HEADERS, parseGitHubCopilotApiKey } from "../registry/oauth/github-copilot"; -import type { Api, Model, Provider, ThinkingConfig } from "../types"; +import type { Api, FetchImpl, Model, Provider, ThinkingConfig } from "../types"; import { isAnthropicOAuthToken, isRecord, toBoolean, toNumber, toPositiveNumber } from "../utils"; import { fetchOpenAICompatibleModels, @@ -56,7 +56,7 @@ function toInputCapabilities(value: unknown): ("text" | "image")[] { return supportsImage ? ["text", "image"] : ["text"]; } -async function fetchModelsDevPayload(fetchImpl: typeof fetch = fetch): Promise { +async function fetchModelsDevPayload(fetchImpl: FetchImpl = fetch): Promise { const response = await fetchImpl(MODELS_DEV_URL, { method: "GET", headers: { Accept: "application/json" }, @@ -195,11 +195,12 @@ function toOllamaNativeBaseUrl(baseUrl: string): string { async function fetchOllamaNativeModels( baseUrl: string, resolveMetadata: (modelId: string) => Promise, + fetchImpl: FetchImpl = fetch, ): Promise[] | null> { const nativeBaseUrl = toOllamaNativeBaseUrl(baseUrl); let response: Response; try { - response = await fetch(`${nativeBaseUrl}/api/tags`, { + response = await fetchImpl(`${nativeBaseUrl}/api/tags`, { method: "GET", headers: { Accept: "application/json" }, }); @@ -318,9 +319,10 @@ function getOllamaThinkingConfig(capabilities: string[] | undefined): ThinkingCo async function fetchOllamaShowMetadata( nativeBaseUrl: string, modelId: string, + fetchImpl: FetchImpl = fetch, ): Promise { try { - const response = await fetch(`${nativeBaseUrl}/api/show`, { + const response = await fetchImpl(`${nativeBaseUrl}/api/show`, { method: "POST", headers: { "Content-Type": "application/json", Accept: "application/json" }, body: JSON.stringify({ model: modelId }), @@ -355,13 +357,16 @@ async function fetchOllamaShowMetadata( * cached so repeated `fetchDynamicModels` calls do not refetch; failed * lookups stay uncached so a later refresh can recover. */ -function createOllamaMetadataResolver(nativeBaseUrl: string): (modelId: string) => Promise { +function createOllamaMetadataResolver( + nativeBaseUrl: string, + fetchImpl?: FetchImpl, +): (modelId: string) => Promise { const cache = new Map>(); return modelId => { const cached = cache.get(modelId); if (cached) return cached; const pending = (async () => { - const metadata = await fetchOllamaShowMetadata(nativeBaseUrl, modelId); + const metadata = await fetchOllamaShowMetadata(nativeBaseUrl, modelId, fetchImpl); if (!metadata) { cache.delete(modelId); return { contextWindow: OLLAMA_FALLBACK_CONTEXT_WINDOW, maxTokens: OLLAMA_DEFAULT_MAX_TOKENS }; @@ -440,7 +445,7 @@ function isLikelyNanoGptTextModelId(id: string): boolean { return !NANO_GPT_NON_TEXT_MODEL_TOKENS.some(token => normalized.includes(token)); } -type SimpleProviderConfig = { apiKey?: string; baseUrl?: string }; +type SimpleProviderConfig = { apiKey?: string; baseUrl?: string; fetch?: FetchImpl }; export function createSimpleOpenAICompletionsOptions( providerId: Parameters[0], @@ -463,6 +468,7 @@ export function createSimpleOpenAICompletionsOptions( const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }), }; @@ -489,6 +495,7 @@ function createSimpleOpenAIResponsesOptions( const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }), }; @@ -521,6 +528,7 @@ function createSimpleAnthropicProviderOptions( baseUrl, }; }, + fetch: config?.fetch, }), }), }; @@ -533,6 +541,7 @@ function createSimpleAnthropicProviderOptions( export interface OpenAIModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function openaiModelManagerOptions(config?: OpenAIModelManagerConfig): ModelManagerOptions<"openai-responses"> { @@ -553,6 +562,7 @@ export function openaiModelManagerOptions(config?: OpenAIModelManagerConfig): Mo const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }), }; @@ -565,6 +575,7 @@ export function openaiModelManagerOptions(config?: OpenAIModelManagerConfig): Mo export interface GroqModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function groqModelManagerOptions(config?: GroqModelManagerConfig): ModelManagerOptions<"openai-completions"> { @@ -578,6 +589,7 @@ export function groqModelManagerOptions(config?: GroqModelManagerConfig): ModelM export interface CerebrasModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function cerebrasModelManagerOptions( @@ -593,6 +605,7 @@ export function cerebrasModelManagerOptions( export interface HuggingfaceModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function huggingfaceModelManagerOptions( @@ -608,6 +621,7 @@ export function huggingfaceModelManagerOptions( export interface NvidiaModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function nvidiaModelManagerOptions( @@ -623,6 +637,7 @@ export function nvidiaModelManagerOptions( export interface XaiModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function xaiModelManagerOptions(config?: XaiModelManagerConfig): ModelManagerOptions<"openai-completions"> { @@ -632,6 +647,7 @@ export function xaiModelManagerOptions(config?: XaiModelManagerConfig): ModelMan export interface XaiOAuthModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } interface XAICuratedModel { @@ -881,6 +897,7 @@ export function isLikelyAimlApiChatModelId(id: string): boolean { export interface AimlApiModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function aimlApiModelManagerOptions( @@ -904,6 +921,7 @@ export function aimlApiModelManagerOptions( const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }), }; @@ -916,6 +934,7 @@ export function aimlApiModelManagerOptions( export interface DeepSeekModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function deepseekModelManagerOptions( @@ -930,6 +949,7 @@ export function deepseekModelManagerOptions( export interface ZhipuCodingPlanModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function zhipuCodingPlanModelManagerOptions( @@ -963,6 +983,7 @@ export function zhipuCodingPlanModelManagerOptions( }, }; }, + fetch: config?.fetch, }), }), }; @@ -1047,6 +1068,7 @@ export function stripFireworksDeepSeekThinkingToggle( export interface FireworksModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } function toFireworksModelName(entry: OpenAICompatibleModelRecord, fallback: string): string { @@ -1082,9 +1104,9 @@ function createModelsDevReferenceMap(models: readonly Model(): Promise>> { +async function loadModelsDevReferences(fetchImpl?: FetchImpl): Promise>> { try { - const payload = await fetchModelsDevPayload(); + const payload = await fetchModelsDevPayload(fetchImpl); return createModelsDevReferenceMap( mapModelsDevToModels(payload as Record, MODELS_DEV_PROVIDER_DESCRIPTORS), ); @@ -1102,7 +1124,7 @@ export function fireworksModelManagerOptions( providerId: "fireworks", ...(apiKey && { fetchDynamicModels: async () => { - const modelsDevReferences = await loadModelsDevReferences<"openai-completions">(); + const modelsDevReferences = await loadModelsDevReferences<"openai-completions">(config?.fetch); return fetchOpenAICompatibleModels({ api: "openai-completions", provider: "fireworks", @@ -1132,6 +1154,7 @@ export function fireworksModelManagerOptions( ), }; }, + fetch: config?.fetch, }); }, }), @@ -1145,6 +1168,7 @@ export function fireworksModelManagerOptions( export interface FirepassModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } /** @@ -1169,6 +1193,7 @@ export function firepassModelManagerOptions( export interface WaferModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } const WAFER_DEFAULT_BASE_URL = "https://pass.wafer.ai/v1"; @@ -1302,6 +1327,7 @@ function createWaferOptions( return wafer?.tier === "pass_included"; }, mapModel: (entry, defaults) => mapWaferModel(providerId, baseUrl, entry, defaults), + fetch: config?.fetch, }), }), }; @@ -1326,6 +1352,7 @@ export function waferServerlessModelManagerOptions( export interface MistralModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function mistralModelManagerOptions( @@ -1341,6 +1368,7 @@ export function mistralModelManagerOptions( export interface OpenCodeModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } function normalizeOpenCodeBasePath(baseUrl: string | undefined, fallbackBasePath: string): string { @@ -1388,6 +1416,7 @@ function openCodeModelManagerOptions( maxTokens: toPositiveNumber(entry.max_completion_tokens, reference.maxTokens), }; }, + fetch: config?.fetch, }), }), }; @@ -1408,6 +1437,7 @@ export function opencodeGoModelManagerOptions(config?: OpenCodeModelManagerConfi export interface OllamaModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function ollamaModelManagerOptions(config?: OllamaModelManagerConfig): ModelManagerOptions<"openai-responses"> { @@ -1415,7 +1445,7 @@ export function ollamaModelManagerOptions(config?: OllamaModelManagerConfig): Mo const baseUrl = normalizeOllamaBaseUrl(config?.baseUrl); const nativeBaseUrl = toOllamaNativeBaseUrl(baseUrl); const references = createBundledReferenceMap<"openai-responses">("ollama" as Parameters[0]); - const resolveMetadata = createOllamaMetadataResolver(nativeBaseUrl); + const resolveMetadata = createOllamaMetadataResolver(nativeBaseUrl, config?.fetch); return { providerId: "ollama", fetchDynamicModels: async () => { @@ -1436,6 +1466,7 @@ export function ollamaModelManagerOptions(config?: OllamaModelManagerConfig): Mo } return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }); if (openAiCompatible && openAiCompatible.length > 0) { await Promise.all( @@ -1454,7 +1485,7 @@ export function ollamaModelManagerOptions(config?: OllamaModelManagerConfig): Mo ); return openAiCompatible; } - const nativeFallback = await fetchOllamaNativeModels(baseUrl, resolveMetadata); + const nativeFallback = await fetchOllamaNativeModels(baseUrl, resolveMetadata, config?.fetch); if (nativeFallback && nativeFallback.length > 0) { for (const model of nativeFallback) applyOllamaReasoningCompat(model); return nativeFallback; @@ -1471,6 +1502,7 @@ export function ollamaModelManagerOptions(config?: OllamaModelManagerConfig): Mo export interface OpenRouterModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function openrouterModelManagerOptions( @@ -1523,6 +1555,7 @@ export function openrouterModelManagerOptions( }), }; }, + fetch: config?.fetch, }), }; } @@ -1594,6 +1627,7 @@ function getZenMuxCacheWritePrice(pricings: Record | undefined) export interface ZenMuxModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function zenmuxModelManagerOptions(config?: ZenMuxModelManagerConfig): ModelManagerOptions { @@ -1630,6 +1664,7 @@ export function zenmuxModelManagerOptions(config?: ZenMuxModelManagerConfig): Mo maxTokens: toPositiveNumber(entry.max_completion_tokens, defaults.maxTokens), }; }, + fetch: config?.fetch, }), }), }; @@ -1642,6 +1677,7 @@ export function zenmuxModelManagerOptions(config?: ZenMuxModelManagerConfig): Mo export interface KiloModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function kiloModelManagerOptions(config?: KiloModelManagerConfig): ModelManagerOptions<"openai-completions"> { @@ -1655,6 +1691,7 @@ export function kiloModelManagerOptions(config?: KiloModelManagerConfig): ModelM provider: "kilo", baseUrl, apiKey, + fetch: config?.fetch, }), }; } @@ -1666,6 +1703,7 @@ export function kiloModelManagerOptions(config?: KiloModelManagerConfig): ModelM export interface AlibabaCodingPlanModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function alibabaCodingPlanModelManagerOptions( @@ -1686,6 +1724,7 @@ export function alibabaCodingPlanModelManagerOptions( const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }; } @@ -1697,6 +1736,7 @@ export function alibabaCodingPlanModelManagerOptions( export interface VercelAiGatewayModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } function normalizeVercelAiGatewayBaseUrls(rawBaseUrl: string | undefined): { baseUrl: string; catalogBaseUrl: string } { @@ -1750,6 +1790,7 @@ export function vercelAiGatewayModelManagerOptions( maxTokens: typeof entry.max_tokens === "number" ? entry.max_tokens : defaults.maxTokens, }; }, + fetch: config?.fetch, }), }; } @@ -1761,6 +1802,7 @@ export function vercelAiGatewayModelManagerOptions( export interface KimiCodeModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function kimiCodeModelManagerOptions( @@ -1801,6 +1843,7 @@ export function kimiCodeModelManagerOptions( }, }; }, + fetch: config?.fetch, }), }), }; @@ -1813,6 +1856,7 @@ export function kimiCodeModelManagerOptions( export interface LmStudioModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function lmStudioModelManagerOptions( @@ -1833,6 +1877,7 @@ export function lmStudioModelManagerOptions( const reference = references.get(defaults.id); return mapWithBundledReference(entry, defaults, reference); }, + fetch: config?.fetch, }), }; } @@ -1844,6 +1889,7 @@ export function lmStudioModelManagerOptions( export interface SyntheticModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function syntheticModelManagerOptions( @@ -1883,6 +1929,7 @@ export function syntheticModelManagerOptions( maxTokens: toPositiveNumber(entry.max_tokens, reference?.maxTokens ?? 8192), }; }, + fetch: config?.fetch, }), }), }; @@ -1895,6 +1942,7 @@ export function syntheticModelManagerOptions( export interface VeniceModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function veniceModelManagerOptions( @@ -1919,6 +1967,7 @@ export function veniceModelManagerOptions( compat: { ...model.compat, supportsUsageInStreaming: false }, }; }, + fetch: config?.fetch, }), }; } @@ -1930,6 +1979,7 @@ export function veniceModelManagerOptions( export interface TogetherModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function togetherModelManagerOptions( @@ -1945,6 +1995,7 @@ export function togetherModelManagerOptions( export interface MoonshotModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function moonshotModelManagerOptions( @@ -1985,6 +2036,7 @@ export function moonshotModelManagerOptions( : undefined), }; }, + fetch: config?.fetch, }), }), }; @@ -1997,6 +2049,7 @@ export function moonshotModelManagerOptions( export interface QwenPortalModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function qwenPortalModelManagerOptions( @@ -2012,6 +2065,7 @@ export function qwenPortalModelManagerOptions( export interface QianfanModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function qianfanModelManagerOptions( @@ -2027,6 +2081,7 @@ export function qianfanModelManagerOptions( export interface CloudflareAiGatewayModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function cloudflareAiGatewayModelManagerOptions( @@ -2050,6 +2105,7 @@ export type XiaomiTokenPlanRegion = "sgp" | "ams" | "cn"; export interface XiaomiModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; providerId?: Provider; tokenPlanRegion?: XiaomiTokenPlanRegion; } @@ -2100,6 +2156,7 @@ export function xiaomiModelManagerOptions( name: toModelName(entry.display_name, model.name), }; }, + fetch: config?.fetch, }); return { providerId, @@ -2124,6 +2181,7 @@ export function xiaomiModelManagerOptions( export interface LiteLLMModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function litellmModelManagerOptions( @@ -2139,7 +2197,7 @@ export function litellmModelManagerOptions( // enrich discovered ids against models.dev — the same reference source the // gateway providers (fireworks et al.) use — instead of a bundled map. fetchDynamicModels: async () => { - const modelsDevReferences = await loadModelsDevReferences<"openai-completions">(); + const modelsDevReferences = await loadModelsDevReferences<"openai-completions">(config?.fetch); return fetchOpenAICompatibleModels({ api: "openai-completions", provider: "litellm", @@ -2147,6 +2205,7 @@ export function litellmModelManagerOptions( apiKey, mapModel: (entry, defaults) => mapWithBundledReference(entry, defaults, modelsDevReferences.get(defaults.id)), + fetch: config?.fetch, }); }, }; @@ -2159,6 +2218,7 @@ export function litellmModelManagerOptions( export interface VllmModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function vllmModelManagerOptions(config?: VllmModelManagerConfig): ModelManagerOptions<"openai-completions"> { @@ -2180,6 +2240,7 @@ export function vllmModelManagerOptions(config?: VllmModelManagerConfig): ModelM contextWindow: toPositiveNumber(entry.max_model_len, model.contextWindow), }; }, + fetch: config?.fetch, }), }; } @@ -2191,6 +2252,7 @@ export function vllmModelManagerOptions(config?: VllmModelManagerConfig): ModelM export interface NanoGptModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function nanoGptModelManagerOptions( @@ -2225,6 +2287,7 @@ export function nanoGptModelManagerOptions( } return isLikelyNanoGptTextModelId(model.id); }, + fetch: config?.fetch, }); if (!models) return null; // Mark base models as reasoning-capable when a :thinking variant existed. @@ -2246,6 +2309,7 @@ export function nanoGptModelManagerOptions( export interface GithubCopilotModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } function inferCopilotApi(modelId: string): Api { @@ -2376,6 +2440,7 @@ export function githubCopilotModelManagerOptions(config?: GithubCopilotModelMana : {}), }; }, + fetch: config?.fetch, }), }), }; @@ -2388,6 +2453,7 @@ export function githubCopilotModelManagerOptions(config?: GithubCopilotModelMana export interface AnthropicModelManagerConfig { apiKey?: string; baseUrl?: string; + fetch?: FetchImpl; } export function anthropicModelManagerOptions( @@ -2398,12 +2464,12 @@ export function anthropicModelManagerOptions( return { providerId: "anthropic", modelsDev: { - fetch: fetchModelsDevPayload, + fetch: () => fetchModelsDevPayload(config?.fetch), map: payload => mapAnthropicModelsDev(payload, baseUrl), }, ...(apiKey && { fetchDynamicModels: async () => { - const modelsDevModels = await fetchModelsDevPayload() + const modelsDevModels = await fetchModelsDevPayload(config?.fetch) .then(payload => mapAnthropicModelsDev(payload, baseUrl)) .catch(() => []); const references = buildAnthropicReferenceMap(modelsDevModels); @@ -2435,6 +2501,7 @@ export function anthropicModelManagerOptions( baseUrl, }; }, + fetch: config?.fetch, }) ?? null ); }, diff --git a/packages/ai/src/registry/api-key-login.ts b/packages/ai/src/registry/api-key-login.ts index 514b6e15a..f11d0e33c 100644 --- a/packages/ai/src/registry/api-key-login.ts +++ b/packages/ai/src/registry/api-key-login.ts @@ -82,6 +82,7 @@ export function createApiKeyLogin(config: ApiKeyLoginConfig): (options: OAuthCon baseUrl: config.validation.baseUrl, model: config.validation.model, signal: options.signal, + fetch: options.fetch, }); } else if (config.validation.kind === "anthropic-messages") { await validateAnthropicCompatibleApiKey({ @@ -90,6 +91,7 @@ export function createApiKeyLogin(config: ApiKeyLoginConfig): (options: OAuthCon baseUrl: config.validation.baseUrl, model: config.validation.model, signal: options.signal, + fetch: options.fetch, }); } else { await validateApiKeyAgainstModelsEndpoint({ @@ -97,6 +99,7 @@ export function createApiKeyLogin(config: ApiKeyLoginConfig): (options: OAuthCon apiKey: trimmed, modelsUrl: config.validation.modelsUrl, signal: options.signal, + fetch: options.fetch, }); } } diff --git a/packages/ai/src/registry/api-key-validation.ts b/packages/ai/src/registry/api-key-validation.ts index e1879036f..daac50146 100644 --- a/packages/ai/src/registry/api-key-validation.ts +++ b/packages/ai/src/registry/api-key-validation.ts @@ -1,9 +1,12 @@ +import type { FetchImpl } from "../types"; + type OpenAICompatibleValidationOptions = { provider: string; apiKey: string; baseUrl: string; model: string; signal?: AbortSignal; + fetch?: FetchImpl; }; type AnthropicCompatibleValidationOptions = { provider: string; @@ -11,6 +14,7 @@ type AnthropicCompatibleValidationOptions = { baseUrl: string; model: string; signal?: AbortSignal; + fetch?: FetchImpl; }; type ModelListValidationOptions = { @@ -18,6 +22,7 @@ type ModelListValidationOptions = { apiKey: string; modelsUrl: string; signal?: AbortSignal; + fetch?: FetchImpl; }; const VALIDATION_TIMEOUT_MS = 15_000; @@ -35,8 +40,9 @@ function normalizeAnthropicCompatibleBaseUrl(baseUrl: string): string { export async function validateOpenAICompatibleApiKey(options: OpenAICompatibleValidationOptions): Promise { const timeoutSignal = AbortSignal.timeout(VALIDATION_TIMEOUT_MS); const signal = options.signal ? AbortSignal.any([options.signal, timeoutSignal]) : timeoutSignal; + const fetchImpl = options.fetch ?? fetch; - const response = await fetch(`${options.baseUrl}/chat/completions`, { + const response = await fetchImpl(`${options.baseUrl}/chat/completions`, { method: "POST", headers: { "Content-Type": "application/json", @@ -75,8 +81,9 @@ export async function validateAnthropicCompatibleApiKey(options: AnthropicCompat const timeoutSignal = AbortSignal.timeout(VALIDATION_TIMEOUT_MS); const signal = options.signal ? AbortSignal.any([options.signal, timeoutSignal]) : timeoutSignal; const baseUrl = normalizeAnthropicCompatibleBaseUrl(options.baseUrl); + const fetchImpl = options.fetch ?? fetch; - const response = await fetch(`${baseUrl}/v1/messages`, { + const response = await fetchImpl(`${baseUrl}/v1/messages`, { method: "POST", headers: { "Content-Type": "application/json", @@ -117,8 +124,9 @@ export async function validateAnthropicCompatibleApiKey(options: AnthropicCompat export async function validateApiKeyAgainstModelsEndpoint(options: ModelListValidationOptions): Promise { const timeoutSignal = AbortSignal.timeout(VALIDATION_TIMEOUT_MS); const signal = options.signal ? AbortSignal.any([options.signal, timeoutSignal]) : timeoutSignal; + const fetchImpl = options.fetch ?? fetch; - const response = await fetch(options.modelsUrl, { + const response = await fetchImpl(options.modelsUrl, { method: "GET", headers: { Authorization: `Bearer ${options.apiKey}`, diff --git a/packages/ai/src/registry/kilo.ts b/packages/ai/src/registry/kilo.ts index 2c593554a..5ba36224b 100644 --- a/packages/ai/src/registry/kilo.ts +++ b/packages/ai/src/registry/kilo.ts @@ -18,7 +18,8 @@ interface KiloDeviceAuthPollResponse { } export async function loginKilo(callbacks: OAuthController): Promise { - const initiateResponse = await fetch(`${KILO_DEVICE_AUTH_BASE_URL}/codes`, { + const fetchImpl = callbacks.fetch ?? fetch; + const initiateResponse = await fetchImpl(`${KILO_DEVICE_AUTH_BASE_URL}/codes`, { method: "POST", headers: { "Content-Type": "application/json" }, }); @@ -49,7 +50,7 @@ export async function loginKilo(callbacks: OAuthController): Promise { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -58,9 +55,10 @@ describe("refreshXAIOAuthToken", () => { const fetchMock = vi.fn(async () => { throw new Error("fetch should not be called when refresh_token is empty"); }); - global.fetch = fetchMock as unknown as typeof fetch; - await expect(refreshXAIOAuthToken("")).rejects.toThrow(/missing refresh_token/); + await expect(refreshXAIOAuthToken("", fetchMock as unknown as typeof fetch)).rejects.toThrow( + /missing refresh_token/, + ); expect(fetchMock).not.toHaveBeenCalled(); }); }); @@ -95,9 +93,8 @@ describe("XAIOAuthFlow.exchangeToken", () => { headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new XAIOAuthFlow({}); + const flow = new XAIOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-abc", "http://127.0.0.1:56121/callback"); await expect(flow.exchangeToken("code-xyz", "state-abc", "http://127.0.0.1:56121/callback")).rejects.toThrow( diff --git a/packages/ai/src/registry/oauth/anthropic.ts b/packages/ai/src/registry/oauth/anthropic.ts index 045afcf4f..22d530ee8 100644 --- a/packages/ai/src/registry/oauth/anthropic.ts +++ b/packages/ai/src/registry/oauth/anthropic.ts @@ -1,6 +1,8 @@ /** * Anthropic OAuth flow (Claude Pro/Max) */ + +import type { FetchImpl } from "../../types"; import { OAuthCallbackFlow } from "./callback-server"; import { generatePKCE } from "./pkce"; import type { OAuthController, OAuthCredentials } from "./types"; @@ -40,9 +42,10 @@ function formatErrorDetails(error: unknown): string { async function postJson( url: string, body: Record, + fetchImpl: FetchImpl, extraHeaders?: Record, ): Promise { - const response = await fetch(url, { + const response = await fetchImpl(url, { method: "POST", headers: { // No Accept header: CC omits it on OAuth token requests. @@ -109,9 +112,12 @@ function extractAccountFromTokenResponse(data: AnthropicTokenResponse): { }; } -async function fetchBootstrapIdentity(accessToken: string): Promise<{ accountId?: string; email?: string }> { +async function fetchBootstrapIdentity( + accessToken: string, + fetchImpl: FetchImpl, +): Promise<{ accountId?: string; email?: string }> { const url = `${BOOTSTRAP_URL}?entrypoint=cli&model=${encodeURIComponent(CLAUDE_CODE_BOOTSTRAP_MODEL)}`; - const response = await fetch(url, { + const response = await fetchImpl(url, { method: "GET", headers: { Accept: "application/json, text/plain, */*", @@ -142,11 +148,14 @@ async function fetchBootstrapIdentity(accessToken: string): Promise<{ accountId? }; } -async function resolveAccountIdentity(data: AnthropicTokenResponse): Promise<{ accountId?: string; email?: string }> { +async function resolveAccountIdentity( + data: AnthropicTokenResponse, + fetchImpl: FetchImpl, +): Promise<{ accountId?: string; email?: string }> { const identity = extractAccountFromTokenResponse(data); if (identity.accountId && identity.email) return identity; try { - const bootstrap = await fetchBootstrapIdentity(data.access_token); + const bootstrap = await fetchBootstrapIdentity(data.access_token, fetchImpl); return { accountId: identity.accountId ?? bootstrap.accountId, email: identity.email ?? bootstrap.email, @@ -159,9 +168,11 @@ async function resolveAccountIdentity(data: AnthropicTokenResponse): Promise<{ a export class AnthropicOAuthFlow extends OAuthCallbackFlow { #verifier: string = ""; #challenge: string = ""; + #fetch: FetchImpl; constructor(ctrl: OAuthController) { super(ctrl, CALLBACK_PORT, CALLBACK_PATH); + this.#fetch = ctrl.fetch ?? fetch; } async generateAuthUrl(state: string, redirectUri: string): Promise<{ url: string; instructions?: string }> { @@ -202,14 +213,18 @@ export class AnthropicOAuthFlow extends OAuthCallbackFlow { let responseBody: string; try { - responseBody = await postJson(TOKEN_URL, { - grant_type: "authorization_code", - client_id: CLIENT_ID, - code: exchangeCode, - state: exchangeState, - redirect_uri: redirectUri, - code_verifier: this.#verifier, - }); + responseBody = await postJson( + TOKEN_URL, + { + grant_type: "authorization_code", + client_id: CLIENT_ID, + code: exchangeCode, + state: exchangeState, + redirect_uri: redirectUri, + code_verifier: this.#verifier, + }, + this.#fetch, + ); } catch (error) { throw new Error( `Token exchange request failed. url=${TOKEN_URL}; redirect_uri=${redirectUri}; response_type=authorization_code; details=${formatErrorDetails(error)}`, @@ -217,7 +232,7 @@ export class AnthropicOAuthFlow extends OAuthCallbackFlow { } const tokenData = parseOAuthTokenResponse(responseBody, "token exchange"); - const { accountId, email } = await resolveAccountIdentity(tokenData); + const { accountId, email } = await resolveAccountIdentity(tokenData, this.#fetch); return { refresh: tokenData.refresh_token, @@ -240,7 +255,11 @@ export async function loginAnthropic(ctrl: OAuthController): Promise { +export async function refreshAnthropicToken( + refreshToken: string, + fetchOverride?: FetchImpl, +): Promise { + const fetchImpl = fetchOverride ?? fetch; let responseBody: string; try { responseBody = await postJson( @@ -250,6 +269,7 @@ export async function refreshAnthropicToken(refreshToken: string): Promise void; + onPrompt: (prompt: { message: string; placeholder?: string; allowEmpty?: boolean }) => Promise; + onProgress?: (message: string) => void; + signal?: AbortSignal; + pollIntervalFloorMs?: number; + pollIntervalScaleMs?: number; + fetch?: FetchImpl; +}; type DeviceCodeResponse = { device_code: string; user_code: string; @@ -106,8 +117,8 @@ export function getGitHubCopilotBaseUrl(enterpriseDomain?: string): string { return `https://${host}`; } -async function fetchJson(url: string, init: RequestInit): Promise { - const response = await fetch(url, init); +async function fetchJson(url: string, init: RequestInit, fetchImpl: FetchImpl): Promise { + const response = await fetchImpl(url, init); if (!response.ok) { const text = await response.text(); throw new Error(`${response.status} ${response.statusText}: ${text}`); @@ -115,20 +126,24 @@ async function fetchJson(url: string, init: RequestInit): Promise { return response.json(); } -async function startDeviceFlow(domain: string): Promise { +async function startDeviceFlow(domain: string, fetchImpl: FetchImpl): Promise { const urls = getUrls(domain); - const data = await fetchJson(urls.deviceCodeUrl, { - method: "POST", - headers: { - Accept: "application/json", - "Content-Type": "application/json", - ...OPENCODE_HEADERS, + const data = await fetchJson( + urls.deviceCodeUrl, + { + method: "POST", + headers: { + Accept: "application/json", + "Content-Type": "application/json", + ...OPENCODE_HEADERS, + }, + body: JSON.stringify({ + client_id: CLIENT_ID, + scope: "read:user", + }), }, - body: JSON.stringify({ - client_id: CLIENT_ID, - scope: "read:user", - }), - }); + fetchImpl, + ); if (!data || typeof data !== "object") { throw new Error("Invalid device code response"); @@ -164,7 +179,8 @@ async function pollForGitHubAccessToken( deviceCode: string, intervalSeconds: number, expiresIn: number, - signal?: AbortSignal, + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, pollIntervalFloorMs = 1000, pollIntervalScaleMs = 1000, ) { @@ -187,19 +203,23 @@ async function pollForGitHubAccessToken( throw new Error("Login cancelled"); } - const raw = await fetchJson(urls.accessTokenUrl, { - method: "POST", - headers: { - Accept: "application/json", - "Content-Type": "application/json", - ...OPENCODE_HEADERS, + const raw = await fetchJson( + urls.accessTokenUrl, + { + method: "POST", + headers: { + Accept: "application/json", + "Content-Type": "application/json", + ...OPENCODE_HEADERS, + }, + body: JSON.stringify({ + client_id: CLIENT_ID, + device_code: deviceCode, + grant_type: "urn:ietf:params:oauth:grant-type:device_code", + }), }, - body: JSON.stringify({ - client_id: CLIENT_ID, - device_code: deviceCode, - grant_type: "urn:ietf:params:oauth:grant-type:device_code", - }), - }); + fetchImpl, + ); if (raw && typeof raw === "object" && typeof (raw as DeviceTokenSuccessResponse).access_token === "string") { return (raw as DeviceTokenSuccessResponse).access_token; @@ -255,12 +275,17 @@ export function refreshGitHubCopilotToken(refreshToken: string, enterpriseDomain * Enable a model for the user's GitHub Copilot account. * This is required for some models (like Claude, Grok) before they can be used. */ -async function enableGitHubCopilotModel(token: string, modelId: string, enterpriseDomain?: string): Promise { +async function enableGitHubCopilotModel( + token: string, + modelId: string, + fetchImpl: FetchImpl, + enterpriseDomain?: string, +): Promise { const baseUrl = getGitHubCopilotBaseUrl(enterpriseDomain); const url = `${baseUrl}/models/${modelId}/policy`; try { - const response = await fetch(url, { + const response = await fetchImpl(url, { method: "POST", headers: { "Content-Type": "application/json", @@ -283,7 +308,8 @@ async function enableGitHubCopilotModel(token: string, modelId: string, enterpri */ async function enableAllGitHubCopilotModels( token: string, - enterpriseDomain?: string, + enterpriseDomain: string | undefined, + fetchImpl: FetchImpl, onProgress?: (model: string, success: boolean) => void, ): Promise { const models = getBundledModels("github-copilot"); @@ -292,7 +318,7 @@ async function enableAllGitHubCopilotModels( const batch = models.slice(i, i + BATCH_SIZE); await Promise.all( batch.map(async model => { - const success = await enableGitHubCopilotModel(token, model.id, enterpriseDomain); + const success = await enableGitHubCopilotModel(token, model.id, fetchImpl, enterpriseDomain); onProgress?.(model.id, success); }), ); @@ -307,14 +333,8 @@ async function enableAllGitHubCopilotModels( * @param options.onProgress - Optional progress callback * @param options.signal - Optional AbortSignal for cancellation */ -export async function loginGitHubCopilot(options: { - onAuth: (url: string, instructions?: string) => void; - onPrompt: (prompt: { message: string; placeholder?: string; allowEmpty?: boolean }) => Promise; - onProgress?: (message: string) => void; - signal?: AbortSignal; - pollIntervalFloorMs?: number; - pollIntervalScaleMs?: number; -}): Promise { +export async function loginGitHubCopilot(options: GitHubCopilotLoginOptions): Promise { + const fetchImpl = options.fetch ?? fetch; const input = await options.onPrompt({ message: "GitHub Enterprise URL/domain (blank for github.com)", placeholder: "company.ghe.com", @@ -334,7 +354,7 @@ export async function loginGitHubCopilot(options: { const domain = normalizedDomain && isPublicGitHubHost(normalizedDomain) ? "github.com" : (normalizedDomain ?? "github.com"); - const device = await startDeviceFlow(domain); + const device = await startDeviceFlow(domain, fetchImpl); options.onAuth(device.verification_uri, `Enter code: ${device.user_code}`); const githubAccessToken = await pollForGitHubAccessToken( @@ -343,6 +363,7 @@ export async function loginGitHubCopilot(options: { device.interval, device.expires_in, options.signal, + fetchImpl, options.pollIntervalFloorMs, options.pollIntervalScaleMs, ); @@ -357,6 +378,6 @@ export async function loginGitHubCopilot(options: { // Enable all models after successful login options.onProgress?.("Enabling models..."); - await enableAllGitHubCopilotModels(githubAccessToken, enterpriseDomain ?? undefined); + await enableAllGitHubCopilotModels(githubAccessToken, enterpriseDomain ?? undefined, fetchImpl); return credentials; } diff --git a/packages/ai/src/registry/oauth/minimax-code.ts b/packages/ai/src/registry/oauth/minimax-code.ts index 23e7db2f7..0101712ce 100644 --- a/packages/ai/src/registry/oauth/minimax-code.ts +++ b/packages/ai/src/registry/oauth/minimax-code.ts @@ -38,6 +38,7 @@ async function loginMiniMaxCodeWithBaseUrl( baseUrl: string, providerName: string, ): Promise { + const fetchImpl = options.fetch ?? fetch; if (!options.onPrompt) { throw new Error("MiniMax Coding Plan login requires onPrompt callback"); } @@ -66,6 +67,7 @@ async function loginMiniMaxCodeWithBaseUrl( baseUrl, model: VALIDATION_MODEL, signal: options.signal, + fetch: fetchImpl, }); return trimmed; } diff --git a/packages/ai/src/registry/oauth/types.ts b/packages/ai/src/registry/oauth/types.ts index e7006c1cb..8cc57daba 100644 --- a/packages/ai/src/registry/oauth/types.ts +++ b/packages/ai/src/registry/oauth/types.ts @@ -1,3 +1,4 @@ +import type { FetchImpl } from "../../types"; import type { OAuthProviderUnion } from "../registry"; export type OAuthCredentials = { @@ -37,6 +38,7 @@ export interface OAuthController { onManualCodeInput?(): Promise; onPrompt?(prompt: OAuthPrompt): Promise; signal?: AbortSignal; + fetch?: FetchImpl; } export interface OAuthLoginCallbacks extends OAuthController { diff --git a/packages/ai/src/registry/oauth/xai-oauth.ts b/packages/ai/src/registry/oauth/xai-oauth.ts index 5eb85552a..1806d05e6 100644 --- a/packages/ai/src/registry/oauth/xai-oauth.ts +++ b/packages/ai/src/registry/oauth/xai-oauth.ts @@ -10,6 +10,7 @@ * rejected on every call site, not just the first. */ +import type { FetchImpl } from "../../types"; import { OAuthCallbackFlow, type OAuthCallbackFlowOptions } from "./callback-server"; import { generatePKCE } from "./pkce"; import type { OAuthController, OAuthCredentials } from "./types"; @@ -70,10 +71,14 @@ export function validateXAIEndpoint(url: string, field: string): string { * * Hermes `_xai_oauth_discovery` L3038-3084. */ -async function xaiOAuthDiscovery(timeoutMs: number = DISCOVERY_TIMEOUT_MS): Promise { +async function xaiOAuthDiscovery( + timeoutMs: number = DISCOVERY_TIMEOUT_MS, + fetchOverride?: FetchImpl, +): Promise { + const fetchImpl = fetchOverride ?? fetch; let response: Response; try { - response = await fetch(XAI_OAUTH_DISCOVERY_URL, { + response = await fetchImpl(XAI_OAUTH_DISCOVERY_URL, { method: "GET", headers: { Accept: "application/json" }, signal: AbortSignal.timeout(timeoutMs), @@ -179,6 +184,7 @@ function buildXAIAuthorizeUrl(opts: BuildXAIAuthorizeUrlOptions): string { */ export class XAIOAuthFlow extends OAuthCallbackFlow { #verifier: string = ""; + #fetch: FetchImpl; constructor(ctrl: OAuthController) { super(ctrl, { @@ -187,6 +193,7 @@ export class XAIOAuthFlow extends OAuthCallbackFlow { callbackHostname: XAI_OAUTH_REDIRECT_HOST, redirectUri: `http://${XAI_OAUTH_REDIRECT_HOST}:${XAI_OAUTH_REDIRECT_PORT}${XAI_OAUTH_REDIRECT_PATH}`, } satisfies OAuthCallbackFlowOptions); + this.#fetch = ctrl.fetch ?? fetch; } async generateAuthUrl(state: string, redirectUri: string): Promise<{ url: string; instructions?: string }> { @@ -194,7 +201,7 @@ export class XAIOAuthFlow extends OAuthCallbackFlow { this.#verifier = pkce.verifier; const nonce = crypto.randomUUID().replace(/-/g, ""); - const discovery = await xaiOAuthDiscovery(); + const discovery = await xaiOAuthDiscovery(DISCOVERY_TIMEOUT_MS, this.#fetch); const url = buildXAIAuthorizeUrl({ authorizationEndpoint: discovery.authorization_endpoint, redirectUri, @@ -210,7 +217,7 @@ export class XAIOAuthFlow extends OAuthCallbackFlow { } async exchangeToken(code: string, _state: string, redirectUri: string): Promise { - const discovery = await xaiOAuthDiscovery(); + const discovery = await xaiOAuthDiscovery(DISCOVERY_TIMEOUT_MS, this.#fetch); const tokenEndpoint = validateXAIEndpoint(discovery.token_endpoint, "token_endpoint"); const body = new URLSearchParams({ @@ -221,7 +228,7 @@ export class XAIOAuthFlow extends OAuthCallbackFlow { code_verifier: this.#verifier, }); - const response = await fetch(tokenEndpoint, { + const response = await this.#fetch(tokenEndpoint, { method: "POST", headers: { "Content-Type": "application/x-www-form-urlencoded", @@ -282,12 +289,13 @@ export async function loginXAIOAuth(ctrl: OAuthController): Promise { +export async function refreshXAIOAuthToken(refreshToken: string, fetchOverride?: FetchImpl): Promise { + const fetchImpl = fetchOverride ?? fetch; if (typeof refreshToken !== "string" || !refreshToken.trim()) { throw new Error("missing refresh_token"); } - const discovery = await xaiOAuthDiscovery(); + const discovery = await xaiOAuthDiscovery(DISCOVERY_TIMEOUT_MS, fetchImpl); const tokenEndpoint = validateXAIEndpoint(discovery.token_endpoint, "token_endpoint"); const body = new URLSearchParams({ @@ -296,7 +304,7 @@ export async function refreshXAIOAuthToken(refreshToken: string): Promise { + const fetchImpl = fetchOverride ?? fetch; // Region-specific Token Plan logins must validate against the selected // cluster. Generic Xiaomi login keeps the historical SGP → AMS → CN fallback. const endpoints = tokenPlanRegion @@ -74,7 +77,7 @@ async function validateXiaomiApiKey( const timeoutSignal = AbortSignal.timeout(VALIDATION_TIMEOUT_MS); const requestSignal = signal ? AbortSignal.any([signal, timeoutSignal]) : timeoutSignal; try { - const response = await fetch(`${ep.baseUrl}/chat/completions`, { + const response = await fetchImpl(`${ep.baseUrl}/chat/completions`, { method: "POST", headers: { "Content-Type": "application/json", @@ -139,6 +142,7 @@ async function validateXiaomiApiKey( * Returns the API key directly (not OAuthCredentials - this isn't OAuth). */ export async function loginXiaomi(options: OAuthController): Promise { + const fetchImpl = options.fetch ?? fetch; if (!options.onPrompt) { throw new Error(`${PROVIDER_NAME} login requires onPrompt callback`); } @@ -159,7 +163,7 @@ export async function loginXiaomi(options: OAuthController): Promise { } options.onProgress?.(`Validating ${PROVIDER_ID} API key...`); - await validateXiaomiApiKey(trimmed, options.signal); + await validateXiaomiApiKey(trimmed, undefined, options.signal, fetchImpl); return trimmed; } @@ -169,6 +173,7 @@ export async function loginXiaomi(options: OAuthController): Promise { * Prompts for a token-plan API key and validates it against the selected region. */ export async function loginXiaomiTokenPlan(options: OAuthController, region: XiaomiTokenPlanRegion): Promise { + const fetchImpl = options.fetch ?? fetch; if (!options.onPrompt) { throw new Error(`Xiaomi Token Plan (${TOKEN_PLAN_REGION_NAMES[region]}) login requires onPrompt callback`); } @@ -189,6 +194,6 @@ export async function loginXiaomiTokenPlan(options: OAuthController, region: Xia } options.onProgress?.(`Validating Xiaomi Token Plan (${TOKEN_PLAN_REGION_NAMES[region]}) API key...`); - await validateXiaomiApiKey(trimmed, options.signal, region); + await validateXiaomiApiKey(trimmed, region, options.signal, fetchImpl); return trimmed; } diff --git a/packages/ai/src/registry/types.ts b/packages/ai/src/registry/types.ts index 1a65244fa..35467bccb 100644 --- a/packages/ai/src/registry/types.ts +++ b/packages/ai/src/registry/types.ts @@ -9,11 +9,11 @@ * line in `./registry.ts`. */ import type { ModelManagerOptions } from "../model-manager"; -import type { Api } from "../types"; +import type { Api, FetchImpl } from "../types"; import type { OAuthCredentials, OAuthLoginCallbacks } from "./oauth/types"; /** Config passed to a provider's runtime model-manager factory. */ -export type ModelManagerConfig = { apiKey?: string; baseUrl?: string }; +export type ModelManagerConfig = { apiKey?: string; baseUrl?: string; fetch?: FetchImpl }; /** * API-key environment fallback: either a single env var name (e.g. diff --git a/packages/ai/src/usage.ts b/packages/ai/src/usage.ts index b49fc2868..e348afb96 100644 --- a/packages/ai/src/usage.ts +++ b/packages/ai/src/usage.ts @@ -5,7 +5,7 @@ * and shared quotas across providers. */ import * as z from "zod/v4"; -import type { Provider } from "./types"; +import type { FetchImpl, Provider } from "./types"; export type UsageUnit = "percent" | "tokens" | "requests" | "usd" | "minutes" | "bytes" | "unknown"; export type UsageStatus = "ok" | "warning" | "exhausted" | "unknown"; @@ -154,7 +154,7 @@ export interface UsageFetchParams { /** Shared runtime utilities for fetchers. */ export interface UsageFetchContext { - fetch: typeof fetch; + fetch: FetchImpl; logger?: UsageLogger; retryWait?: (delayMs: number, signal?: AbortSignal) => Promise; } diff --git a/packages/ai/src/utils/discovery/gemini.ts b/packages/ai/src/utils/discovery/gemini.ts index 09e0163fc..c1c0c27f0 100644 --- a/packages/ai/src/utils/discovery/gemini.ts +++ b/packages/ai/src/utils/discovery/gemini.ts @@ -1,7 +1,7 @@ import * as z from "zod/v4"; import { getBundledModels } from "../../models"; import { UNK_CONTEXT_WINDOW, UNK_MAX_TOKENS } from "../../provider-models/discovery-constants"; -import type { Model } from "../../types"; +import type { FetchImpl, Model } from "../../types"; const GOOGLE_GENERATIVE_AI_BASE_URL = "https://generativelanguage.googleapis.com/v1beta"; const DEFAULT_PAGE_SIZE = 100; @@ -52,7 +52,7 @@ export interface GeminiDiscoveryOptions { /** Optional abort signal for HTTP requests. */ signal?: AbortSignal; /** Optional fetch implementation override for tests. */ - fetch?: typeof fetch; + fetch?: FetchImpl; } /** diff --git a/packages/ai/src/utils/discovery/openai-compatible.ts b/packages/ai/src/utils/discovery/openai-compatible.ts index 3dc654c75..24e74afd8 100644 --- a/packages/ai/src/utils/discovery/openai-compatible.ts +++ b/packages/ai/src/utils/discovery/openai-compatible.ts @@ -1,6 +1,6 @@ import * as z from "zod/v4"; import { UNK_CONTEXT_WINDOW, UNK_MAX_TOKENS } from "../../provider-models/discovery-constants"; -import type { Api, Model, Provider } from "../../types"; +import type { Api, FetchImpl, Model, Provider } from "../../types"; const MODELS_PATH = "/models"; @@ -81,7 +81,7 @@ export interface FetchOpenAICompatibleModelsOptions { /** Optional AbortSignal for request cancellation. */ signal?: AbortSignal; /** Optional fetch implementation override for testing/custom runtimes. */ - fetch?: typeof globalThis.fetch; + fetch?: FetchImpl; /** * Optional post-normalization filter. * Return false to skip a model. diff --git a/packages/ai/test/anthropic-client.test.ts b/packages/ai/test/anthropic-client.test.ts index 259b13b8a..46abba978 100644 --- a/packages/ai/test/anthropic-client.test.ts +++ b/packages/ai/test/anthropic-client.test.ts @@ -5,6 +5,7 @@ import { AnthropicMessagesClient, } from "@oh-my-pi/pi-ai/providers/anthropic-client"; import type { MessageCreateParamsStreaming } from "@oh-my-pi/pi-ai/providers/anthropic-wire"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const params: MessageCreateParamsStreaming = { model: "claude-sonnet-4-5", @@ -15,7 +16,7 @@ const params: MessageCreateParamsStreaming = { type FetchCall = { url: string; init: RequestInit }; -function createFetchMock(responses: Array): { calls: FetchCall[]; fetch: typeof fetch } { +function createFetchMock(responses: Array): { calls: FetchCall[]; fetch: FetchImpl } { const calls: FetchCall[] = []; const fetchImpl = (async (input: string | URL | Request, init?: RequestInit) => { calls.push({ url: String(input), init: init ?? {} }); diff --git a/packages/ai/test/anthropic-oauth.test.ts b/packages/ai/test/anthropic-oauth.test.ts index 6b6e91c50..30322edc2 100644 --- a/packages/ai/test/anthropic-oauth.test.ts +++ b/packages/ai/test/anthropic-oauth.test.ts @@ -7,10 +7,7 @@ import { } from "@oh-my-pi/pi-ai/utils/anthropic-auth"; import { withEnv } from "./helpers"; -const originalFetch = global.fetch; - afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -49,9 +46,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-123", "http://localhost:54545/callback"); const result = await flow.exchangeToken("code-123", "state-123", "http://localhost:54545/callback"); @@ -80,9 +76,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-123", "http://localhost:54545/callback"); await flow.exchangeToken("code-123#state-override", "state-123", "http://localhost:54545/callback"); @@ -107,9 +102,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-123", "http://localhost:54545/callback"); await flow.exchangeToken("code-123#", "state-explicit", "http://localhost:54545/callback"); @@ -135,9 +129,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const result = await refreshAnthropicToken("refresh-123"); + const result = await refreshAnthropicToken("refresh-123", fetchMock as unknown as typeof fetch); expect(result.access).toBe("new-access-token"); expect(result.refresh).toBe("new-refresh-token"); @@ -160,9 +153,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-123", "http://localhost:54545/callback"); const result = await flow.exchangeToken("code-123", "state-123", "http://localhost:54545/callback"); @@ -185,9 +177,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const result = await refreshAnthropicToken("refresh-123"); + const result = await refreshAnthropicToken("refresh-123", fetchMock as unknown as typeof fetch); expect(result.accountId).toBe("aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee"); expect(result.email).toBe("refreshed@example.com"); @@ -222,9 +213,8 @@ describe("anthropic oauth alignment", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-noaccount", "http://localhost:54545/callback"); const result = await flow.exchangeToken("code-noaccount", "state-noaccount", "http://localhost:54545/callback"); @@ -251,9 +241,8 @@ describe("anthropic oauth alignment", () => { headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const flow = new AnthropicOAuthFlow({}); + const flow = new AnthropicOAuthFlow({ fetch: fetchMock as unknown as typeof fetch }); await flow.generateAuthUrl("state-noaccount", "http://localhost:54545/callback"); const result = await flow.exchangeToken("code-noaccount", "state-noaccount", "http://localhost:54545/callback"); diff --git a/packages/ai/test/anthropic-stream-timeout.test.ts b/packages/ai/test/anthropic-stream-timeout.test.ts index 422a94235..94775eef8 100644 --- a/packages/ai/test/anthropic-stream-timeout.test.ts +++ b/packages/ai/test/anthropic-stream-timeout.test.ts @@ -4,8 +4,6 @@ import type { AnthropicMessagesClientLike } from "@oh-my-pi/pi-ai/providers/anth import type { Context, Model } from "@oh-my-pi/pi-ai/types"; import { waitForDelayOrAbort } from "./helpers"; -const originalFetch = global.fetch; - const model: Model<"anthropic-messages"> = { id: "claude-sonnet-4-5", name: "Claude Sonnet 4.5", @@ -167,7 +165,6 @@ async function resolveAfterMicrotasks(promise: Promise, errorMessage: stri } afterEach(() => { - global.fetch = originalFetch; vi.useRealTimers(); vi.restoreAllMocks(); }); diff --git a/packages/ai/test/azure-openai-responses-stream.test.ts b/packages/ai/test/azure-openai-responses-stream.test.ts index ca44715f9..7d463bc9f 100644 --- a/packages/ai/test/azure-openai-responses-stream.test.ts +++ b/packages/ai/test/azure-openai-responses-stream.test.ts @@ -3,9 +3,7 @@ import { type AzureOpenAIResponsesOptions, streamAzureOpenAIResponses, } from "@oh-my-pi/pi-ai/providers/azure-openai-responses"; -import type { Context, Model, Tool } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model, Tool } from "@oh-my-pi/pi-ai/types"; const azureModel: Model<"azure-openai-responses"> = { id: "gpt-5-mini", @@ -79,7 +77,6 @@ async function captureAzurePayload( } afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -175,7 +172,7 @@ describe("azure openai responses streaming", () => { }); it("surfaces nested response.failed provider errors", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createSseResponse([ { type: "response.failed", @@ -184,12 +181,12 @@ describe("azure openai responses streaming", () => { }, }, ]), - ) as unknown as typeof fetch; + ); const result = await streamAzureOpenAIResponses( azureModel, { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }] }, - { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1" }, + { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1", fetch: fetchMock }, ).result(); expect(result.stopReason).toBe("error"); @@ -197,7 +194,7 @@ describe("azure openai responses streaming", () => { }); it("surfaces response.failed incomplete reasons", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createSseResponse([ { type: "response.failed", @@ -206,12 +203,12 @@ describe("azure openai responses streaming", () => { }, }, ]), - ) as unknown as typeof fetch; + ); const result = await streamAzureOpenAIResponses( azureModel, { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }] }, - { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1" }, + { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1", fetch: fetchMock }, ).result(); expect(result.stopReason).toBe("error"); @@ -219,7 +216,7 @@ describe("azure openai responses streaming", () => { }); it("surfaces response.completed failed status_details errors", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createSseResponse([ { type: "response.completed", @@ -231,12 +228,12 @@ describe("azure openai responses streaming", () => { }, }, ]), - ) as unknown as typeof fetch; + ); const result = await streamAzureOpenAIResponses( azureModel, { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }] }, - { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1" }, + { apiKey: "test-key", azureBaseUrl: azureModel.baseUrl, azureApiVersion: "v1", fetch: fetchMock }, ).result(); expect(result.stopReason).toBe("error"); diff --git a/packages/ai/test/claude-usage-retry.test.ts b/packages/ai/test/claude-usage-retry.test.ts index 52377a412..4af8de43b 100644 --- a/packages/ai/test/claude-usage-retry.test.ts +++ b/packages/ai/test/claude-usage-retry.test.ts @@ -1,4 +1,5 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import type { UsageFetchContext } from "@oh-my-pi/pi-ai/usage"; import { claudeUsageProvider } from "@oh-my-pi/pi-ai/usage/claude"; @@ -13,7 +14,7 @@ function jsonResponse(status: number, body: unknown, headers: Record { attempt += 1; if (attempt < 3) return jsonResponse(429, { error: "rate_limited" }); return jsonResponse(200, VALID_PAYLOAD); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); expect(report).not.toBeNull(); @@ -57,7 +58,7 @@ describe("claudeUsageProvider retry contract", () => { attempt += 1; if (attempt === 1) return jsonResponse(503, { error: "unavailable" }); return jsonResponse(200, VALID_PAYLOAD); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); expect(report).not.toBeNull(); @@ -69,7 +70,7 @@ describe("claudeUsageProvider retry contract", () => { const fetchMock = (async () => { attempt += 1; return jsonResponse(401, { error: "unauthorized" }); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); expect(report).toBeNull(); @@ -81,7 +82,7 @@ describe("claudeUsageProvider retry contract", () => { const fetchMock = (async () => { attempt += 1; return jsonResponse(404, { error: "not_found" }); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); expect(report).toBeNull(); @@ -93,7 +94,7 @@ describe("claudeUsageProvider retry contract", () => { const fetchMock = (async () => { attempt += 1; return jsonResponse(429, { error: "rate_limited" }); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); expect(report).toBeNull(); @@ -110,7 +111,7 @@ describe("claudeUsageProvider retry contract", () => { return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "1" }); } return jsonResponse(200, VALID_PAYLOAD); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, retryWait)); expect(report).not.toBeNull(); @@ -131,7 +132,7 @@ describe("claudeUsageProvider retry contract", () => { return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "60" }); } return jsonResponse(200, VALID_PAYLOAD); - }) as unknown as typeof fetch; + }) as FetchImpl; const controller = new AbortController(); const retryWait = vi.fn(async (delayMs: number, signal?: AbortSignal) => { @@ -170,7 +171,7 @@ describe("claudeUsageProvider retry contract", () => { return jsonResponse(200, {}); } return jsonResponse(429, { error: "rate_limited" }); - }) as unknown as typeof fetch; + }) as FetchImpl; const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); // The 200 set lastPayload but had no usage data; 429s mean no further diff --git a/packages/ai/test/firepass.live.ts b/packages/ai/test/firepass.live.ts index 2a678ef70..c1657b440 100644 --- a/packages/ai/test/firepass.live.ts +++ b/packages/ai/test/firepass.live.ts @@ -10,7 +10,7 @@ */ import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const apiKey = process.env.FIREPASS_API_KEY; if (!apiKey) { @@ -29,14 +29,14 @@ interface CapturedRequest { body: string | null; } -const originalFetch = global.fetch; +const originalFetch = fetch; const captured: { value: CapturedRequest | null } = { value: null }; -type FetchInput = Parameters[0]; -global.fetch = (async (input: FetchInput, init?: RequestInit) => { +type FetchInput = string | URL | Request; +const fetchImpl: FetchImpl = async (input: FetchInput, init?: RequestInit) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; captured.value = { url, body: typeof init?.body === "string" ? init.body : null }; - return originalFetch(input as Parameters[0], init); -}) as typeof global.fetch; + return originalFetch(input, init); +}; const context: Context = { systemPrompt: ["Reply with exactly two words."], @@ -51,6 +51,7 @@ async function runEffort(label: string, reasoning: "xhigh" | undefined) { // Intentionally omit maxTokens here so we can also assert that the Kimi-family // safety net (openai-completions.ts isKimi) injects the catalog default. ...(reasoning ? { reasoning } : {}), + fetch: fetchImpl, }); let text = ""; let stopReason: string | undefined; diff --git a/packages/ai/test/firepass.test.ts b/packages/ai/test/firepass.test.ts index b6b8a9e72..15de3f7d1 100644 --- a/packages/ai/test/firepass.test.ts +++ b/packages/ai/test/firepass.test.ts @@ -6,16 +6,10 @@ * (`kimi-k2.6-turbo`) and the openai-completions provider translates it to the wire * form at request time. */ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function sseResponse(events: unknown[]): Response { const payload = `${events.map(e => `data: ${typeof e === "string" ? e : JSON.stringify(e)}`).join("\n\n")}\n\n`; @@ -38,14 +32,14 @@ describe("Fire Pass provider", () => { it("translates the friendly id to the router wire id when calling chat completions", async () => { const model = getBundledModel<"openai-completions">("firepass", "kimi-k2.6-turbo"); const captured: { body: string | null } = { body: null }; - global.fetch = (async (_input: unknown, init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input: string | URL | Request, init?: RequestInit) => { captured.body = typeof init?.body === "string" ? init.body : null; return sseResponse([ { choices: [{ delta: { content: "ok" }, index: 0 }] }, { choices: [{ delta: {}, finish_reason: "stop", index: 0 }] }, "[DONE]", ]); - }) as typeof global.fetch; + }; const context: Context = { systemPrompt: [], @@ -53,6 +47,7 @@ describe("Fire Pass provider", () => { }; const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey: "fpk_test", + fetch: fetchMock, }); for await (const _event of stream) { /* drain */ @@ -74,14 +69,14 @@ describe("Fire Pass provider", () => { expect(model.compat?.reasoningEffortMap?.xhigh).toBeUndefined(); const captured: { body: string | null } = { body: null }; - global.fetch = (async (_input: unknown, init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input: string | URL | Request, init?: RequestInit) => { captured.body = typeof init?.body === "string" ? init.body : null; return sseResponse([ { choices: [{ delta: { content: "ok" }, index: 0 }] }, { choices: [{ delta: {}, finish_reason: "stop", index: 0 }] }, "[DONE]", ]); - }) as typeof global.fetch; + }; const context: Context = { systemPrompt: [], @@ -90,6 +85,7 @@ describe("Fire Pass provider", () => { const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey: "fpk_test", reasoning: "xhigh", + fetch: fetchMock, }); for await (const _event of stream) { /* drain */ @@ -110,14 +106,14 @@ describe("Fire Pass provider", () => { expect(model.maxTokens).toBeGreaterThan(0); const captured: { body: string | null } = { body: null }; - global.fetch = (async (_input: unknown, init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input: string | URL | Request, init?: RequestInit) => { captured.body = typeof init?.body === "string" ? init.body : null; return sseResponse([ { choices: [{ delta: { content: "ok" }, index: 0 }] }, { choices: [{ delta: {}, finish_reason: "stop", index: 0 }] }, "[DONE]", ]); - }) as typeof global.fetch; + }; const context: Context = { systemPrompt: [], @@ -125,6 +121,7 @@ describe("Fire Pass provider", () => { }; const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey: "fpk_test", + fetch: fetchMock, // Intentionally omit maxTokens — the provider must inject the catalog default. }); for await (const _event of stream) { @@ -143,14 +140,14 @@ describe("Fire Pass provider", () => { id: "accounts/fireworks/routers/kimi-k2p6-turbo", }; const captured: { body: string | null } = { body: null }; - global.fetch = (async (_input: unknown, init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input: string | URL | Request, init?: RequestInit) => { captured.body = typeof init?.body === "string" ? init.body : null; return sseResponse([ { choices: [{ delta: { content: "ok" }, index: 0 }] }, { choices: [{ delta: {}, finish_reason: "stop", index: 0 }] }, "[DONE]", ]); - }) as typeof global.fetch; + }; const context: Context = { systemPrompt: [], @@ -158,6 +155,7 @@ describe("Fire Pass provider", () => { }; const stream = streamOpenAICompletions(model, context, { apiKey: "fpk_test", + fetch: fetchMock, }); for await (const _event of stream) { /* drain */ diff --git a/packages/ai/test/github-copilot-anthropic-auth.test.ts b/packages/ai/test/github-copilot-anthropic-auth.test.ts index b1918e458..7566c6cfb 100644 --- a/packages/ai/test/github-copilot-anthropic-auth.test.ts +++ b/packages/ai/test/github-copilot-anthropic-auth.test.ts @@ -4,10 +4,7 @@ import { OPENCODE_HEADERS } from "@oh-my-pi/pi-ai/registry/oauth/github-copilot" import type { Context, Model } from "@oh-my-pi/pi-ai/types"; import { buildAnthropicUrl } from "@oh-my-pi/pi-ai/utils/anthropic-auth"; -const originalFetch = global.fetch; - afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -94,17 +91,18 @@ describe("Anthropic Copilot auth config", () => { it("sends OpenCode Go Anthropic requests with X-Api-Key", async () => { const requestedApiKeys: Array = []; const requestedAuthorizations: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedApiKeys.push(getRequestHeader(input, init, "X-Api-Key")); requestedAuthorizations.push(getRequestHeader(input, init, "Authorization")); return new Response(JSON.stringify({ error: { type: "authentication_error", message: "Unauthorized" } }), { status: 401, headers: { "Content-Type": "application/json" }, }); - }) as unknown as typeof fetch; + }); const result = await streamAnthropic(makeOpenCodeGoQwen37Model(), testContext, { apiKey: "opencode_test_key", + fetch: fetchMock as unknown as typeof fetch, }).result(); expect(result.stopReason).toBe("error"); @@ -238,17 +236,18 @@ describe("Anthropic Copilot auth config", () => { it("forwards initiatorOverride to Copilot message requests", async () => { const requestedInitiators: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedInitiators.push(getRequestHeader(input, init, "X-Initiator")); return new Response(JSON.stringify({ error: { type: "authentication_error", message: "Unauthorized" } }), { status: 401, headers: { "Content-Type": "application/json" }, }); - }) as unknown as typeof fetch; + }); const model = makeCopilotClaudeModel(); const result = await streamAnthropic(model, testContext, { apiKey: "ghu_test_copilot_token", + fetch: fetchMock as unknown as typeof fetch, initiatorOverride: "agent", }).result(); diff --git a/packages/ai/test/github-copilot-login.test.ts b/packages/ai/test/github-copilot-login.test.ts index 5f2ca2faa..9beac62c1 100644 --- a/packages/ai/test/github-copilot-login.test.ts +++ b/packages/ai/test/github-copilot-login.test.ts @@ -1,11 +1,9 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { loginGitHubCopilot } from "@oh-my-pi/pi-ai/registry/oauth/github-copilot"; -const originalFetch = global.fetch; const FAST_POLL_OPTIONS = { pollIntervalFloorMs: 0, pollIntervalScaleMs: 1 } as const; afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -56,11 +54,11 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const onAuth = vi.fn(); const credentials = await loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth, onPrompt: mockOnPrompt(""), }); @@ -93,10 +91,10 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt("ghe.example.com"), }); @@ -125,10 +123,10 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt(" "), }); @@ -192,10 +190,10 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }); @@ -221,10 +219,10 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; await expect( loginGitHubCopilot({ + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }), @@ -248,11 +246,11 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; await expect( loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }), @@ -279,10 +277,10 @@ describe("loginGitHubCopilot", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ ...FAST_POLL_OPTIONS, + fetch: fetchMock as unknown as typeof fetch, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }); diff --git a/packages/ai/test/github-copilot-model-limits.test.ts b/packages/ai/test/github-copilot-model-limits.test.ts index 20e6248e5..8e45cbd34 100644 --- a/packages/ai/test/github-copilot-model-limits.test.ts +++ b/packages/ai/test/github-copilot-model-limits.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; @@ -7,13 +7,6 @@ import { createModelManager } from "@oh-my-pi/pi-ai/model-manager"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { githubCopilotModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); - function getHeaderValue(headers: unknown, key: string): string | undefined { if (!headers) return undefined; if (headers instanceof Headers) { @@ -45,7 +38,7 @@ async function discoverCopilotModels( expectedBaseUrl = "https://api.githubcopilot.com", expectedAuthorizationToken = apiKey, ) { - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); expect(url).toBe(`${expectedBaseUrl}/models`); expect(init?.method).toBe("GET"); @@ -55,9 +48,7 @@ async function discoverCopilotModels( headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; - - const options = githubCopilotModelManagerOptions({ apiKey }); + const options = githubCopilotModelManagerOptions({ apiKey, fetch: fetchMock }); expect(options.fetchDynamicModels).toBeDefined(); const models = await options.fetchDynamicModels?.(); expect(models).not.toBeNull(); @@ -237,7 +228,7 @@ describe("github copilot model limits mapping", () => { it("keeps discovered context window through full model resolution for bundled models", async () => { const tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-copilot-models-")); try { - global.fetch = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); expect(url).toBe("https://api.githubcopilot.com/models"); expect(init?.method).toBe("GET"); @@ -260,9 +251,9 @@ describe("github copilot model limits mapping", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }) as unknown as typeof fetch; + }); - const options = githubCopilotModelManagerOptions({ apiKey: "copilot-test-key" }); + const options = githubCopilotModelManagerOptions({ apiKey: "copilot-test-key", fetch: fetchMock }); const manager = createModelManager({ ...options, cacheDbPath: path.join(tempDir, "models.db"), diff --git a/packages/ai/test/github-copilot-openai-base-url.test.ts b/packages/ai/test/github-copilot-openai-base-url.test.ts index 930e33188..6d3aeec69 100644 --- a/packages/ai/test/github-copilot-openai-base-url.test.ts +++ b/packages/ai/test/github-copilot-openai-base-url.test.ts @@ -4,10 +4,7 @@ import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-comple import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; - afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -46,13 +43,16 @@ const enterpriseApiKey = JSON.stringify({ token: testToken, enterpriseUrl: "ghe. describe("GitHub Copilot OpenAI transport base URL", () => { it("uses model baseUrl for chat completions", async () => { const requestedUrls: string[] = []; - global.fetch = vi.fn(async (input: string | URL | Request) => { + const fetchMock = vi.fn(async (input: string | URL | Request) => { requestedUrls.push(getRequestUrl(input)); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-4o") as Model<"openai-completions">; - const result = await streamOpenAICompletions(model, testContext, { apiKey: testToken }).result(); + const result = await streamOpenAICompletions(model, testContext, { + apiKey: testToken, + fetch: fetchMock as unknown as typeof fetch, + }).result(); expect(result.stopReason).toBe("error"); expect(requestedUrls[0]).toBe("https://api.githubcopilot.com/chat/completions"); @@ -60,13 +60,16 @@ describe("GitHub Copilot OpenAI transport base URL", () => { it("uses model baseUrl for responses API", async () => { const requestedUrls: string[] = []; - global.fetch = vi.fn(async (input: string | URL | Request) => { + const fetchMock = vi.fn(async (input: string | URL | Request) => { requestedUrls.push(getRequestUrl(input)); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-5-mini") as Model<"openai-responses">; - const result = await streamOpenAIResponses(model, testContext, { apiKey: testToken }).result(); + const result = await streamOpenAIResponses(model, testContext, { + apiKey: testToken, + fetch: fetchMock as unknown as typeof fetch, + }).result(); expect(result.stopReason).toBe("error"); expect(requestedUrls[0]).toBe("https://api.githubcopilot.com/responses"); @@ -75,14 +78,17 @@ describe("GitHub Copilot OpenAI transport base URL", () => { it("routes structured enterprise credentials to the enterprise chat completions host", async () => { const requestedUrls: string[] = []; const requestedAuthHeaders: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedUrls.push(getRequestUrl(input)); requestedAuthHeaders.push(getRequestHeader(input, init, "Authorization")); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-4o") as Model<"openai-completions">; - const result = await streamOpenAICompletions(model, testContext, { apiKey: enterpriseApiKey }).result(); + const result = await streamOpenAICompletions(model, testContext, { + apiKey: enterpriseApiKey, + fetch: fetchMock as unknown as typeof fetch, + }).result(); expect(result.stopReason).toBe("error"); expect(requestedUrls[0]).toBe("https://copilot-api.ghe.example.com/chat/completions"); @@ -92,14 +98,17 @@ describe("GitHub Copilot OpenAI transport base URL", () => { it("routes structured enterprise credentials to the enterprise responses host", async () => { const requestedUrls: string[] = []; const requestedAuthHeaders: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedUrls.push(getRequestUrl(input)); requestedAuthHeaders.push(getRequestHeader(input, init, "Authorization")); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-5-mini") as Model<"openai-responses">; - const result = await streamOpenAIResponses(model, testContext, { apiKey: enterpriseApiKey }).result(); + const result = await streamOpenAIResponses(model, testContext, { + apiKey: enterpriseApiKey, + fetch: fetchMock as unknown as typeof fetch, + }).result(); expect(result.stopReason).toBe("error"); expect(requestedUrls[0]).toBe("https://copilot-api.ghe.example.com/responses"); @@ -108,14 +117,15 @@ describe("GitHub Copilot OpenAI transport base URL", () => { it("forwards initiatorOverride to chat completions requests", async () => { const requestedInitiators: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedInitiators.push(getRequestHeader(input, init, "X-Initiator")); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-4o") as Model<"openai-completions">; const result = await streamOpenAICompletions(model, testContext, { apiKey: testToken, + fetch: fetchMock as unknown as typeof fetch, initiatorOverride: "agent", }).result(); @@ -125,14 +135,15 @@ describe("GitHub Copilot OpenAI transport base URL", () => { it("forwards initiatorOverride to responses requests", async () => { const requestedInitiators: Array = []; - global.fetch = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { requestedInitiators.push(getRequestHeader(input, init, "X-Initiator")); return createUnauthorizedResponse(); - }) as unknown as typeof fetch; + }); const model = getBundledModel("github-copilot", "gpt-5-mini") as Model<"openai-responses">; const result = await streamOpenAIResponses(model, testContext, { apiKey: testToken, + fetch: fetchMock as unknown as typeof fetch, initiatorOverride: "agent", }).result(); diff --git a/packages/ai/test/google-antigravity-usage.test.ts b/packages/ai/test/google-antigravity-usage.test.ts index 799f079d7..2beb9d54a 100644 --- a/packages/ai/test/google-antigravity-usage.test.ts +++ b/packages/ai/test/google-antigravity-usage.test.ts @@ -5,6 +5,7 @@ * different model entries, and handles mixed-case tier names. */ import { describe, expect, it } from "bun:test"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import type { UsageFetchContext, UsageFetchParams } from "@oh-my-pi/pi-ai/usage"; import { antigravityUsageProvider } from "@oh-my-pi/pi-ai/usage/google-antigravity"; @@ -27,16 +28,16 @@ function makeCredential(overrides?: Partial) { } satisfies UsageFetchParams["credential"]; } -function fakeFetch(json: unknown): typeof fetch { +function fakeFetch(json: unknown): FetchImpl { const fn = async () => new Response(JSON.stringify(json), { status: 200, headers: { "content-type": "application/json" }, }); - return fn as unknown as typeof fetch; + return fn; } -function makeCtx(fetchImpl?: typeof fetch): UsageFetchContext { +function makeCtx(fetchImpl?: FetchImpl): UsageFetchContext { return { fetch: fetchImpl ?? fakeFetch({}) }; } diff --git a/packages/ai/test/google-gemini-cli-3x-thinking.test.ts b/packages/ai/test/google-gemini-cli-3x-thinking.test.ts index 117e3aba5..1b28ce5d7 100644 --- a/packages/ai/test/google-gemini-cli-3x-thinking.test.ts +++ b/packages/ai/test/google-gemini-cli-3x-thinking.test.ts @@ -1,9 +1,8 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; -import { Effort } from "@oh-my-pi/pi-ai"; +import { describe, expect, it } from "bun:test"; +import { Effort, type FetchImpl } from "@oh-my-pi/pi-ai"; import { enrichModelThinking } from "@oh-my-pi/pi-ai/model-thinking"; import { streamSimple } from "@oh-my-pi/pi-ai/stream"; import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; interface GeminiCliThinkingConfig { thinkingLevel?: string; @@ -44,19 +43,22 @@ function extractThinking(bodyText: string | undefined): GeminiCliThinkingConfig } describe("google-gemini-cli Gemini 3.x thinking mapping", () => { - afterEach(() => { - vi.restoreAllMocks(); - }); + const createFetchMock = + (capture: (body: string | undefined) => void): FetchImpl => + (_input, init) => { + capture(typeof init?.body === "string" ? init.body : undefined); + return Promise.resolve(new Response('{"error":{"message":"bad request"}}', { status: 400 })); + }; it("uses thinkingLevel for gemini-3.1-pro-preview when the effort is supported", async () => { let requestBody: string | undefined; - using _hook = hookFetch((_input, init) => { - requestBody = typeof init?.body === "string" ? init.body : undefined; - return new Response('{"error":{"message":"bad request"}}', { status: 400 }); + const fetchMock = createFetchMock(body => { + requestBody = body; }); const stream = streamSimple(createModel("gemini-3.1-pro-preview"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), reasoning: Effort.High, + fetch: fetchMock, }); await stream.result(); @@ -67,15 +69,15 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { it("rejects unsupported gemini-3.1-pro-preview efforts instead of promoting them", () => { let requestBody: string | undefined; - using _hook = hookFetch((_input, init) => { - requestBody = typeof init?.body === "string" ? init.body : undefined; - return new Response('{"error":{"message":"bad request"}}', { status: 400 }); + const fetchMock = createFetchMock(body => { + requestBody = body; }); expect(() => streamSimple(createModel("gemini-3.1-pro-preview"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), reasoning: Effort.Medium, + fetch: fetchMock, }), ).toThrow(/Supported efforts: low, high/); expect(requestBody).toBeUndefined(); @@ -83,14 +85,14 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { it("uses thinkingLevel for gemini-3.1-flash-preview", async () => { let requestBody: string | undefined; - using _hook = hookFetch((_input, init) => { - requestBody = typeof init?.body === "string" ? init.body : undefined; - return new Response('{"error":{"message":"bad request"}}', { status: 400 }); + const fetchMock = createFetchMock(body => { + requestBody = body; }); const stream = streamSimple(createModel("gemini-3.1-flash-preview"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), reasoning: Effort.Medium, + fetch: fetchMock, }); await stream.result(); @@ -101,14 +103,14 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { it("keeps thinkingBudget for gemini-2.5-pro", async () => { let requestBody: string | undefined; - using _hook = hookFetch((_input, init) => { - requestBody = typeof init?.body === "string" ? init.body : undefined; - return new Response('{"error":{"message":"bad request"}}', { status: 400 }); + const fetchMock = createFetchMock(body => { + requestBody = body; }); const stream = streamSimple(createModel("gemini-2.5-pro"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), reasoning: Effort.Medium, + fetch: fetchMock, }); await stream.result(); diff --git a/packages/ai/test/google-gemini-cli-alignment.test.ts b/packages/ai/test/google-gemini-cli-alignment.test.ts index a2367454c..38a0ca844 100644 --- a/packages/ai/test/google-gemini-cli-alignment.test.ts +++ b/packages/ai/test/google-gemini-cli-alignment.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it } from "bun:test"; import * as geminiCliProvider from "@oh-my-pi/pi-ai/providers/google-gemini-cli"; import { ANTIGRAVITY_SYSTEM_INSTRUCTION, @@ -8,8 +8,7 @@ import { streamGoogleGeminiCli, } from "@oh-my-pi/pi-ai/providers/google-gemini-cli"; import { getOAuthApiKey } from "@oh-my-pi/pi-ai/registry/oauth"; -import type { Context, Model, TJsonSchema } from "@oh-my-pi/pi-ai/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { Context, FetchImpl, Model, TJsonSchema } from "@oh-my-pi/pi-ai/types"; function createModel(provider: "google-gemini-cli" | "google-antigravity"): Model<"google-gemini-cli"> { return { @@ -227,10 +226,10 @@ describe("Google Gemini CLI alignment", () => { }); it("adds anthropic-beta for Antigravity Claude reasoning models without relying on id suffix", async () => { let requestHeaders: Headers | undefined; - using _hook = hookFetch(async (_url, init) => { + const fetchMock: FetchImpl = async (_url, init) => { requestHeaders = new Headers(init?.headers); return new Response('{"error":{"message":"bad request"}}', { status: 400 }); - }); + }; const model: Model<"google-gemini-cli"> = { ...createModel("google-antigravity"), @@ -241,6 +240,7 @@ describe("Google Gemini CLI alignment", () => { const result = await streamGoogleGeminiCli(model, createContext(), { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("error"); @@ -251,24 +251,21 @@ describe("Google Gemini CLI alignment", () => { }); describe("retry guardrails", () => { - afterEach(() => { - vi.restoreAllMocks(); - }); - it("does not treat explicit HTTP failures as network retry errors", async () => { let fetchCalls = 0; - using _hook = hookFetch(async () => { + const fetchMock: FetchImpl = async () => { fetchCalls += 1; return new Response('{"error":{"message":"busy"}}', { status: 503, headers: { "retry-after": "120" }, }); - }); + }; const model = createModel("google-gemini-cli"); const stream = streamGoogleGeminiCli(model, createContext(), { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), maxRetryDelayMs: 1000, + fetch: fetchMock, }); const result = await stream.result(); diff --git a/packages/ai/test/google-system-prompt.test.ts b/packages/ai/test/google-system-prompt.test.ts index f12d1e23b..a381d0e6c 100644 --- a/packages/ai/test/google-system-prompt.test.ts +++ b/packages/ai/test/google-system-prompt.test.ts @@ -1,7 +1,6 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { streamGoogle } from "@oh-my-pi/pi-ai/providers/google"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const model: Model<"google-generative-ai"> = { id: "gemini-3-pro-preview", @@ -20,27 +19,21 @@ async function captureGooglePayload( context: Context, ): Promise<{ config: { systemInstruction?: unknown }; contents: unknown[] }> { let captured: { config: { systemInstruction?: unknown }; contents: unknown[] } | undefined; - // Intercept the outgoing REST call so the streamGoogle promise resolves cleanly without - // hitting the network. The test only validates `onPayload` (which fires before fetch). - using _hook = hookFetch( - async () => new Response("", { status: 200, headers: { "content-type": "text/event-stream" } }), - ); + const fetchMock: FetchImpl = async () => + new Response("", { status: 200, headers: { "content-type": "text/event-stream" } }); await streamGoogle(model, context, { apiKey: "test-key", onPayload: payload => { captured = payload as { config: { systemInstruction?: unknown }; contents: unknown[] }; }, + fetch: fetchMock, }).result(); expect(captured).toBeDefined(); return captured!; } -afterEach(() => { - vi.restoreAllMocks(); -}); - describe("Google provider system prompts", () => { it("sends every system prompt block as systemInstruction text parts", async () => { const payload = await captureGooglePayload({ diff --git a/packages/ai/test/helpers/fetch-mock.ts b/packages/ai/test/helpers/fetch-mock.ts new file mode 100644 index 000000000..f0c78a3a3 --- /dev/null +++ b/packages/ai/test/helpers/fetch-mock.ts @@ -0,0 +1,8 @@ +import type { FetchImpl } from "../../src/types"; + +type FetchHandler = (input: string | URL | Request, init?: RequestInit) => Response | Promise; + +/** Wrap a fetch handler as a {@link FetchImpl}, normalizing sync `Response` returns. */ +export function mockFetch(fn: FetchHandler): FetchImpl { + return async (input, init) => fn(input, init); +} diff --git a/packages/ai/test/issue-1203-repro.test.ts b/packages/ai/test/issue-1203-repro.test.ts index 0a6b3fc52..25cbcc0f8 100644 --- a/packages/ai/test/issue-1203-repro.test.ts +++ b/packages/ai/test/issue-1203-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function createSseResponse(events: unknown[]): Response { const payload = `${events @@ -19,11 +13,11 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { +function createMockFetch(events: unknown[]): FetchImpl { async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { return createSseResponse(events); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return Object.assign(mockFetch, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -61,7 +55,7 @@ function stopChunk(model: Model<"openai-completions">): unknown { describe("issue #1203 - MiniMax Coding Plan CN think tags", () => { it("parses minimax-code-cn content into a thinking block", async () => { const model = getBundledModel("minimax-code-cn", "MiniMax-M2.5") as Model<"openai-completions">; - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ minimaxChunk(model, ""), minimaxChunk(model, "hidden reasoning"), minimaxChunk(model, ""), @@ -70,7 +64,10 @@ describe("issue #1203 - MiniMax Coding Plan CN think tags", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toEqual([ { type: "thinking", thinking: "hidden reasoning", thinkingSignature: undefined }, diff --git a/packages/ai/test/issue-1399-repro.test.ts b/packages/ai/test/issue-1399-repro.test.ts index fc523c553..fd6eb4115 100644 --- a/packages/ai/test/issue-1399-repro.test.ts +++ b/packages/ai/test/issue-1399-repro.test.ts @@ -4,8 +4,7 @@ import * as os from "node:os"; import * as path from "node:path"; import { streamBedrock } from "@oh-my-pi/pi-ai/providers/amazon-bedrock"; import { clearAwsCredentialCache } from "@oh-my-pi/pi-ai/providers/aws-credentials"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const model: Model<"bedrock-converse-stream"> = { id: "zai.glm-5", @@ -82,12 +81,12 @@ describe("issue #1399: Bedrock bearer token precedence", () => { clearAwsCredentialCache(); let requestHeaders: Headers | undefined; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { requestHeaders = new Headers(init?.headers); return new Response('{"message":"unauthorized"}', { status: 401 }); - }); + }; - const result = await streamBedrock(model, context, {}).result(); + const result = await streamBedrock(model, context, { fetch: fetchMock }).result(); expect(requestHeaders?.get("authorization")).toBe("Bearer bedrock-api-key"); expect(requestHeaders?.has("x-amz-date")).toBe(false); @@ -117,12 +116,12 @@ describe("issue #1399: Bedrock bearer token precedence", () => { clearAwsCredentialCache(); let requestHeaders: Headers | undefined; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { requestHeaders = new Headers(init?.headers); return new Response('{"message":"unauthorized"}', { status: 401 }); - }); + }; - const result = await streamBedrock(model, context, { apiKey: "" }).result(); + const result = await streamBedrock(model, context, { apiKey: "", fetch: fetchMock }).result(); expect(requestHeaders?.get("authorization")).toBe("Bearer bedrock-api-key"); expect(requestHeaders?.get("authorization")).not.toBe("Bearer "); @@ -150,12 +149,12 @@ describe("issue #1399: Bedrock bearer token precedence", () => { clearAwsCredentialCache(); let requestHeaders: Headers | undefined; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { requestHeaders = new Headers(init?.headers); return new Response('{"message":"unauthorized"}', { status: 401 }); - }); + }; - const result = await streamBedrock(model, context, { apiKey: "" }).result(); + const result = await streamBedrock(model, context, { apiKey: "", fetch: fetchMock }).result(); const authorization = requestHeaders?.get("authorization"); expect(authorization).toStartWith("AWS4-HMAC-SHA256 "); diff --git a/packages/ai/test/issue-1617-repro.test.ts b/packages/ai/test/issue-1617-repro.test.ts index 6caf43635..9efbf9c3d 100644 --- a/packages/ai/test/issue-1617-repro.test.ts +++ b/packages/ai/test/issue-1617-repro.test.ts @@ -13,23 +13,18 @@ * dynamic-fetch path (which reads the bundled models.json reference) does * not regress after a /v1/models cache refresh. */ -import { afterEach, describe, expect, test } from "bun:test"; +import { describe, expect, test } from "bun:test"; import { MODELS_DEV_PROVIDER_DESCRIPTORS, type ModelsDevModel, opencodeGoModelManagerOptions, opencodeZenModelManagerOptions, } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const OPENCODE_ZEN_BASE = "https://opencode.ai/zen/v1"; const OPENCODE_GO_BASE = "https://opencode.ai/zen/go/v1"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - describe("opencode-zen/-go resolver routes MiniMax M3 to openai-completions (issue #1617)", () => { const zenDescriptor = MODELS_DEV_PROVIDER_DESCRIPTORS.find(d => d.providerId === "opencode-zen"); const goDescriptor = MODELS_DEV_PROVIDER_DESCRIPTORS.find(d => d.providerId === "opencode-go"); @@ -64,24 +59,26 @@ describe("opencode-zen/-go resolver routes MiniMax M3 to openai-completions (iss test("opencode-zen /v1/models refresh routes a freshly-discovered M3 to openai-completions", async () => { let requestedUrl = ""; - const mockFetch = async (input: string | Request | URL): Promise => { - requestedUrl = input instanceof Request ? input.url : String(input); - return new Response( - JSON.stringify({ - data: [ - { - id: "minimax-m3-free", - name: "MiniMax M3 Free", - context_length: 200000, - }, - ], - }), - { headers: { "content-type": "application/json" } }, - ); - }; - global.fetch = Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + const mockFetch: FetchImpl = Object.assign( + async (input: string | Request | URL): Promise => { + requestedUrl = input instanceof Request ? input.url : String(input); + return new Response( + JSON.stringify({ + data: [ + { + id: "minimax-m3-free", + name: "MiniMax M3 Free", + context_length: 200000, + }, + ], + }), + { headers: { "content-type": "application/json" } }, + ); + }, + { preconnect: fetch.preconnect }, + ); - const options = opencodeZenModelManagerOptions({ apiKey: "opencode-test-key" }); + const options = opencodeZenModelManagerOptions({ apiKey: "opencode-test-key", fetch: mockFetch }); const models = await options.fetchDynamicModels?.(); const m3 = models?.find(model => model.id === "minimax-m3-free"); @@ -92,24 +89,26 @@ describe("opencode-zen/-go resolver routes MiniMax M3 to openai-completions (iss test("opencode-go /v1/models refresh routes a freshly-discovered M3 to openai-completions", async () => { let requestedUrl = ""; - const mockFetch = async (input: string | Request | URL): Promise => { - requestedUrl = input instanceof Request ? input.url : String(input); - return new Response( - JSON.stringify({ - data: [ - { - id: "minimax-m3", - name: "MiniMax M3", - context_length: 200000, - }, - ], - }), - { headers: { "content-type": "application/json" } }, - ); - }; - global.fetch = Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + const mockFetch: FetchImpl = Object.assign( + async (input: string | Request | URL): Promise => { + requestedUrl = input instanceof Request ? input.url : String(input); + return new Response( + JSON.stringify({ + data: [ + { + id: "minimax-m3", + name: "MiniMax M3", + context_length: 200000, + }, + ], + }), + { headers: { "content-type": "application/json" } }, + ); + }, + { preconnect: fetch.preconnect }, + ); - const options = opencodeGoModelManagerOptions({ apiKey: "opencode-test-key" }); + const options = opencodeGoModelManagerOptions({ apiKey: "opencode-test-key", fetch: mockFetch }); const models = await options.fetchDynamicModels?.(); const m3 = models?.find(model => model.id === "minimax-m3"); diff --git a/packages/ai/test/issue-1776-repro.test.ts b/packages/ai/test/issue-1776-repro.test.ts index f929221ff..8a9c9b5e9 100644 --- a/packages/ai/test/issue-1776-repro.test.ts +++ b/packages/ai/test/issue-1776-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function createSseResponse(events: unknown[]): Response { const payload = `${events @@ -19,11 +13,11 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { +function createMockFetch(events: unknown[]): FetchImpl { async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { return createSseResponse(events); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return Object.assign(mockFetch, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -74,13 +68,16 @@ function stopChunk(model: Model<"openai-completions">): unknown { describe("issue #1776 - MiniMax object-shaped tool arguments", () => { it("preserves object-shaped streamed tool arguments without a serialization round-trip", async () => { const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "bash", arguments: { command: "printf '%s\\n' ok" } }), stopChunk(model), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("toolUse"); expect(result.content).toEqual([ @@ -90,14 +87,17 @@ describe("issue #1776 - MiniMax object-shaped tool arguments", () => { it("still assembles tool arguments streamed as the standard JSON-string deltas", async () => { const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "bash", arguments: '{"command":' }), toolCallChunk(model, { arguments: ' "printf ok"}' }), stopChunk(model), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("toolUse"); expect(result.content).toEqual([ diff --git a/packages/ai/test/issue-1838-repro.test.ts b/packages/ai/test/issue-1838-repro.test.ts index 12e51a374..cb23c1dfc 100644 --- a/packages/ai/test/issue-1838-repro.test.ts +++ b/packages/ai/test/issue-1838-repro.test.ts @@ -32,16 +32,10 @@ * (OpenRouter, OpenCode, Kilo, Fireworks, …) translates thinking into its * own native format and would reject the extra key. */ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { AssistantMessage, Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { AssistantMessage, Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function abortedSignal(): AbortSignal { const controller = new AbortController(); @@ -56,9 +50,9 @@ function sseDoneResponse(): Response { }); } -function mockFetch(): typeof fetch { +function mockFetch(): FetchImpl { const fn = async (_input: string | URL | Request, _init?: RequestInit): Promise => sseDoneResponse(); - return Object.assign(fn, { preconnect: originalFetch.preconnect }); + return Object.assign(fn, { preconnect: fetch.preconnect }); } function moonshotKimiModel(id: string, reasoning = true): Model<"openai-completions"> { @@ -95,12 +89,12 @@ async function capturePayload( context: Context = basicContext(), ): Promise { const { promise, resolve } = Promise.withResolvers(); - global.fetch = mockFetch(); streamOpenAICompletions(model, context, { apiKey: "test-key", signal: abortedSignal(), onPayload: payload => resolve(payload), ...opts, + fetch: mockFetch(), }); return promise; } diff --git a/packages/ai/test/issue-1846-repro.test.ts b/packages/ai/test/issue-1846-repro.test.ts index 39df0d735..e81365d54 100644 --- a/packages/ai/test/issue-1846-repro.test.ts +++ b/packages/ai/test/issue-1846-repro.test.ts @@ -5,7 +5,7 @@ import { AuthStorage, SqliteAuthCredentialStore } from "@oh-my-pi/pi-ai/auth-sto import { xiaomiModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { convertMessages, detectCompat } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { getOAuthProviders } from "@oh-my-pi/pi-ai/registry/oauth"; -import type { AssistantMessage, Model, ThinkingContent, ToolCall } from "@oh-my-pi/pi-ai/types"; +import type { AssistantMessage, FetchImpl, Model, ThinkingContent, ToolCall } from "@oh-my-pi/pi-ai/types"; const TP_KEY = "tp-ci1p8t1w4e1sbxgyc8v65tnrjbzro287igmvyf25van9mt76"; const SGP_BASE_URL = "https://token-plan-sgp.xiaomimimo.com/v1"; @@ -61,14 +61,10 @@ describe("issue #1846: Xiaomi Token Plan provider support", () => { it("logs into the selected Token Plan region and stores that provider key", async () => { const seen: string[] = []; let authUrl = ""; - const fetchMock = Object.assign( - async (input: string | URL | Request) => { - seen.push(String(input)); - return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }, - { preconnect() {} }, - ); - vi.spyOn(globalThis, "fetch").mockImplementation(fetchMock); + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request) => { + seen.push(String(input)); + return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); + }); const store = new SqliteAuthCredentialStore(new Database(":memory:")); const storage = new AuthStorage(store); await storage.reload(); @@ -78,6 +74,7 @@ describe("issue #1846: Xiaomi Token Plan provider support", () => { authUrl = info.url; }, onPrompt: async () => TP_KEY, + fetch: fetchMock, }); expect(seen).toEqual([`${SGP_BASE_URL}/chat/completions`]); @@ -87,20 +84,17 @@ describe("issue #1846: Xiaomi Token Plan provider support", () => { }); it("discovers Token Plan models under the regional provider id", async () => { - const fetchMock = Object.assign( - async (_input: string | URL | Request) => { - return new Response(JSON.stringify({ data: [{ id: "mimo-v2.5-pro", name: "MiMo V2.5 Pro" }] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }, - { preconnect() {} }, - ); - vi.spyOn(globalThis, "fetch").mockImplementation(fetchMock); + const fetchMock: FetchImpl = vi.fn(async (_input: string | URL | Request) => { + return new Response(JSON.stringify({ data: [{ id: "mimo-v2.5-pro", name: "MiMo V2.5 Pro" }] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }); + }); const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY, providerId: "xiaomi-token-plan-sgp", tokenPlanRegion: "sgp", + fetch: fetchMock, }); const models = await opts.fetchDynamicModels?.(); diff --git a/packages/ai/test/issue-2080-repro.test.ts b/packages/ai/test/issue-2080-repro.test.ts index 36f9c657f..44d38c584 100644 --- a/packages/ai/test/issue-2080-repro.test.ts +++ b/packages/ai/test/issue-2080-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function createSseResponse(events: unknown[]): Response { const payload = `${events @@ -19,11 +13,11 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { +function createMockFetch(events: unknown[]): FetchImpl { async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { return createSseResponse(events); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return Object.assign(mockFetch, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -83,7 +77,7 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); // Two chunks; each carries a slice of the `input` string. The // concatenation forms the real hashline patch. - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "edit", arguments: { input: "[foo.ts#A1B2]\nreplace 91..91:\n+ " }, @@ -95,7 +89,10 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toEqual([ { @@ -115,7 +112,7 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { // that re-emit the full args on every delta. `startsWith` collapses // the merge to the latest cumulative snapshot instead of duplicating // the shared prefix. - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "edit", arguments: { input: "[foo.ts#A1B2]\nreplace 91..91:" }, @@ -127,7 +124,10 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toEqual([ { @@ -141,14 +141,17 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { it("preserves keys that only appear in earlier chunks instead of dropping them with later chunks", async () => { const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "edit", arguments: { input: "[foo.ts#A1B2]\ndelete 5" } }), toolCallChunk(model, { arguments: { dryRun: true } }), stopChunk(model), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toEqual([ { @@ -169,7 +172,7 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { // reconstructing the args the way the proxy does (concat + parse) and comparing // against the source-side merged result. const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "edit", arguments: { input: "[foo.ts#A1B2]\nreplace 91..91:\n+ " }, @@ -181,7 +184,7 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { "[DONE]", ]); - const s = streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }); + const s = streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", fetch: fetchMock }); let accumulated = ""; let toolCallEndArgs: unknown; for await (const event of s) { @@ -204,13 +207,13 @@ describe("issue #2080 - MiniMax multi-chunk object tool arguments", () => { // moves emission to `finishToolCallBlock`. The single-chunk path stays correct end-to-end: // the proxy still concatenates ("" then the final delta) and parses to the same args. const model = getBundledModel<"openai-completions">("minimax-code-cn", "MiniMax-M3"); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ toolCallChunk(model, { name: "edit", arguments: { input: "[foo.ts#A1B2]\ndelete 5" } }), stopChunk(model), "[DONE]", ]); - const s = streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }); + const s = streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", fetch: fetchMock }); let accumulated = ""; for await (const event of s) { if (event.type === "toolcall_delta") accumulated += event.delta; diff --git a/packages/ai/test/issue-2105-repro.test.ts b/packages/ai/test/issue-2105-repro.test.ts index 0e274be50..44354c056 100644 --- a/packages/ai/test/issue-2105-repro.test.ts +++ b/packages/ai/test/issue-2105-repro.test.ts @@ -15,53 +15,45 @@ describe("AIML API built-in provider (issue #2105)", () => { }); test("uses the OpenAI-compatible completions transport and AIML API base URL", async () => { - const previousFetch = global.fetch; const calls: Array<{ url: string; authorization: string | null }> = []; - global.fetch = Object.assign( - async (input: string | URL | Request, init?: RequestInit) => { - const headers = new Headers(init?.headers); - calls.push({ url: input.toString(), authorization: headers.get("authorization") }); - return new Response( - JSON.stringify({ - data: [ - { id: "alibaba/qwen-image", name: "Qwen Image" }, - { id: "claude-sonnet-4-5", name: "Claude Sonnet 4.5" }, - { id: "google/veo-3.1-first-last-image-to-video", name: "Veo Video" }, - { id: "gpt-4o", name: "GPT-4o" }, - { id: "gpt-4o-mini-tts", name: "GPT-4o Mini TTS" }, - { id: "text-embedding-3-large", name: "Text Embedding 3 Large" }, - ], - }), - { - status: 200, - headers: { "Content-Type": "application/json" }, - }, - ); - }, - { preconnect: previousFetch.preconnect }, - ); - try { - const options = aimlApiModelManagerOptions({ apiKey: "aiml-test-key" }); - const models = await options.fetchDynamicModels?.(); - - expect(options.providerId).toBe("aimlapi"); - expect(calls).toEqual([ + const fetchMock = (async (input: string | URL | Request, init?: RequestInit) => { + const headers = new Headers(init?.headers); + calls.push({ url: input.toString(), authorization: headers.get("authorization") }); + return new Response( + JSON.stringify({ + data: [ + { id: "alibaba/qwen-image", name: "Qwen Image" }, + { id: "claude-sonnet-4-5", name: "Claude Sonnet 4.5" }, + { id: "google/veo-3.1-first-last-image-to-video", name: "Veo Video" }, + { id: "gpt-4o", name: "GPT-4o" }, + { id: "gpt-4o-mini-tts", name: "GPT-4o Mini TTS" }, + { id: "text-embedding-3-large", name: "Text Embedding 3 Large" }, + ], + }), { - url: "https://api.aimlapi.com/v1/models", - authorization: "Bearer aiml-test-key", + status: 200, + headers: { "Content-Type": "application/json" }, }, - ]); - expect(models?.find(model => model.id === "gpt-4o")).toMatchObject({ - id: "gpt-4o", - name: "GPT-4o", - api: "openai-completions", - provider: "aimlapi", - baseUrl: "https://api.aimlapi.com/v1", - }); - expect(models?.map(model => model.id)).toEqual(["claude-sonnet-4-5", "gpt-4o"]); - } finally { - global.fetch = previousFetch; - } + ); + }) as typeof fetch; + const options = aimlApiModelManagerOptions({ apiKey: "aiml-test-key", fetch: fetchMock }); + const models = await options.fetchDynamicModels?.(); + + expect(options.providerId).toBe("aimlapi"); + expect(calls).toEqual([ + { + url: "https://api.aimlapi.com/v1/models", + authorization: "Bearer aiml-test-key", + }, + ]); + expect(models?.find(model => model.id === "gpt-4o")).toMatchObject({ + id: "gpt-4o", + name: "GPT-4o", + api: "openai-completions", + provider: "aimlapi", + baseUrl: "https://api.aimlapi.com/v1", + }); + expect(models?.map(model => model.id)).toEqual(["claude-sonnet-4-5", "gpt-4o"]); }); test("filters AIML API discovery to chat-compatible model IDs", () => { diff --git a/packages/ai/test/issue-2113-repro.test.ts b/packages/ai/test/issue-2113-repro.test.ts index fa3267cd7..81ed3bb2f 100644 --- a/packages/ai/test/issue-2113-repro.test.ts +++ b/packages/ai/test/issue-2113-repro.test.ts @@ -16,19 +16,13 @@ * The fix marks every `kimi-k2.x` id as reasoning + vision in the * moonshot discovery mapper and stamps default thinking metadata. */ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai/effort"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { moonshotModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import type { AssistantMessage, Context, Model } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - function moonshotKimiModel(id: string, reasoning: boolean): Model<"openai-completions"> { return { ...getBundledModel("openai", "gpt-4o-mini"), @@ -84,16 +78,15 @@ async function runHiTurn( model: Model<"openai-completions">, ): Promise<{ captured: CapturedRequest; assistant: AssistantMessage }> { const captured: CapturedRequest = { url: "", body: {} }; - const fetchImpl = async (input: string | URL | Request, init?: RequestInit): Promise => { + const fetchMock = (async (input: string | URL | Request, init?: RequestInit): Promise => { const url = typeof input === "string" ? input : input instanceof URL ? input.href : input.url; captured.url = url; const raw = typeof init?.body === "string" ? init.body : ""; captured.body = raw ? (JSON.parse(raw) as Record) : {}; return buildMockMoonshotResponse(); - }; - global.fetch = Object.assign(fetchImpl, { preconnect: originalFetch.preconnect }); + }) as typeof fetch; - const stream = streamOpenAICompletions(model, basicContext(), { apiKey: "test-key" }); + const stream = streamOpenAICompletions(model, basicContext(), { apiKey: "test-key", fetch: fetchMock }); for await (const _ of stream) { // drain until terminal event } @@ -103,8 +96,7 @@ async function runHiTurn( describe("issue #2113 — moonshot kimi-k2.6 discovery and wire format", () => { it("moonshot discovery mapper marks kimi-k2.6 as reasoning + vision with thinking metadata", async () => { - const opts = moonshotModelManagerOptions({ apiKey: "test-key" }); - const fetchImpl = async (input: string | URL | Request): Promise => { + const fetchMock = (async (input: string | URL | Request): Promise => { const url = typeof input === "string" ? input : input instanceof URL ? input.href : input.url; expect(url).toContain("api.moonshot.ai/v1/models"); const body = { @@ -119,10 +111,9 @@ describe("issue #2113 — moonshot kimi-k2.6 discovery and wire format", () => { status: 200, headers: { "content-type": "application/json" }, }); - }; - global.fetch = Object.assign(fetchImpl, { preconnect: originalFetch.preconnect }); + }) as typeof fetch; - const models = await opts.fetchDynamicModels?.(); + const models = await moonshotModelManagerOptions({ apiKey: "test-key", fetch: fetchMock }).fetchDynamicModels?.(); expect(models).toBeDefined(); const byId = new Map(models?.map(m => [m.id, m])); @@ -158,18 +149,18 @@ describe("issue #2113 — moonshot kimi-k2.6 discovery and wire format", () => { it("wire body includes thinking.keep='all' when reasoning is explicitly requested", async () => { const model = moonshotKimiModel("kimi-k2.6", true); const captured: CapturedRequest = { url: "", body: {} }; - const fetchImpl = async (input: string | URL | Request, init?: RequestInit): Promise => { + const fetchMock = (async (input: string | URL | Request, init?: RequestInit): Promise => { const url = typeof input === "string" ? input : input instanceof URL ? input.href : input.url; captured.url = url; const raw = typeof init?.body === "string" ? init.body : ""; captured.body = raw ? (JSON.parse(raw) as Record) : {}; return buildMockMoonshotResponse(); - }; - global.fetch = Object.assign(fetchImpl, { preconnect: originalFetch.preconnect }); + }) as typeof fetch; const stream = streamOpenAICompletions(model, basicContext(), { apiKey: "test-key", reasoning: "high", + fetch: fetchMock, }); for await (const _ of stream) { // drain diff --git a/packages/ai/test/issue-772-repro.test.ts b/packages/ai/test/issue-772-repro.test.ts index 76594649d..6d3bece66 100644 --- a/packages/ai/test/issue-772-repro.test.ts +++ b/packages/ai/test/issue-772-repro.test.ts @@ -1,7 +1,7 @@ import { describe, expect, it } from "bun:test"; import { xiaomiModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { loginXiaomi } from "@oh-my-pi/pi-ai/registry/oauth/xiaomi"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const TOKEN_PLAN_SGP_HOST = "token-plan-sgp.xiaomimimo.com"; const STANDARD_HOST = "api.xiaomimimo.com"; @@ -9,15 +9,16 @@ const STANDARD_HOST = "api.xiaomimimo.com"; describe("issue-772: Xiaomi MiMo token-plan (tp-) keys", () => { it("loginXiaomi validates tp- keys against the SGP token-plan host first", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response("{}", { status: 200 }); - }); + }; await loginXiaomi({ onAuth: () => {}, onPrompt: async () => "tp-test-key", onProgress: () => {}, + fetch: fetchMock, }); expect(seen).toHaveLength(1); @@ -28,15 +29,15 @@ describe("issue-772: Xiaomi MiMo token-plan (tp-) keys", () => { it("xiaomiModelManagerOptions discovers models from the SGP token-plan host when given a tp- key", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response(JSON.stringify({ data: [] }), { status: 200, headers: { "content-type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: "tp-test-key" }); + const opts = xiaomiModelManagerOptions({ apiKey: "tp-test-key", fetch: fetchMock }); await opts.fetchDynamicModels?.(); expect(seen.length).toBeGreaterThan(0); @@ -47,15 +48,15 @@ describe("issue-772: Xiaomi MiMo token-plan (tp-) keys", () => { it("xiaomiModelManagerOptions still uses the standard host for sk- keys", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response(JSON.stringify({ data: [] }), { status: 200, headers: { "content-type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: "sk-test-key" }); + const opts = xiaomiModelManagerOptions({ apiKey: "sk-test-key", fetch: fetchMock }); await opts.fetchDynamicModels?.(); expect(seen.length).toBeGreaterThan(0); diff --git a/packages/ai/test/issue-827-repro.test.ts b/packages/ai/test/issue-827-repro.test.ts index 70094e71e..56a19c73c 100644 --- a/packages/ai/test/issue-827-repro.test.ts +++ b/packages/ai/test/issue-827-repro.test.ts @@ -7,18 +7,12 @@ * — when a forced tool_choice is sent to a Kimi reasoning model, we strip * reasoning for that single turn rather than dropping `tool_choice` outright. */ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import type { Context, Model, Tool } from "@oh-my-pi/pi-ai/types"; import * as z from "zod/v4"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - const echoTool: Tool = { name: "echo", description: "Echo input", diff --git a/packages/ai/test/issue-847-repro.test.ts b/packages/ai/test/issue-847-repro.test.ts index fccf56034..c6137bdf5 100644 --- a/packages/ai/test/issue-847-repro.test.ts +++ b/packages/ai/test/issue-847-repro.test.ts @@ -1,17 +1,15 @@ import { afterEach, describe, expect, test, vi } from "bun:test"; import { ollamaModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; - -const originalFetch = global.fetch; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); describe("ollama provider context window discovery (issue #847)", () => { test("uses /api/show context_length for unbundled cloud models like deepseek-v4-flash:cloud", async () => { const showCalls: string[] = []; - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/v1/models") { return new Response( @@ -52,7 +50,7 @@ describe("ollama provider context window discovery (issue #847)", () => { throw new Error(`Unexpected URL: ${url}`); }) as unknown as typeof fetch; - const options = ollamaModelManagerOptions(); + const options = ollamaModelManagerOptions({ fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const deepseek = models?.find(m => m.id === "deepseek-v4-flash:cloud"); @@ -64,7 +62,7 @@ describe("ollama provider context window discovery (issue #847)", () => { test("caches /api/show results across repeated fetchDynamicModels calls", async () => { const showCalls: string[] = []; - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/v1/models") { return new Response( @@ -86,14 +84,14 @@ describe("ollama provider context window discovery (issue #847)", () => { throw new Error(`Unexpected URL: ${url}`); }) as unknown as typeof fetch; - const options = ollamaModelManagerOptions(); + const options = ollamaModelManagerOptions({ fetch: fetchMock }); await options.fetchDynamicModels?.(); await options.fetchDynamicModels?.(); expect(showCalls).toEqual(["deepseek-v4-flash:cloud"]); }); test("falls back to 128k when /api/show is unavailable", async () => { - global.fetch = vi.fn(async input => { + const fetchMock: FetchImpl = vi.fn(async input => { const url = String(input); if (url === "http://127.0.0.1:11434/v1/models") { return new Response(JSON.stringify({ object: "list", data: [{ id: "mystery:1b", object: "model" }] }), { @@ -107,7 +105,7 @@ describe("ollama provider context window discovery (issue #847)", () => { throw new Error(`Unexpected URL: ${url}`); }) as unknown as typeof fetch; - const options = ollamaModelManagerOptions(); + const options = ollamaModelManagerOptions({ fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const mystery = models?.find(m => m.id === "mystery:1b"); expect(mystery?.contextWindow).toBe(128000); diff --git a/packages/ai/test/issue-887-repro.test.ts b/packages/ai/test/issue-887-repro.test.ts index 463a3e56d..1f5b92b45 100644 --- a/packages/ai/test/issue-887-repro.test.ts +++ b/packages/ai/test/issue-887-repro.test.ts @@ -8,7 +8,7 @@ * descriptor must override these specific ids to openai-completions so that * regenerated models.json keeps the correct routing. */ -import { afterEach, describe, expect, test } from "bun:test"; +import { describe, expect, test } from "bun:test"; import { MODELS_DEV_PROVIDER_DESCRIPTORS, type ModelsDevModel, @@ -17,12 +17,6 @@ import { const OPENCODE_GO_BASE = "https://opencode.ai/zen/go/v1"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - describe("opencode-go resolver routes 404-ing ids to openai-completions (issue #887)", () => { const descriptor = MODELS_DEV_PROVIDER_DESCRIPTORS.find(d => d.providerId === "opencode-go"); @@ -51,7 +45,7 @@ describe("opencode-go resolver routes 404-ing ids to openai-completions (issue # test("runtime /v1/models refresh preserves qwen3.7-max Anthropic transport", async () => { let requestedUrl = ""; - const mockFetch = async (input: string | Request | URL): Promise => { + const fetchMock = (async (input: string | Request | URL): Promise => { requestedUrl = input instanceof Request ? input.url : String(input); return new Response( JSON.stringify({ @@ -59,10 +53,9 @@ describe("opencode-go resolver routes 404-ing ids to openai-completions (issue # }), { headers: { "content-type": "application/json" } }, ); - }; - global.fetch = Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + }) as typeof fetch; - const options = opencodeGoModelManagerOptions({ apiKey: "opencode-test-key" }); + const options = opencodeGoModelManagerOptions({ apiKey: "opencode-test-key", fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const qwenMax = models?.find(model => model.id === "qwen3.7-max"); diff --git a/packages/ai/test/issue-911-repro.test.ts b/packages/ai/test/issue-911-repro.test.ts index f039c8e58..8f23c7cd0 100644 --- a/packages/ai/test/issue-911-repro.test.ts +++ b/packages/ai/test/issue-911-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function createSseResponse(events: unknown[]): Response { const payload = `${events @@ -19,11 +13,10 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { - async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { +function createMockFetch(events: unknown[]): FetchImpl { + return (async (_input: string | URL | Request, _init?: RequestInit): Promise => { return createSseResponse(events); - } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + }) as typeof fetch; } function baseContext(): Context { @@ -51,7 +44,7 @@ describe("issue #911 - Mistral Medium 3.5 array content parts", () => { }; it("normalizes array-of-parts delta.content into the assembled text without [object Object]", async () => { - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-mistral-1", object: "chat.completion.chunk", @@ -86,7 +79,10 @@ describe("issue #911 - Mistral Medium 3.5 array content parts", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) @@ -97,7 +93,7 @@ describe("issue #911 - Mistral Medium 3.5 array content parts", () => { }); it("handles mixed string and array-of-parts content shapes within one stream", async () => { - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-mistral-2", object: "chat.completion.chunk", @@ -132,7 +128,10 @@ describe("issue #911 - Mistral Medium 3.5 array content parts", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) diff --git a/packages/ai/test/issue-912-repro.test.ts b/packages/ai/test/issue-912-repro.test.ts index e7a5aa940..60507b796 100644 --- a/packages/ai/test/issue-912-repro.test.ts +++ b/packages/ai/test/issue-912-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { stream } from "@oh-my-pi/pi-ai/stream"; import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - function makeCopilotResponsesModel(baseUrl: string): Model<"openai-responses"> { return { id: "gpt-5-mini", @@ -58,7 +52,7 @@ describe("issue #912 — github-copilot abort propagation", () => { // signal — this is the regression vector. Real Bun fetch normally // propagates abort to the body reader, but we cannot rely on it for // every transport (HTTP/2, intermediaries, native sockets). - global.fetch = (async (input: string | URL | Request, _init?: RequestInit): Promise => { + const fetchMock = (async (input: string | URL | Request, _init?: RequestInit): Promise => { const url = typeof input === "string" ? input : input instanceof URL ? input.href : input.url; if (!url.endsWith("/responses")) { throw new Error(`Unexpected fetch to ${url}`); @@ -88,6 +82,7 @@ describe("issue #912 — github-copilot abort propagation", () => { const providerStream = stream(model, makeContext(), { apiKey: JSON.stringify({ token: "ghu_test_token", enterpriseUrl: undefined }), signal: controller.signal, + fetch: fetchMock, }); await fetchInvoked.promise; diff --git a/packages/ai/test/issue-945-repro.test.ts b/packages/ai/test/issue-945-repro.test.ts index d1221c2a1..2d24f2065 100644 --- a/packages/ai/test/issue-945-repro.test.ts +++ b/packages/ai/test/issue-945-repro.test.ts @@ -1,15 +1,9 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import type { Context, Model, Tool } from "@oh-my-pi/pi-ai/types"; import * as z from "zod/v4"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - const echoTool: Tool = { name: "echo", description: "Echo input", diff --git a/packages/ai/test/issue-955-repro.test.ts b/packages/ai/test/issue-955-repro.test.ts index cd59363b3..b8df44778 100644 --- a/packages/ai/test/issue-955-repro.test.ts +++ b/packages/ai/test/issue-955-repro.test.ts @@ -1,14 +1,8 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import type { Context, Model } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - const context: Context = { systemPrompt: ["stable instructions", "cacheable policy"], messages: [{ role: "user", content: "hello", timestamp: Date.now() }], diff --git a/packages/ai/test/issue-957-repro.test.ts b/packages/ai/test/issue-957-repro.test.ts index d26e1f9ea..0ce011327 100644 --- a/packages/ai/test/issue-957-repro.test.ts +++ b/packages/ai/test/issue-957-repro.test.ts @@ -5,10 +5,7 @@ import * as path from "node:path"; import { AuthStorage, SqliteAuthCredentialStore } from "@oh-my-pi/pi-ai/auth-storage"; import * as kimiOauth from "@oh-my-pi/pi-ai/registry/oauth/kimi"; -const originalFetch = global.fetch; - afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -31,20 +28,24 @@ describe("issue #957 - Kimi OAuth refresh", () => { const issuedAt = 1_700_000_000_000; vi.spyOn(Date, "now").mockReturnValue(issuedAt); vi.spyOn(kimiOauth, "getKimiCommonHeaders").mockReturnValue(kimiHeadersStub); - - global.fetch = (async (_input: string | URL, init?: RequestInit) => { - const params = new URLSearchParams(String(init?.body)); - expect(params.get("grant_type")).toBe("refresh_token"); - expect(params.get("refresh_token")).toBe("refresh-0"); - return new Response( - JSON.stringify({ - access_token: "access-1", - refresh_token: "refresh-1", - expires_in: 60 * 60, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - }) as unknown as typeof fetch; + vi.spyOn(globalThis, "fetch").mockImplementation( + Object.assign( + async (_input: string | URL | Request, init?: RequestInit) => { + const params = new URLSearchParams(String(init?.body)); + expect(params.get("grant_type")).toBe("refresh_token"); + expect(params.get("refresh_token")).toBe("refresh-0"); + return new Response( + JSON.stringify({ + access_token: "access-1", + refresh_token: "refresh-1", + expires_in: 60 * 60, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + }, + { preconnect: fetch.preconnect }, + ), + ); const refreshed = await kimiOauth.refreshKimiToken("refresh-0"); @@ -62,20 +63,25 @@ describe("issue #957 - Kimi OAuth refresh", () => { vi.spyOn(kimiOauth, "getKimiCommonHeaders").mockReturnValue(kimiHeadersStub); let refreshCalls = 0; - global.fetch = (async (_input: string | URL, init?: RequestInit) => { - refreshCalls += 1; - const params = new URLSearchParams(String(init?.body)); - expect(params.get("grant_type")).toBe("refresh_token"); - expect(params.get("refresh_token")).toBe("refresh-stored"); - return new Response( - JSON.stringify({ - access_token: "access-refreshed", - refresh_token: "refresh-refreshed", - expires_in: 60 * 60, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - }) as unknown as typeof fetch; + vi.spyOn(globalThis, "fetch").mockImplementation( + Object.assign( + async (_input: string | URL | Request, init?: RequestInit) => { + refreshCalls += 1; + const params = new URLSearchParams(String(init?.body)); + expect(params.get("grant_type")).toBe("refresh_token"); + expect(params.get("refresh_token")).toBe("refresh-stored"); + return new Response( + JSON.stringify({ + access_token: "access-refreshed", + refresh_token: "refresh-refreshed", + expires_in: 60 * 60, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + }, + { preconnect: fetch.preconnect }, + ), + ); const tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-issue-957-")); const store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); diff --git a/packages/ai/test/issue-959-repro.test.ts b/packages/ai/test/issue-959-repro.test.ts index 87fe2219c..f3146ec43 100644 --- a/packages/ai/test/issue-959-repro.test.ts +++ b/packages/ai/test/issue-959-repro.test.ts @@ -1,13 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function createSseResponse(events: unknown[]): Response { const payload = `${events @@ -19,11 +13,11 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { +function createMockFetch(events: unknown[]): FetchImpl { async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { return createSseResponse(events); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return Object.assign(mockFetch, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -44,7 +38,7 @@ describe("issue #959 - deepseek chat-template token leakage", () => { }; it("strips leaked deepseek chat-template markers from visible text for deepseek providers", async () => { - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-deepseek-1", object: "chat.completion.chunk", @@ -76,7 +70,10 @@ describe("issue #959 - deepseek chat-template token leakage", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(block => block.type === "text") .map(block => (block as { text: string }).text) @@ -87,7 +84,7 @@ describe("issue #959 - deepseek chat-template token leakage", () => { }); it("holds partial deepseek markers across chunks before stripping them", async () => { - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-deepseek-2", object: "chat.completion.chunk", @@ -112,7 +109,10 @@ describe("issue #959 - deepseek chat-template token leakage", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(block => block.type === "text") .map(block => (block as { text: string }).text) diff --git a/packages/ai/test/issue-969-repro.test.ts b/packages/ai/test/issue-969-repro.test.ts index c33749f92..e89145202 100644 --- a/packages/ai/test/issue-969-repro.test.ts +++ b/packages/ai/test/issue-969-repro.test.ts @@ -1,14 +1,8 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai/effort"; import { getSupportedEfforts } from "@oh-my-pi/pi-ai/model-thinking"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const testContext: Context = { messages: [{ role: "user", content: "hello", timestamp: 0 }], @@ -46,7 +40,7 @@ describe("issue #969 — custom thinking metadata must preserve explicit xhigh", it("uses the configured xhigh effort for custom OpenAI-compatible models", async () => { const model = customOpenAICompatModel(); let payload: Record | undefined; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, init?: RequestInit): Promise => { payload = JSON.parse(typeof init?.body === "string" ? init.body : "{}") as Record; return createSseResponse([ @@ -67,13 +61,14 @@ describe("issue #969 — custom thinking metadata must preserve explicit xhigh", "[DONE]", ]); }, - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); expect(getSupportedEfforts(model)).toContain(Effort.XHigh); const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key", reasoning: "xhigh", + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("stop"); diff --git a/packages/ai/test/kilo-login.test.ts b/packages/ai/test/kilo-login.test.ts index 23e22d91d..372eaa6b4 100644 --- a/packages/ai/test/kilo-login.test.ts +++ b/packages/ai/test/kilo-login.test.ts @@ -1,16 +1,10 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { loginKilo } from "@oh-my-pi/pi-ai/registry/kilo"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("kilo oauth login", () => { it("returns OAuth credentials when device authorization is approved", async () => { - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); if (url === "https://api.kilo.ai/api/device-auth/codes") { expect(init?.method).toBe("POST"); @@ -31,10 +25,9 @@ describe("kilo oauth login", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; const onAuth = vi.fn(); - const credentials = await loginKilo({ onAuth }); + const credentials = await loginKilo({ onAuth, fetch: fetchMock }); expect(onAuth).toHaveBeenCalledWith({ url: "https://kilo.ai/verify", @@ -47,13 +40,15 @@ describe("kilo oauth login", () => { }); it("surfaces rate-limit errors from device authorization start", async () => { - global.fetch = vi.fn(async () => new Response(null, { status: 429 })) as unknown as typeof fetch; + const fetchMock: FetchImpl = vi.fn(async () => new Response(null, { status: 429 })); - await expect(loginKilo({})).rejects.toThrow("Too many pending authorization requests. Please try again later."); + await expect(loginKilo({ fetch: fetchMock })).rejects.toThrow( + "Too many pending authorization requests. Please try again later.", + ); }); it("surfaces denied device authorization state", async () => { - const fetchMock = vi.fn(async (input: string | URL) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request) => { const url = typeof input === "string" ? input : input.toString(); if (url === "https://api.kilo.ai/api/device-auth/codes") { return new Response( @@ -70,8 +65,7 @@ describe("kilo oauth login", () => { } throw new Error(`Unexpected URL: ${url}`); }); - global.fetch = fetchMock as unknown as typeof fetch; - await expect(loginKilo({})).rejects.toThrow("Authorization was denied"); + await expect(loginKilo({ fetch: fetchMock })).rejects.toThrow("Authorization was denied"); }); }); diff --git a/packages/ai/test/minimax-code-login.test.ts b/packages/ai/test/minimax-code-login.test.ts index 754783112..efcf9b0b0 100644 --- a/packages/ai/test/minimax-code-login.test.ts +++ b/packages/ai/test/minimax-code-login.test.ts @@ -1,20 +1,21 @@ import { describe, expect, it } from "bun:test"; import { loginMiniMaxCode, loginMiniMaxCodeCn } from "@oh-my-pi/pi-ai/registry/oauth/minimax-code"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("MiniMax Coding Plan login", () => { it("opens the international platform and validates against the international API", async () => { const authUrls: string[] = []; const validationUrls: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { validationUrls.push(String(input)); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; const apiKey = await loginMiniMaxCode({ onAuth: info => authUrls.push(info.url), onPrompt: async () => " sk-intl ", + fetch: fetchMock, }); expect(apiKey).toBe("sk-intl"); @@ -26,14 +27,15 @@ describe("MiniMax Coding Plan login", () => { const authUrls: string[] = []; const validationUrls: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { validationUrls.push(String(input)); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; const apiKey = await loginMiniMaxCodeCn({ onAuth: info => authUrls.push(info.url), onPrompt: async () => " sk-cn ", + fetch: fetchMock, }); expect(apiKey).toBe("sk-cn"); diff --git a/packages/ai/test/nanogpt-login.test.ts b/packages/ai/test/nanogpt-login.test.ts index 0b301b18d..7a0635da2 100644 --- a/packages/ai/test/nanogpt-login.test.ts +++ b/packages/ai/test/nanogpt-login.test.ts @@ -1,16 +1,10 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { loginNanoGPT } from "@oh-my-pi/pi-ai/registry/nanogpt"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("nanogpt login", () => { it("validates API key without requiring a specific model entitlement", async () => { - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); expect(url).toBe("https://nano-gpt.com/api/v1/models"); expect(init?.method).toBe("GET"); @@ -20,10 +14,10 @@ describe("nanogpt login", () => { headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; const apiKey = await loginNanoGPT({ onPrompt: async () => "sk-nano-test", + fetch: fetchMock, }); expect(apiKey).toBe("sk-nano-test"); @@ -31,17 +25,17 @@ describe("nanogpt login", () => { }); it("surfaces validation errors from models endpoint", async () => { - const fetchMock = vi.fn(async () => { + const fetchMock: FetchImpl = vi.fn(async () => { return new Response('{"code":"invalid_api_key"}', { status: 401, headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; await expect( loginNanoGPT({ onPrompt: async () => "sk-nano-test", + fetch: fetchMock, }), ).rejects.toThrow("NanoGPT API key validation failed (401)"); }); diff --git a/packages/ai/test/nanogpt-model-limits.test.ts b/packages/ai/test/nanogpt-model-limits.test.ts index 512a74970..d8266d284 100644 --- a/packages/ai/test/nanogpt-model-limits.test.ts +++ b/packages/ai/test/nanogpt-model-limits.test.ts @@ -1,20 +1,13 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai/effort"; import { nanoGptModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); - async function discoverNanoGptModels( payload: unknown, apiKey = "nanogpt-test-key", expectedBaseUrl = "https://nano-gpt.com/api/v1", ) { - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); expect(url).toBe(`${expectedBaseUrl}/models`); expect(init?.method).toBe("GET"); @@ -27,9 +20,7 @@ async function discoverNanoGptModels( headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; - - const options = nanoGptModelManagerOptions({ apiKey, baseUrl: expectedBaseUrl }); + const options = nanoGptModelManagerOptions({ apiKey, baseUrl: expectedBaseUrl, fetch: fetchMock }); expect(options.fetchDynamicModels).toBeDefined(); const models = await options.fetchDynamicModels?.(); expect(models).not.toBeNull(); diff --git a/packages/ai/test/oauth-deepseek.test.ts b/packages/ai/test/oauth-deepseek.test.ts index 3ecf301be..9d1efd233 100644 --- a/packages/ai/test/oauth-deepseek.test.ts +++ b/packages/ai/test/oauth-deepseek.test.ts @@ -1,26 +1,20 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; - +import { describe, expect, it, vi } from "bun:test"; import { loginDeepSeek, normalizeDeepSeekApiKey } from "@oh-my-pi/pi-ai/registry/deepseek"; import type { OAuthController } from "@oh-my-pi/pi-ai/registry/oauth/types"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); - -function makeController(paste: string): OAuthController { +function makeController(paste: string, fetchMock: FetchImpl): OAuthController { return { onAuth: () => {}, onPrompt: async () => paste, + fetch: fetchMock, }; } describe("loginDeepSeek validation", () => { it("validates against GET /v1/models and returns the trimmed key on 200", async () => { const calls: Array<{ url: string; method?: string; auth?: string | null }> = []; - const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; const headers = new Headers(init?.headers ?? {}); calls.push({ url, method: init?.method, auth: headers.get("authorization") }); @@ -35,9 +29,8 @@ describe("loginDeepSeek validation", () => { { status: 200, headers: { "Content-Type": "application/json" } }, ); }); - global.fetch = fetchMock as unknown as typeof fetch; - const key = await loginDeepSeek(makeController(" sk-valid-key ")); + const key = await loginDeepSeek(makeController(" sk-valid-key ", fetchMock)); expect(key).toBe("sk-valid-key"); expect(fetchMock).toHaveBeenCalledTimes(1); @@ -47,34 +40,31 @@ describe("loginDeepSeek validation", () => { }); it("throws a validation error when /v1/models returns 401", async () => { - const fetchMock = vi.fn( + const fetchMock: FetchImpl = vi.fn( async () => new Response("invalid api key", { status: 401, headers: { "Content-Type": "text/plain" } }), ); - global.fetch = fetchMock as unknown as typeof fetch; - await expect(loginDeepSeek(makeController("sk-bad-key"))).rejects.toThrow(/deepseek.*401/i); + await expect(loginDeepSeek(makeController("sk-bad-key", fetchMock))).rejects.toThrow(/deepseek.*401/i); expect(fetchMock).toHaveBeenCalledTimes(1); }); it("strips a pasted 'Bearer ' prefix before validating", async () => { const seen: string[] = []; - const fetchMock = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { seen.push(new Headers(init?.headers ?? {}).get("authorization") ?? ""); return new Response(JSON.stringify({ object: "list", data: [] }), { status: 200 }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const key = await loginDeepSeek(makeController("Bearer sk-with-prefix")); + const key = await loginDeepSeek(makeController("Bearer sk-with-prefix", fetchMock)); expect(key).toBe("sk-with-prefix"); expect(seen[0]).toBe("Bearer sk-with-prefix"); // exactly one Bearer, not nested }); it("rejects an empty paste before touching the network", async () => { - const fetchMock = vi.fn(async () => new Response("", { status: 200 })); - global.fetch = fetchMock as unknown as typeof fetch; + const fetchMock: FetchImpl = vi.fn(async () => new Response("", { status: 200 })); - await expect(loginDeepSeek(makeController(" "))).rejects.toThrow(/API key is required/i); + await expect(loginDeepSeek(makeController(" ", fetchMock))).rejects.toThrow(/API key is required/i); expect(fetchMock).not.toHaveBeenCalled(); }); }); diff --git a/packages/ai/test/ollama-cloud-provider.test.ts b/packages/ai/test/ollama-cloud-provider.test.ts index 033b7b8de..69be295f9 100644 --- a/packages/ai/test/ollama-cloud-provider.test.ts +++ b/packages/ai/test/ollama-cloud-provider.test.ts @@ -1,10 +1,9 @@ import { afterEach, describe, expect, test, vi } from "bun:test"; import { ollamaCloudModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/ollama"; import { completeSimple, getEnvApiKey, stream, streamSimple } from "@oh-my-pi/pi-ai/stream"; -import type { Context, Model, Tool } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, Tool } from "@oh-my-pi/pi-ai/types"; const originalApiKey = Bun.env.OLLAMA_CLOUD_API_KEY; -const originalFetch = global.fetch; const cloudModel: Model<"ollama-chat"> = { id: "gpt-oss:120b", @@ -52,7 +51,6 @@ afterEach(() => { } else { Bun.env.OLLAMA_CLOUD_API_KEY = originalApiKey; } - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -63,7 +61,7 @@ describe("ollama-cloud provider support", () => { }); test("discovers ollama-cloud models from native cloud endpoints", async () => { - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); const headers = new Headers(init?.headers); expect(headers.get("Authorization")).toBe("Bearer cloud-test-key"); @@ -92,9 +90,9 @@ describe("ollama-cloud provider support", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key" }); + const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key", fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const gpt = models?.find(model => model.id === "gpt-oss:120b"); const qwen = models?.find(model => model.id === "qwen3:32b"); @@ -108,14 +106,11 @@ describe("ollama-cloud provider support", () => { expect(gpt?.input).toEqual(["text", "image"]); expect(qwen?.name).toBe("Qwen 3 32B"); expect(qwen?.input).toEqual(["text", "image"]); - expect(global.fetch).toHaveBeenCalledWith( - "https://ollama.com/api/tags", - expect.objectContaining({ method: "GET" }), - ); + expect(fetchMock).toHaveBeenCalledWith("https://ollama.com/api/tags", expect.objectContaining({ method: "GET" })); }); test("tolerates individual /api/show failures during model discovery", async () => { - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); if (url === "https://ollama.com/api/tags") { return new Response( @@ -136,9 +131,9 @@ describe("ollama-cloud provider support", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key" }); + const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key", fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const ids = models?.map(m => m.id).sort(); @@ -148,7 +143,7 @@ describe("ollama-cloud provider support", () => { }); test("falls back to bundled metadata when /api/show metadata is unavailable", async () => { - global.fetch = vi.fn(async (input, _init) => { + const fetchMock: FetchImpl = vi.fn(async (input, _init) => { const url = String(input); if (url === "https://ollama.com/api/tags") { return new Response( @@ -162,9 +157,9 @@ describe("ollama-cloud provider support", () => { return new Response(null, { status: 500 }); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key" }); + const options = ollamaCloudModelManagerOptions({ apiKey: "cloud-test-key", fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const model = models?.find(candidate => candidate.id === "gpt-oss:120b"); @@ -177,7 +172,7 @@ describe("ollama-cloud provider support", () => { }); test("streams native chat responses with thinking, text, and usage mapping", async () => { - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { expect(String(input)).toBe("https://ollama.com/api/chat"); const headers = new Headers(init?.headers); expect(headers.get("Authorization")).toBe("Bearer cloud-test-key"); @@ -205,14 +200,14 @@ describe("ollama-cloud provider support", () => { eval_count: 4, }, ]); - }) as unknown as typeof fetch; + }); const response = stream( cloudModel, { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }, - { apiKey: "cloud-test-key" }, + { apiKey: "cloud-test-key", fetch: fetchMock }, ); const eventTypes: string[] = []; @@ -235,7 +230,7 @@ describe("ollama-cloud provider support", () => { }); test("supports ollama-cloud through streamSimple option mapping", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createNdjsonResponse([ { model: "gpt-oss:120b", @@ -244,12 +239,12 @@ describe("ollama-cloud provider support", () => { }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 2, eval_count: 4 }, ]), - ) as unknown as typeof fetch; + ); const response = await streamSimple( cloudModel, { messages: [{ role: "user", content: "Say hi", timestamp: Date.now() }] }, - { apiKey: "cloud-test-key", toolChoice: "auto" }, + { apiKey: "cloud-test-key", toolChoice: "auto", fetch: fetchMock }, ).result(); expect(response.stopReason).toBe("stop"); @@ -259,7 +254,7 @@ describe("ollama-cloud provider support", () => { }); test("supports ollama-cloud through completeSimple top-level contract", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createNdjsonResponse([ { model: "gpt-oss:120b", @@ -268,12 +263,12 @@ describe("ollama-cloud provider support", () => { }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 3, eval_count: 5 }, ]), - ) as unknown as typeof fetch; + ); const response = await completeSimple( cloudModel, { messages: [{ role: "user", content: "Finish this", timestamp: Date.now() }] }, - { apiKey: "cloud-test-key" }, + { apiKey: "cloud-test-key", fetch: fetchMock }, ); expect(response.stopReason).toBe("stop"); @@ -282,7 +277,7 @@ describe("ollama-cloud provider support", () => { expect(response.usage.output).toBe(5); }); test("streams tool calls and maps native tool stop reasons", async () => { - global.fetch = vi.fn(async () => + const fetchMock: FetchImpl = vi.fn(async () => createNdjsonResponse([ { model: "gpt-oss:120b", @@ -309,7 +304,7 @@ describe("ollama-cloud provider support", () => { eval_count: 2, }, ]), - ) as unknown as typeof fetch; + ); const response = stream( cloudModel, @@ -317,7 +312,7 @@ describe("ollama-cloud provider support", () => { messages: [{ role: "user", content: "Read README", timestamp: Date.now() }], tools: [readFileTool], }, - { apiKey: "cloud-test-key" }, + { apiKey: "cloud-test-key", fetch: fetchMock }, ); const eventTypes: string[] = []; for await (const event of response) { @@ -337,7 +332,7 @@ describe("ollama-cloud provider support", () => { test("converts replay history, tools, and images into native ollama chat payloads", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { @@ -347,7 +342,7 @@ describe("ollama-cloud provider support", () => { }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 3, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); const context: Context = { messages: [ @@ -388,7 +383,7 @@ describe("ollama-cloud provider support", () => { tools: [readFileTool], }; - await stream(cloudModel, context, { apiKey: "cloud-test-key" }).result(); + await stream(cloudModel, context, { apiKey: "cloud-test-key", fetch: fetchMock }).result(); const messages = requestBody?.messages as Array> | undefined; expect(requestBody?.model).toBe("gpt-oss:120b"); @@ -417,13 +412,13 @@ describe("ollama-cloud provider support", () => { test("strips `thinking` from assistant history messages on ollama-cloud", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { model: "gpt-oss:120b", message: { role: "assistant", content: "ok" }, done: false }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 1, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); const context: Context = { messages: [ @@ -460,7 +455,7 @@ describe("ollama-cloud provider support", () => { tools: [readFileTool], }; - await stream(cloudModel, context, { apiKey: "cloud-test-key" }).result(); + await stream(cloudModel, context, { apiKey: "cloud-test-key", fetch: fetchMock }).result(); const messages = requestBody?.messages as Array> | undefined; const assistant = messages?.find(message => message.role === "assistant"); @@ -476,7 +471,7 @@ describe("ollama-cloud provider support", () => { test("emits one Ollama system message per ordered system prompt entry", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { @@ -486,7 +481,7 @@ describe("ollama-cloud provider support", () => { }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 3, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); await stream( cloudModel, @@ -494,7 +489,7 @@ describe("ollama-cloud provider support", () => { systemPrompt: ["Stable instruction.", "Extra policy."], messages: [{ role: "user", content: "Hello", timestamp: Date.now() }], }, - { apiKey: "cloud-test-key" }, + { apiKey: "cloud-test-key", fetch: fetchMock }, ).result(); const messages = requestBody?.messages as Array> | undefined; @@ -507,54 +502,54 @@ describe("ollama-cloud provider support", () => { describe("mapToolChoice", () => { test("omits tool_choice when undefined or auto", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { model: "gpt-oss:120b", message: { role: "assistant", content: "ok" }, done: false }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 1, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); await stream( cloudModel, { messages: [{ role: "user", content: "hi", timestamp: Date.now() }], tools: [readFileTool] }, - { apiKey: "cloud-test-key", toolChoice: "auto" }, + { apiKey: "cloud-test-key", toolChoice: "auto", fetch: fetchMock }, ).result(); expect(requestBody?.tool_choice).toBeUndefined(); }); test("passes tool_choice: none when ToolChoice is none", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { model: "gpt-oss:120b", message: { role: "assistant", content: "ok" }, done: false }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 1, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); await stream( cloudModel, { messages: [{ role: "user", content: "hi", timestamp: Date.now() }], tools: [readFileTool] }, - { apiKey: "cloud-test-key", toolChoice: "none" }, + { apiKey: "cloud-test-key", toolChoice: "none", fetch: fetchMock }, ).result(); expect(requestBody?.tool_choice).toBe("none"); }); test("passes tool_choice: required when ToolChoice is required or any", async () => { let requestBody: Record | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as Record; return createNdjsonResponse([ { model: "gpt-oss:120b", message: { role: "assistant", content: "ok" }, done: false }, { model: "gpt-oss:120b", done: true, done_reason: "stop", prompt_eval_count: 1, eval_count: 1 }, ]); - }) as unknown as typeof fetch; + }); await stream( cloudModel, { messages: [{ role: "user", content: "hi", timestamp: Date.now() }], tools: [readFileTool] }, - { apiKey: "cloud-test-key", toolChoice: "required" }, + { apiKey: "cloud-test-key", toolChoice: "required", fetch: fetchMock }, ).result(); expect(requestBody?.tool_choice).toBe("required"); }); diff --git a/packages/ai/test/ollama-provider.test.ts b/packages/ai/test/ollama-provider.test.ts index ccf7ec666..746fcb559 100644 --- a/packages/ai/test/ollama-provider.test.ts +++ b/packages/ai/test/ollama-provider.test.ts @@ -1,15 +1,8 @@ -import { afterEach, describe, expect, test, vi } from "bun:test"; +import { describe, expect, test, vi } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai/effort"; import { ollamaModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { streamOllama } from "@oh-my-pi/pi-ai/providers/ollama"; -import type { Context, Model, Tool } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { Context, FetchImpl, Model, Tool } from "@oh-my-pi/pi-ai/types"; interface OllamaRequestBody { tools?: Array<{ function: { name: string } }>; @@ -18,7 +11,7 @@ interface OllamaRequestBody { describe("ollama local provider discovery", () => { test("applies /api/show context and thinking capabilities to OpenAI-compatible local models", async () => { - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/v1/models") { return new Response( @@ -41,9 +34,9 @@ describe("ollama local provider discovery", () => { ); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - const options = ollamaModelManagerOptions(); + const options = ollamaModelManagerOptions({ fetch: fetchMock }); const models = await options.fetchDynamicModels?.(); const model = models?.find(candidate => candidate.id === "deepseek-v4:latest"); @@ -55,7 +48,7 @@ describe("ollama local provider discovery", () => { }); test("remaps Ollama's unsupported reasoning levels and skips non-reasoning models", async () => { - global.fetch = vi.fn(async (input, init) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/v1/models") { return new Response( @@ -81,9 +74,9 @@ describe("ollama local provider discovery", () => { ); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - const models = await ollamaModelManagerOptions().fetchDynamicModels?.(); + const models = await ollamaModelManagerOptions({ fetch: fetchMock }).fetchDynamicModels?.(); const reasoningModel = models?.find(candidate => candidate.id === "gemma4:e4b"); const plainModel = models?.find(candidate => candidate.id === "llama-plain:latest"); @@ -100,13 +93,13 @@ describe("ollama local provider discovery", () => { describe("ollama tool forcing", () => { test("limits named forced tool requests to the selected tool", async () => { let requestBody: OllamaRequestBody | undefined; - global.fetch = vi.fn(async (_input, init) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { requestBody = JSON.parse(String(init?.body ?? "{}")) as OllamaRequestBody; return new Response(`${JSON.stringify({ done: true })}\n`, { status: 200, headers: { "Content-Type": "application/x-ndjson" }, }); - }) as unknown as typeof fetch; + }); const model = { id: "ggml-org/gemma-3-1b-it/GGUF", @@ -139,6 +132,7 @@ describe("ollama tool forcing", () => { for await (const event of streamOllama(model, context, { apiKey: "test-key", toolChoice: { type: "function", name: "write" }, + fetch: fetchMock, })) { eventTypes.push(event.type); } diff --git a/packages/ai/test/openai-codex-stream.test.ts b/packages/ai/test/openai-codex-stream.test.ts index 9d3f6d727..cda33fe5a 100644 --- a/packages/ai/test/openai-codex-stream.test.ts +++ b/packages/ai/test/openai-codex-stream.test.ts @@ -6,10 +6,9 @@ import { prewarmOpenAICodexResponses, streamOpenAICodexResponses, } from "@oh-my-pi/pi-ai/providers/openai-codex-responses"; -import type { Context, Model, ProviderSessionState } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, ProviderSessionState } from "@oh-my-pi/pi-ai/types"; import { getAgentDir, setAgentDir, TempDir } from "@oh-my-pi/pi-utils"; -const originalFetch = global.fetch; const originalAgentDir = getAgentDir(); const originalWebSocket = global.WebSocket; const originalCodexWebSocketRetryBudget = Bun.env.PI_CODEX_WEBSOCKET_RETRY_BUDGET; @@ -31,7 +30,6 @@ function restoreEnv(name: string, value: string | undefined): void { } afterEach(() => { - global.fetch = originalFetch; global.WebSocket = originalWebSocket; setAgentDir(originalAgentDir); restoreEnv("PI_CODEX_WEBSOCKET_RETRY_BUDGET", originalCodexWebSocketRetryBudget); @@ -256,10 +254,10 @@ describe("openai-codex streaming", () => { const context = createCodexTestContext(); const requestedUrls: string[] = []; const sse = createCompletedCodexSse("Hello"); - global.fetch = vi.fn(async (input: string | URL) => { + const fetchMock = vi.fn(async (input: string | URL) => { requestedUrls.push(typeof input === "string" ? input : input.toString()); return new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }); - }) as unknown as typeof fetch; + }); for (const baseUrl of [ undefined, @@ -268,7 +266,10 @@ describe("openai-codex streaming", () => { "https://chatgpt.com/backend-api/codex/responses", ]) { const model = { ...createCodexTestModel(baseUrl), preferWebsockets: false }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); expect(result.stopReason).toBe("stop"); } @@ -297,12 +298,15 @@ describe("openai-codex streaming", () => { `data: ${JSON.stringify({ type: "response.output_item.done", item: { type: "function_call", id: "fc_1", call_id: "call_1", name: "read_file", arguments: '{"path":"README.md"}' } })}`, `data: ${JSON.stringify({ type: "response.completed", response: { status: "completed", usage: { input_tokens: 5, output_tokens: 3, total_tokens: 8, input_tokens_details: { cached_tokens: 0 } } } })}`, ].join("\n\n")}\n\n`; - global.fetch = vi.fn( + const fetchMock = vi.fn( async () => new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }), - ) as unknown as typeof fetch; + ); const model = { ...createCodexTestModel("https://chatgpt.com/backend-api"), preferWebsockets: false }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); const toolCall = result.content.find(c => c.type === "toolCall"); if (toolCall?.type !== "toolCall") throw new Error("expected a finalized toolCall block"); @@ -316,13 +320,14 @@ describe("openai-codex streaming", () => { setAgentDir(tempDir.path()); const token = createCodexTestToken(); const context = createCodexTestContext(); - global.fetch = ((input: string | URL | Request, init?: RequestInit) => - Promise.resolve(createNoProgressCodexSse(getRequestSignal(input, init)))) as typeof fetch; + const fetchMock: FetchImpl = (input: string | URL | Request, init?: RequestInit) => + Promise.resolve(createNoProgressCodexSse(getRequestSignal(input, init))); const controller = new AbortController(); setTimeout(() => controller.abort(), 30); const model = { ...createCodexTestModel("https://chatgpt.com/backend-api"), preferWebsockets: false }; const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, signal: controller.signal, }).result(); @@ -511,7 +516,6 @@ describe("openai-codex streaming", () => { headers: { "content-type": "text/event-stream" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; class QueueOverflowWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -532,6 +536,7 @@ describe("openai-codex streaming", () => { const providerSessionState = new Map(); const model = createCodexTestModel("https://chatgpt.com/backend-api"); const result = await streamOpenAICodexResponses(model, createCodexTestContext(), { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-queue-overflow-session", providerSessionState, @@ -674,8 +679,6 @@ describe("openai-codex streaming", () => { return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", name: "GPT-5.1 Codex", @@ -694,7 +697,7 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const streamResult = streamOpenAICodexResponses(model, context, { apiKey: token }); + const streamResult = streamOpenAICodexResponses(model, context, { apiKey: token, fetch: fetchMock as FetchImpl }); let sawTextDelta = false; let sawDone = false; @@ -737,7 +740,6 @@ describe("openai-codex streaming", () => { headers: { "content-type": "text/event-stream" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -758,6 +760,7 @@ describe("openai-codex streaming", () => { }; const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, serviceTier: "priority", }).result(); @@ -807,7 +810,6 @@ describe("openai-codex streaming", () => { } return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -827,7 +829,10 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); expect(result.stopReason).toBe("error"); expect(result.errorMessage).toContain("terminal completion event"); }); @@ -850,9 +855,9 @@ describe("openai-codex streaming", () => { `data: ${JSON.stringify({ type: "response.failed", code: "server_error", message: "late failure after terminal event" })}`, ].join("\n\n")}\n\n`; - global.fetch = vi.fn( + const fetchMock = vi.fn( async () => new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }), - ) as unknown as typeof fetch; + ); const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -871,7 +876,10 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); expect(result.stopReason).toBe("stop"); expect(result.content.find(block => block.type === "text")?.text).toBe("Hello"); }); @@ -907,7 +915,6 @@ describe("openai-codex streaming", () => { } return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -927,7 +934,10 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); expect(fetchMock).toHaveBeenCalledTimes(1); expect(result.stopReason).toBe("error"); expect((result.errorMessage ?? "").toLowerCase()).toContain("rate limit"); @@ -971,7 +981,6 @@ describe("openai-codex streaming", () => { } return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -991,7 +1000,10 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const result = await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + fetch: fetchMock as FetchImpl, + }).result(); expect(fetchMock).toHaveBeenCalledTimes(2); expect(result.stopReason).toBe("stop"); expect(result.content.find(block => block.type === "text")?.text).toBe("Hello after retry"); @@ -1074,8 +1086,6 @@ describe("openai-codex streaming", () => { return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", name: "GPT-5.1 Codex", @@ -1094,7 +1104,11 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - const streamResult = streamOpenAICodexResponses(model, context, { apiKey: token, sessionId }); + const streamResult = streamOpenAICodexResponses(model, context, { + apiKey: token, + sessionId, + fetch: fetchMock as FetchImpl, + }); await streamResult.result(); }); it("keeps prompt_cache_key separate from Codex conversation headers", async () => { @@ -1108,7 +1122,7 @@ describe("openai-codex streaming", () => { let capturedHeaders: Headers | undefined; let capturedBody: Record | undefined; - global.fetch = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { const url = typeof input === "string" ? input : input.toString(); if (url === "https://api.github.com/repos/openai/codex/releases/latest") { return new Response(JSON.stringify({ tag_name: "rust-v0.0.0" }), { status: 200 }); @@ -1126,9 +1140,10 @@ describe("openai-codex streaming", () => { }); } return new Response("not found", { status: 404 }); - }) as unknown as typeof fetch; + }); await streamOpenAICodexResponses(model, createCodexTestContext(), { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId, promptCacheKey, @@ -1209,8 +1224,6 @@ describe("openai-codex streaming", () => { return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const model = enrichModelThinking({ id: "gpt-5.3-codex", name: "GPT-5.3 Codex", @@ -1230,6 +1243,7 @@ describe("openai-codex streaming", () => { }; const streamResult = streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, reasoning: "minimal", }); @@ -1301,8 +1315,6 @@ describe("openai-codex streaming", () => { return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; - const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", name: "GPT-5.1 Codex", @@ -1322,7 +1334,7 @@ describe("openai-codex streaming", () => { }; // No sessionId provided - const streamResult = streamOpenAICodexResponses(model, context, { apiKey: token }); + const streamResult = streamOpenAICodexResponses(model, context, { apiKey: token, fetch: fetchMock as FetchImpl }); await streamResult.result(); }); @@ -1353,7 +1365,6 @@ describe("openai-codex streaming", () => { } return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; class FailingWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { super(url, options); @@ -1388,6 +1399,7 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); const streamResult = streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-session", providerSessionState, @@ -1421,7 +1433,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { return new Response(sse, { headers: { "content-type": "text/event-stream" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; let constructorCount = 0; class FailingConnectWebSocket extends MockWebSocket { @@ -1457,6 +1468,7 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-fatal-fallback-session", providerSessionState, @@ -1496,7 +1508,6 @@ describe("openai-codex streaming", () => { expect(headers.get("x-models-etag")).toBe("models-etag-1"); return new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; class HandshakeWebSocket extends MockWebSocket { handshakeHeaders = { @@ -1540,11 +1551,13 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); await streamOpenAICodexResponses(websocketModel, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-handshake-session", providerSessionState, }).result(); await streamOpenAICodexResponses(sseModel, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-handshake-session", providerSessionState, @@ -1566,7 +1579,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; class ServiceTierWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -1621,6 +1633,7 @@ describe("openai-codex streaming", () => { }; const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, serviceTier: "priority", sessionId: "ws-service-tier-session", @@ -1645,7 +1658,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; class DeltaWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -1686,6 +1698,7 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "First question", timestamp: Date.now() }], }; const firstResponse = await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-delta-session", providerSessionState, @@ -1699,6 +1712,7 @@ describe("openai-codex streaming", () => { ], }; await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-delta-session", providerSessionState, @@ -1749,7 +1763,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; class PreviousResponseMissingWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -1807,6 +1820,7 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "First question", timestamp: Date.now() }], }; const firstResponse = await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-expired-previous-response-session", providerSessionState, @@ -1821,6 +1835,7 @@ describe("openai-codex streaming", () => { }; const secondResponse = await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-expired-previous-response-session", providerSessionState, @@ -1869,7 +1884,6 @@ describe("openai-codex streaming", () => { capturedBodies.push(JSON.parse(String(init?.body)) as Record); return new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", name: "GPT-5.1 Codex", @@ -1887,8 +1901,12 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }; - await streamOpenAICodexResponses(model, context, { apiKey: token }).result(); - await streamOpenAICodexResponses(model, context, { apiKey: token, textVerbosity: "high" }).result(); + await streamOpenAICodexResponses(model, context, { apiKey: token, fetch: fetchMock as FetchImpl }).result(); + await streamOpenAICodexResponses(model, context, { + apiKey: token, + textVerbosity: "high", + fetch: fetchMock as FetchImpl, + }).result(); expect((capturedBodies[0]?.text as { verbosity?: string } | undefined)?.verbosity).toBe("low"); expect((capturedBodies[1]?.text as { verbosity?: string } | undefined)?.verbosity).toBe("high"); @@ -1908,7 +1926,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; class WebSocketV2HeaderProbe extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -1945,6 +1962,7 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-v2-session", providerSessionState, @@ -1968,7 +1986,6 @@ describe("openai-codex streaming", () => { headers: { "content-type": "text/event-stream" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; let sendCount = 0; class IdleWebSocket extends MockWebSocket { @@ -2010,6 +2027,7 @@ describe("openai-codex streaming", () => { const controller = new AbortController(); setTimeout(() => controller.abort(), 30); const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-idle-timeout-session", providerSessionState, @@ -2027,7 +2045,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not run once the websocket stream becomes replay-unsafe"); }); - global.fetch = fetchMock as unknown as typeof fetch; let sendCount = 0; let interval: NodeJS.Timeout | undefined; @@ -2077,6 +2094,7 @@ describe("openai-codex streaming", () => { const model = createCodexTestModel("https://chatgpt.com/backend-api"); const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, createCodexTestContext(), { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-no-progress-session", providerSessionState, @@ -2104,7 +2122,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not run for degenerate tool-call arguments"); }); - global.fetch = fetchMock as unknown as typeof fetch; let sendCount = 0; let closeCount = 0; @@ -2147,6 +2164,7 @@ describe("openai-codex streaming", () => { const model = createCodexTestModel("https://chatgpt.com/backend-api"); const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, createCodexTestContext(), { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-whitespace-arguments-session", providerSessionState, @@ -2168,7 +2186,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not run when the websocket recovers"); }); - global.fetch = fetchMock as unknown as typeof fetch; let connectionCount = 0; let closeCount = 0; @@ -2237,6 +2254,7 @@ describe("openai-codex streaming", () => { const model = createCodexTestModel("https://chatgpt.com/backend-api"); const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, createCodexTestContext(), { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-whitespace-recovery-session", providerSessionState, @@ -2268,7 +2286,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called when websocket retry succeeds"); }); - global.fetch = fetchMock as unknown as typeof fetch; let constructorCount = 0; const requestTypes: string[] = []; @@ -2317,6 +2334,7 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-retry-close-session", providerSessionState, @@ -2356,7 +2374,6 @@ describe("openai-codex streaming", () => { } return new Response("not found", { status: 404 }); }); - global.fetch = fetchMock as unknown as typeof fetch; class UnavailableBeforeStreamWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -2391,6 +2408,7 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); const result = await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-unavailable-session", providerSessionState, @@ -2421,7 +2439,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; const sentTypesByConnection: string[][] = []; let constructorCount = 0; @@ -2505,6 +2522,7 @@ describe("openai-codex streaming", () => { const providerSessionState = new Map(); const firstResult = await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-abort-reset-session", providerSessionState, @@ -2516,6 +2534,7 @@ describe("openai-codex streaming", () => { secondAbortController.abort(); }; const secondResult = await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-abort-reset-session", signal: secondAbortController.signal, @@ -2524,6 +2543,7 @@ describe("openai-codex streaming", () => { expect(secondResult.stopReason).toBe("aborted"); const thirdResult = await streamOpenAICodexResponses(model, thirdContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-abort-reset-session", providerSessionState, @@ -2547,7 +2567,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; const sentTypes: string[] = []; let constructorCount = 0; @@ -2622,6 +2641,7 @@ describe("openai-codex streaming", () => { const providerSessionState = new Map(); const firstResult = await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-error-reset-session", providerSessionState, @@ -2629,6 +2649,7 @@ describe("openai-codex streaming", () => { expect(firstResult.role).toBe("assistant"); const secondResult = await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-error-reset-session", providerSessionState, @@ -2637,6 +2658,7 @@ describe("openai-codex streaming", () => { expect(secondResult.errorMessage).toContain("simulated request error"); const thirdResult = await streamOpenAICodexResponses(model, thirdContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-error-reset-session", providerSessionState, @@ -2669,7 +2691,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn( async () => new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }), ); - global.fetch = fetchMock as unknown as typeof fetch; class MalformedMessageWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -2703,6 +2724,7 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-malformed-json-session", providerSessionState: new Map(), @@ -2736,7 +2758,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn( async () => new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }), ); - global.fetch = fetchMock as unknown as typeof fetch; class BufferedCloseWebSocket extends MockWebSocket { constructor(url: string, options?: { headers?: WsHeaders }) { @@ -2783,6 +2804,7 @@ describe("openai-codex streaming", () => { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], }, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-buffered-close-session", providerSessionState: new Map(), @@ -2819,7 +2841,6 @@ describe("openai-codex streaming", () => { sseModelsEtags.push(headers.get("x-models-etag")); return new Response(sse, { status: 200, headers: { "content-type": "text/event-stream" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; const requestTypes: string[] = []; class DivergedAppendWebSocket extends MockWebSocket { @@ -2877,16 +2898,19 @@ describe("openai-codex streaming", () => { const providerSessionState = new Map(); await streamOpenAICodexResponses(websocketModel, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-diverged-session", providerSessionState, }).result(); await streamOpenAICodexResponses(websocketModel, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-diverged-session", providerSessionState, }).result(); await streamOpenAICodexResponses(sseModel, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-diverged-session", providerSessionState, @@ -2911,7 +2935,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; let constructorCount = 0; let sendCount = 0; @@ -2951,7 +2974,11 @@ describe("openai-codex streaming", () => { }; const providerSessionState = new Map(); - await prewarmOpenAICodexResponses(model, { apiKey: token, sessionId: "ws-reuse-session", providerSessionState }); + await prewarmOpenAICodexResponses(model, { + apiKey: token, + sessionId: "ws-reuse-session", + providerSessionState, + }); const firstContext: Context = { systemPrompt: ["You are a helpful assistant."], @@ -2966,11 +2993,13 @@ describe("openai-codex streaming", () => { }; await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-reuse-session", providerSessionState, }).result(); await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-reuse-session", providerSessionState, @@ -3018,7 +3047,6 @@ describe("openai-codex streaming", () => { callCount += 1; return new Response(sse, { status: 200, headers: responseHeaders }); }); - global.fetch = fetchMock as unknown as typeof fetch; const model: Model<"openai-codex-responses"> = { id: "gpt-5.1-codex", @@ -3040,11 +3068,13 @@ describe("openai-codex streaming", () => { const providerSessionState = new Map(); await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "turn-state-session", providerSessionState, }).result(); await streamOpenAICodexResponses(model, context, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "turn-state-session", providerSessionState, @@ -3074,7 +3104,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; let constructorCount = 0; let sendCount = 0; @@ -3112,6 +3141,7 @@ describe("openai-codex streaming", () => { }; await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-idle-reuse-session", providerSessionState, @@ -3123,6 +3153,7 @@ describe("openai-codex streaming", () => { await new Promise(resolve => setTimeout(resolve, 30)); const second = await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-idle-reuse-session", providerSessionState, @@ -3153,7 +3184,6 @@ describe("openai-codex streaming", () => { const fetchMock = vi.fn(async () => { throw new Error("SSE fallback should not be called"); }); - global.fetch = fetchMock as unknown as typeof fetch; let constructorCount = 0; let sendCount = 0; @@ -3208,6 +3238,7 @@ describe("openai-codex streaming", () => { }; const first = await streamOpenAICodexResponses(model, firstContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-stale-frame-session", providerSessionState, @@ -3215,6 +3246,7 @@ describe("openai-codex streaming", () => { expect(first.stopReason).toBe("stop"); const second = await streamOpenAICodexResponses(model, secondContext, { + fetch: fetchMock as FetchImpl, apiKey: token, sessionId: "ws-stale-frame-session", providerSessionState, diff --git a/packages/ai/test/openai-codex-usage.test.ts b/packages/ai/test/openai-codex-usage.test.ts index f7f927a9d..6509ad545 100644 --- a/packages/ai/test/openai-codex-usage.test.ts +++ b/packages/ai/test/openai-codex-usage.test.ts @@ -7,6 +7,7 @@ * widget lose per-model visibility. */ import { describe, expect, it } from "bun:test"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { openaiCodexUsageProvider } from "@oh-my-pi/pi-ai/usage/openai-codex"; const accessTokenFixture = (() => { @@ -44,7 +45,7 @@ function makePayload() { }; } -function fakeFetch(payload: unknown): typeof fetch { +function fakeFetch(payload: unknown): FetchImpl { const fn = async () => new Response(JSON.stringify(payload), { status: 200, headers: { "content-type": "application/json" } }); return fn as unknown as typeof fetch; diff --git a/packages/ai/test/openai-completions-compat.test.ts b/packages/ai/test/openai-completions-compat.test.ts index c26b35bc4..6d499fa7a 100644 --- a/packages/ai/test/openai-completions-compat.test.ts +++ b/packages/ai/test/openai-completions-compat.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { applyOpenRouterRoutingVariant, @@ -7,13 +7,7 @@ import { streamOpenAICompletions, } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { type ResolvedOpenAICompat, resolveOpenAICompat } from "@oh-my-pi/pi-ai/providers/openai-completions-compat"; -import type { AssistantMessage, Context, Model, OpenAICompat } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { AssistantMessage, Context, FetchImpl, Model, OpenAICompat } from "@oh-my-pi/pi-ai/types"; function createAbortedSignal(): AbortSignal { const controller = new AbortController(); @@ -46,12 +40,12 @@ function createSseResponse(events: unknown[]): Response { }); } -function createMockFetch(events: unknown[]): typeof fetch { +function createMockFetch(events: unknown[]): FetchImpl { async function mockFetch(_input: string | URL | Request, _init?: RequestInit): Promise { return createSseResponse(events); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return Object.assign(mockFetch, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -71,9 +65,10 @@ async function captureOpenAICompletionsPayload( context: Context = baseContext(), ): Promise { const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions(model, context, { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), onPayload: payload => resolve(payload), }); @@ -407,7 +402,7 @@ describe("openai-completions compatibility", () => { ...getBundledModel("openai", "gpt-4o-mini"), api: "openai-completions", }; - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-test", object: "chat.completion.chunk", @@ -435,7 +430,10 @@ describe("openai-completions compatibility", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("stop"); expect(result.usage.input).toBe(10); expect(result.usage.output).toBe(3); @@ -469,7 +467,7 @@ describe("openai-completions compatibility", () => { ...getBundledModel("openai", "gpt-4o-mini"), api: "openai-completions", }; - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-end", object: "chat.completion.chunk", @@ -487,7 +485,10 @@ describe("openai-completions compatibility", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("stop"); expect(result.content[0]).toMatchObject({ type: "text", text: "done" }); }); @@ -505,9 +506,10 @@ describe("openai-completions compatibility", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), onPayload: payload => resolve(payload), }); @@ -526,7 +528,7 @@ describe("openai-completions compatibility", () => { ...getBundledModel("openai", "gpt-4o-mini"), api: "openai-completions", }; - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-reasoning-text", object: "chat.completion.chunk", @@ -549,7 +551,10 @@ describe("openai-completions compatibility", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toContainEqual({ type: "thinking", thinking: "inspect tool output", @@ -769,7 +774,7 @@ describe("kimi model detection via detectCompat", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions( model, { @@ -788,6 +793,7 @@ describe("kimi model detection via detectCompat", () => { }, { apiKey: "test-key", + fetch: fetchMock, reasoning: "high", signal: createAbortedSignal(), onPayload: payload => resolve(payload), @@ -835,7 +841,7 @@ describe("kimi model detection via detectCompat", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions( model, { @@ -854,6 +860,7 @@ describe("kimi model detection via detectCompat", () => { }, { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), onPayload: payload => resolve(payload), }, @@ -905,7 +912,7 @@ describe("kimi model detection via detectCompat", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions( model, { @@ -924,6 +931,7 @@ describe("kimi model detection via detectCompat", () => { }, { apiKey: "test-key", + fetch: fetchMock, reasoning: "high", // Forced tool choice triggers `disableReasoningOnForcedToolChoice` // for Kimi, suppressing reasoning_effort on the wire body. @@ -994,7 +1002,7 @@ describe("kimi model detection via detectCompat", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions( model, { @@ -1013,6 +1021,7 @@ describe("kimi model detection via detectCompat", () => { }, { apiKey: "test-key", + fetch: fetchMock, reasoning: "high", signal: createAbortedSignal(), onPayload: payload => resolve(payload), @@ -1077,7 +1086,7 @@ describe("kimi model detection via detectCompat", () => { }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions( model, { @@ -1096,6 +1105,7 @@ describe("kimi model detection via detectCompat", () => { }, { apiKey: "test-key", + fetch: fetchMock, reasoning, signal: createAbortedSignal(), onPayload: payload => resolve(payload), @@ -1246,7 +1256,7 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { it("strips leaked <\uff5cDSML\uff5c...\uff5c> markers from visible content", async () => { const model = nvidiaDeepseekModel(); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-nim-1", object: "chat.completion.chunk", @@ -1269,7 +1279,10 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) @@ -1281,7 +1294,7 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { it("holds back partial token split across chunks", async () => { const model = nvidiaDeepseekModel(); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-nim-2", object: "chat.completion.chunk", @@ -1306,7 +1319,10 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) @@ -1316,7 +1332,7 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { it("flushes a dangling partial open delimiter at end of stream", async () => { const model = nvidiaDeepseekModel(); - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-nim-3", object: "chat.completion.chunk", @@ -1336,7 +1352,10 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { // At end-of-stream we have no way to know whether the partial is a real token, // so we emit it verbatim rather than swallow legitimate text forever. - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) @@ -1352,7 +1371,7 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { baseUrl: "https://integrate.api.nvidia.com/v1", id: "meta/llama-3.3-70b-instruct", }; - global.fetch = createMockFetch([ + const fetchMock = createMockFetch([ { id: "chatcmpl-nim-4", object: "chat.completion.chunk", @@ -1370,7 +1389,10 @@ describe("NVIDIA NIM DeepSeek special-token stripping", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); const text = result.content .filter(b => b.type === "text") .map(b => (b as { text: string }).text) @@ -1457,9 +1479,10 @@ describe("openrouterVariant request integration", () => { it("appends the configured variant suffix to params.model for OpenRouter requests", async () => { const model = getBundledModel("openrouter", "anthropic/claude-sonnet-4") as Model<"openai-completions">; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), openrouterVariant: "nitro", onPayload: payload => resolve(payload), @@ -1475,9 +1498,10 @@ describe("openrouterVariant request integration", () => { id: `${base.id}:online`, }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), openrouterVariant: "nitro", onPayload: payload => resolve(payload), @@ -1492,9 +1516,10 @@ describe("openrouterVariant request integration", () => { api: "openai-completions", }; const { promise, resolve } = Promise.withResolvers(); - global.fetch = createMockFetch(["[DONE]"]); + const fetchMock = createMockFetch(["[DONE]"]); streamOpenAICompletions(model, baseContext(), { apiKey: "test-key", + fetch: fetchMock, signal: createAbortedSignal(), openrouterVariant: "nitro", onPayload: payload => resolve(payload), diff --git a/packages/ai/test/openai-completions-disable-reasoning.test.ts b/packages/ai/test/openai-completions-disable-reasoning.test.ts index 4d3e4d19c..fa933659d 100644 --- a/packages/ai/test/openai-completions-disable-reasoning.test.ts +++ b/packages/ai/test/openai-completions-disable-reasoning.test.ts @@ -1,18 +1,12 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai/effort"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const testContext: Context = { messages: [{ role: "user", content: "hello", timestamp: 0 }], }; -afterEach(() => { - global.fetch = originalFetch; -}); - function createSseResponse(events: unknown[]): Response { const payload = `${events.map(event => `data: ${typeof event === "string" ? event : JSON.stringify(event)}`).join("\n\n")}\n\n`; return new Response(payload, { @@ -53,7 +47,7 @@ function createFireworksReasoningEffortModel(): Model<"openai-completions"> { async function captureDisableReasoningPayload(model: Model<"openai-completions">): Promise> { let payload: Record | undefined; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, init?: RequestInit): Promise => { payload = JSON.parse(typeof init?.body === "string" ? init.body : "{}") as Record; return createSseResponse([ @@ -74,11 +68,12 @@ async function captureDisableReasoningPayload(model: Model<"openai-completions"> "[DONE]", ]); }, - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key", + fetch: fetchMock, disableReasoning: true, }).result(); diff --git a/packages/ai/test/openai-completions-progress-chunk.test.ts b/packages/ai/test/openai-completions-progress-chunk.test.ts index db4b1058d..7e92a8990 100644 --- a/packages/ai/test/openai-completions-progress-chunk.test.ts +++ b/packages/ai/test/openai-completions-progress-chunk.test.ts @@ -1,13 +1,11 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { getOpenAICompletionsStreamIdleTimeoutFallbackMs, isOpenAICompletionsProgressChunk, streamOpenAICompletions, } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const openAICompletionsModel = { ...(getBundledModel("openai", "gpt-4o-mini") as Model<"openai-completions">), @@ -80,9 +78,6 @@ function createKeepaliveOnlyCompletionsResponse(modelId: string, signal: AbortSi }); } -afterEach(() => { - global.fetch = originalFetch; -}); describe("getOpenAICompletionsStreamIdleTimeoutFallbackMs", () => { it("widens GLM 5.1 coding-plan stream watchdogs", () => { const model = { @@ -278,13 +273,14 @@ describe("isOpenAICompletionsProgressChunk", () => { }); describe("provider integration", () => { it("times out a completions stream whose keepalives never make progress", async () => { - global.fetch = ((input: string | URL | Request, init?: RequestInit) => + const fetchMock: FetchImpl = (input: string | URL | Request, init?: RequestInit) => Promise.resolve( createKeepaliveOnlyCompletionsResponse(openAICompletionsModel.id, getRequestSignal(input, init)), - )) as typeof fetch; + ); const result = await streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", + fetch: fetchMock, streamFirstEventTimeoutMs: 1_000, streamIdleTimeoutMs: 20, }).result(); diff --git a/packages/ai/test/openai-completions-upstream-provider.test.ts b/packages/ai/test/openai-completions-upstream-provider.test.ts index 52e2f4e54..067281388 100644 --- a/packages/ai/test/openai-completions-upstream-provider.test.ts +++ b/packages/ai/test/openai-completions-upstream-provider.test.ts @@ -1,9 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const model = { ...(getBundledModel("openai", "gpt-4o-mini") as Model<"openai-completions">), @@ -26,16 +24,12 @@ function chunk(extra: Record): Record { return { id: "gen-1", object: "chat.completion.chunk", created: 0, model: model.id, ...extra }; } -afterEach(() => { - global.fetch = originalFetch; -}); - describe("openai-completions upstream provider capture", () => { // Contract: aggregators (OpenRouter, …) report the upstream provider that served // the request via a top-level `provider` field on every chunk. We surface it on // the assistant message so telemetry/session logs can attribute routing. it("records the aggregator-reported upstream provider from the stream", async () => { - global.fetch = ((_input: string | URL | Request, _init?: RequestInit) => + const fetchMock: FetchImpl = () => Promise.resolve( createSseResponse([ chunk({ provider: "Anthropic", choices: [{ index: 0, delta: { content: "Hi" } }] }), @@ -45,9 +39,12 @@ describe("openai-completions upstream provider capture", () => { usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, }), ]), - )) as typeof fetch; + ); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.upstreamProvider).toBe("Anthropic"); expect(result.stopReason).toBe("stop"); @@ -55,7 +52,7 @@ describe("openai-completions upstream provider capture", () => { }); it("leaves upstreamProvider undefined when no provider field is present", async () => { - global.fetch = ((_input: string | URL | Request, _init?: RequestInit) => + const fetchMock: FetchImpl = () => Promise.resolve( createSseResponse([ chunk({ choices: [{ index: 0, delta: { content: "Hi" } }] }), @@ -64,9 +61,12 @@ describe("openai-completions upstream provider capture", () => { usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, }), ]), - )) as typeof fetch; + ); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.upstreamProvider).toBeUndefined(); }); diff --git a/packages/ai/test/openai-first-event-timeout.test.ts b/packages/ai/test/openai-first-event-timeout.test.ts index 4886cf79c..ca39097de 100644 --- a/packages/ai/test/openai-first-event-timeout.test.ts +++ b/packages/ai/test/openai-first-event-timeout.test.ts @@ -1,14 +1,12 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamAzureOpenAIResponses } from "@oh-my-pi/pi-ai/providers/azure-openai-responses"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; import { streamSimple } from "@oh-my-pi/pi-ai/stream"; -import type { Context, Model, TextContent } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, TextContent } from "@oh-my-pi/pi-ai/types"; import { waitForDelayOrAbort } from "./helpers"; -const originalFetch = global.fetch; - const openAIResponsesModel = getBundledModel("openai", "gpt-5-mini") as Model<"openai-responses">; const openAICompletionsModel = { ...(getBundledModel("openai", "gpt-4o-mini") as Model<"openai-completions">), @@ -97,12 +95,12 @@ function createHangingSseResponse(signal: AbortSignal | undefined): Response { }); } -function createHangingFetch(): typeof fetch { +function createHangingFetch(): FetchImpl { async function mockFetch(input: string | URL | Request, init?: RequestInit): Promise { return createHangingSseResponse(getRequestSignal(input, init)); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return mockFetch as typeof fetch; } function createSseResponse(events: unknown[]): Response { @@ -169,14 +167,14 @@ function createDelayedFetch( delayMs: number, responseFactory: () => Response, onRequest?: (input: string | URL | Request, init: RequestInit | undefined) => void, -): typeof fetch { +): FetchImpl { async function mockFetch(input: string | URL | Request, init?: RequestInit): Promise { onRequest?.(input, init); await waitForDelayOrAbort(delayMs, getRequestSignal(input, init)); return responseFactory(); } - return Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + return mockFetch as typeof fetch; } function createOpenAIResponsesSuccessResponse(): Response { @@ -253,28 +251,33 @@ function createOllamaChatSuccessResponse(): Response { } async function expectFirstEventTimeout( - run: (streamFirstEventTimeoutMs: number) => Promise<{ stopReason: string; errorMessage?: string }>, + run: ( + streamFirstEventTimeoutMs: number, + fetchMock: FetchImpl, + ) => Promise<{ stopReason: string; errorMessage?: string }>, expectedMessage: string, ): Promise { - global.fetch = createHangingFetch(); - - const result = await run(20); + const fetchMock = createHangingFetch(); + const result = await run(20, fetchMock); expect(result.stopReason).toBe("error"); expect(result.errorMessage).toBe(expectedMessage); } async function expectRequestSetupTimeout( - run: (streamFirstEventTimeoutMs: number) => Promise<{ stopReason: string; errorMessage?: string }>, + run: ( + streamFirstEventTimeoutMs: number, + fetchMock: FetchImpl, + ) => Promise<{ stopReason: string; errorMessage?: string }>, expectedMessage: string, responseFactory: () => Response, ): Promise { const timeoutHeaders: string[] = []; - global.fetch = createDelayedFetch(30, responseFactory, (input, init) => { + const fetchMock = createDelayedFetch(30, responseFactory, (input, init) => { timeoutHeaders.push(getRequestHeader(input, init, "X-Stainless-Timeout") ?? ""); }); - const result = await run(20); + const result = await run(20, fetchMock); expect(result.stopReason).toBe("error"); expect(result.errorMessage).toBe(expectedMessage); @@ -285,14 +288,15 @@ async function expectCallerAbort( run: ( signal: AbortSignal, streamFirstEventTimeoutMs: number, + fetchMock: FetchImpl, ) => Promise<{ stopReason: string; errorMessage?: string }>, unexpectedMessage: string, ): Promise { - global.fetch = createHangingFetch(); + const fetchMock = createHangingFetch(); const controller = new AbortController(); setTimeout(() => controller.abort(), 5); - const result = await run(controller.signal, 50); + const result = await run(controller.signal, 50, fetchMock); expect(result.stopReason).toBe("aborted"); expect(result.errorMessage).not.toBe(unexpectedMessage); @@ -306,7 +310,10 @@ function getFirstTextContent(result: { content: unknown[] }): TextContent | unde } async function expectDelayedRequestSetupSucceeds( - run: (streamFirstEventTimeoutMs: number) => Promise<{ stopReason: string; content: unknown[] }>, + run: ( + streamFirstEventTimeoutMs: number, + fetchMock: FetchImpl, + ) => Promise<{ stopReason: string; content: unknown[] }>, responseFactory: () => Response, ): Promise { // The watchdog must cover request setup (connection + first byte). We simulate @@ -315,35 +322,32 @@ async function expectDelayedRequestSetupSucceeds( // parallel test processes) can never trip the watchdog on a request that is // supposed to succeed. The complementary firing tests pin the budget < latency // case with their own short timeouts, so the contrast is preserved. - global.fetch = createDelayedFetch(30, responseFactory); - - const result = await run(5_000); + const fetchMock = createDelayedFetch(30, responseFactory); + const result = await run(5_000, fetchMock); expect(result.stopReason).toBe("stop"); expect(getFirstTextContent(result)).toMatchObject({ type: "text", text: "Hello delayed" }); } -afterEach(() => { - global.fetch = originalFetch; -}); - describe("OpenAI-family first-event timeouts", () => { it("surfaces the OpenAI responses first-event timeout message instead of a generic abort", async () => { await expectFirstEventTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI responses stream timed out while waiting for the first event", ); }); it("times out OpenAI responses before the stream opens and forwards the budget to the SDK request", async () => { await expectRequestSetupTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI responses stream timed out while waiting for the first event", createOpenAIResponsesSuccessResponse, @@ -356,13 +360,14 @@ describe("OpenAI-family first-event timeouts", () => { const timeoutHeaders: string[] = []; Bun.env.PI_OPENAI_STREAM_IDLE_TIMEOUT_MS = "1500"; Bun.env.PI_STREAM_FIRST_EVENT_TIMEOUT_MS = "20"; - global.fetch = createDelayedFetch(30, createOpenAIResponsesSuccessResponse, (input, init) => { + const fetchMock = createDelayedFetch(30, createOpenAIResponsesSuccessResponse, (input, init) => { timeoutHeaders.push(getRequestHeader(input, init, "X-Stainless-Timeout") ?? ""); }); try { const result = await streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("stop"); @@ -387,11 +392,12 @@ describe("OpenAI-family first-event timeouts", () => { const previousGenericFirstEventTimeout = Bun.env.PI_STREAM_FIRST_EVENT_TIMEOUT_MS; Bun.env.PI_OPENAI_STREAM_IDLE_TIMEOUT_MS = "1500"; Bun.env.PI_STREAM_FIRST_EVENT_TIMEOUT_MS = "20"; - global.fetch = createDelayedFetch(30, createOllamaChatSuccessResponse); + const fetchMock = createDelayedFetch(30, createOllamaChatSuccessResponse); try { const result = await streamSimple(ollamaChatModel, baseContext(), { apiKey: "test-key", + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("stop"); @@ -416,7 +422,7 @@ describe("OpenAI-family first-event timeouts", () => { const timeoutHeaders: string[] = []; Bun.env.PI_OPENAI_STREAM_FIRST_EVENT_TIMEOUT_MS = "1500"; Bun.env.PI_STREAM_FIRST_EVENT_TIMEOUT_MS = "20"; - global.fetch = createDelayedFetch(30, createOpenAIResponsesSuccessResponse, (input, init) => { + const fetchMock = createDelayedFetch(30, createOpenAIResponsesSuccessResponse, (input, init) => { timeoutHeaders.push(getRequestHeader(input, init, "X-Stainless-Timeout") ?? ""); }); @@ -424,6 +430,7 @@ describe("OpenAI-family first-event timeouts", () => { const result = await streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", streamIdleTimeoutMs: 5_000, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("stop"); @@ -444,13 +451,14 @@ describe("OpenAI-family first-event timeouts", () => { }); it("times out OpenAI responses streams that only emit no-progress status events", async () => { - global.fetch = ((input: string | URL | Request, init?: RequestInit) => - Promise.resolve(createNoProgressOpenAIResponsesStream(getRequestSignal(input, init)))) as typeof fetch; + const fetchMock: FetchImpl = (input: string | URL | Request, init?: RequestInit) => + Promise.resolve(createNoProgressOpenAIResponsesStream(getRequestSignal(input, init))); const result = await streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs: 1_000, streamIdleTimeoutMs: 20, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("error"); @@ -467,7 +475,7 @@ describe("OpenAI-family first-event timeouts", () => { }); it("forwards streamSimple per-call timeout options to OpenAI-family providers", async () => { - global.fetch = createHangingFetch(); + const fetchMock = createHangingFetch(); const controller = new AbortController(); const abortTimer = setTimeout(() => controller.abort(new Error("fallback abort")), 200); abortTimer.unref(); @@ -478,6 +486,7 @@ describe("OpenAI-family first-event timeouts", () => { signal: controller.signal, streamFirstEventTimeoutMs: 20, streamIdleTimeoutMs: 20, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("error"); @@ -489,20 +498,22 @@ describe("OpenAI-family first-event timeouts", () => { it("surfaces the OpenAI completions first-event timeout message", async () => { await expectFirstEventTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI completions stream timed out while waiting for the first event", ); }); it("times out OpenAI completions before the stream opens and forwards the budget to the SDK request", async () => { await expectRequestSetupTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI completions stream timed out while waiting for the first event", () => createOpenAICompletionsSuccessResponse(openAICompletionsModel.id), @@ -511,24 +522,26 @@ describe("OpenAI-family first-event timeouts", () => { it("surfaces the Azure OpenAI responses first-event timeout message", async () => { await expectFirstEventTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamAzureOpenAIResponses(azureOpenAIResponsesModel, baseContext(), { apiKey: "test-key", azureBaseUrl: azureOpenAIResponsesModel.baseUrl, azureApiVersion: "v1", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "Azure OpenAI responses stream timed out while waiting for the first event", ); }); it("times out Azure OpenAI responses before the stream opens and forwards the budget to the SDK request", async () => { await expectRequestSetupTimeout( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamAzureOpenAIResponses(azureOpenAIResponsesModel, baseContext(), { apiKey: "test-key", azureBaseUrl: azureOpenAIResponsesModel.baseUrl, azureApiVersion: "v1", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "Azure OpenAI responses stream timed out while waiting for the first event", createOpenAIResponsesSuccessResponse, @@ -536,8 +549,8 @@ describe("OpenAI-family first-event timeouts", () => { }); it("times out Azure responses streams that only emit no-progress status events", async () => { - global.fetch = ((input: string | URL | Request, init?: RequestInit) => - Promise.resolve(createNoProgressOpenAIResponsesStream(getRequestSignal(input, init)))) as typeof fetch; + const fetchMock: FetchImpl = (input: string | URL | Request, init?: RequestInit) => + Promise.resolve(createNoProgressOpenAIResponsesStream(getRequestSignal(input, init))); const controller = new AbortController(); const abortTimer = setTimeout(() => controller.abort(new Error("fallback abort")), 200); abortTimer.unref(); @@ -550,6 +563,7 @@ describe("OpenAI-family first-event timeouts", () => { signal: controller.signal, streamFirstEventTimeoutMs: 1_000, streamIdleTimeoutMs: 20, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("error"); @@ -561,11 +575,12 @@ describe("OpenAI-family first-event timeouts", () => { it("keeps caller aborts as aborted for OpenAI responses", async () => { await expectCallerAbort( - (signal, streamFirstEventTimeoutMs) => + (signal, streamFirstEventTimeoutMs, fetchMock) => streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", signal, streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI responses stream timed out while waiting for the first event", ); @@ -573,11 +588,12 @@ describe("OpenAI-family first-event timeouts", () => { it("keeps caller aborts as aborted for OpenAI completions", async () => { await expectCallerAbort( - (signal, streamFirstEventTimeoutMs) => + (signal, streamFirstEventTimeoutMs, fetchMock) => streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", signal, streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "OpenAI completions stream timed out while waiting for the first event", ); @@ -585,13 +601,14 @@ describe("OpenAI-family first-event timeouts", () => { it("keeps caller aborts as aborted for Azure OpenAI responses", async () => { await expectCallerAbort( - (signal, streamFirstEventTimeoutMs) => + (signal, streamFirstEventTimeoutMs, fetchMock) => streamAzureOpenAIResponses(azureOpenAIResponsesModel, baseContext(), { apiKey: "test-key", azureBaseUrl: azureOpenAIResponsesModel.baseUrl, azureApiVersion: "v1", signal, streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), "Azure OpenAI responses stream timed out while waiting for the first event", ); @@ -599,10 +616,11 @@ describe("OpenAI-family first-event timeouts", () => { it("does not arm the first-event watchdog before OpenAI responses stream setup finishes", async () => { await expectDelayedRequestSetupSucceeds( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), createOpenAIResponsesSuccessResponse, ); @@ -610,10 +628,11 @@ describe("OpenAI-family first-event timeouts", () => { it("does not arm the first-event watchdog before OpenAI completions stream setup finishes", async () => { await expectDelayedRequestSetupSucceeds( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), () => createOpenAICompletionsSuccessResponse(openAICompletionsModel.id), ); @@ -621,12 +640,13 @@ describe("OpenAI-family first-event timeouts", () => { it("does not arm the first-event watchdog before Azure OpenAI responses setup finishes", async () => { await expectDelayedRequestSetupSucceeds( - streamFirstEventTimeoutMs => + (streamFirstEventTimeoutMs, fetchMock) => streamAzureOpenAIResponses(azureOpenAIResponsesModel, baseContext(), { apiKey: "test-key", azureBaseUrl: azureOpenAIResponsesModel.baseUrl, azureApiVersion: "v1", streamFirstEventTimeoutMs, + fetch: fetchMock, }).result(), createOpenAIResponsesSuccessResponse, ); diff --git a/packages/ai/test/openai-responses-cache-affinity.test.ts b/packages/ai/test/openai-responses-cache-affinity.test.ts index 8cd6d2fc7..68c8fab50 100644 --- a/packages/ai/test/openai-responses-cache-affinity.test.ts +++ b/packages/ai/test/openai-responses-cache-affinity.test.ts @@ -1,9 +1,8 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { type OpenAIResponsesOptions, streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; -const originalFetch = global.fetch; const model = getBundledModel("openai", "gpt-5-mini") as Model<"openai-responses">; function createSseResponse(events: unknown[]): Response { @@ -26,7 +25,7 @@ async function captureOpenAIResponseHeaders( clientRequestId: null as string | null, body: null as Record | null, }; - const fetchMock = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { captured.sessionId = getHeader(init?.headers, "session_id"); captured.clientRequestId = getHeader(init?.headers, "x-client-request-id"); captured.body = typeof init?.body === "string" ? (JSON.parse(init.body) as Record) : null; @@ -61,13 +60,12 @@ async function captureOpenAIResponseHeaders( }, ]); }); - global.fetch = Object.assign(fetchMock, { preconnect: originalFetch.preconnect }) as typeof fetch; const context: Context = { systemPrompt: ["stable system", "stable durable context"], messages: [{ role: "user", content: "hi", timestamp: Date.now() }], }; - const stream = streamOpenAIResponses(model, context, { apiKey: "test-key", ...options }); + const stream = streamOpenAIResponses(model, context, { apiKey: "test-key", ...options, fetch: fetchMock }); for await (const event of stream) { if (event.type === "done" || event.type === "error") break; @@ -77,7 +75,6 @@ async function captureOpenAIResponseHeaders( } afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); diff --git a/packages/ai/test/openai-responses-omit-max-output-tokens.test.ts b/packages/ai/test/openai-responses-omit-max-output-tokens.test.ts index 423ffa78d..36432158f 100644 --- a/packages/ai/test/openai-responses-omit-max-output-tokens.test.ts +++ b/packages/ai/test/openai-responses-omit-max-output-tokens.test.ts @@ -1,15 +1,13 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamSimple } from "@oh-my-pi/pi-ai/stream"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const baseModel = getBundledModel("openai", "gpt-4o-mini") as Model<"openai-responses">; -function mockSseFetch(): Record { +function mockSseFetch(): { fetchMock: FetchImpl; captured: Record } { const captured: Record = {}; - const fetchMock = vi.fn(async (_url: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_url: string | URL | Request, init?: RequestInit) => { const body = typeof init?.body === "string" ? (JSON.parse(init.body) as Record) : {}; Object.assign(captured, body); const event = { @@ -29,8 +27,7 @@ function mockSseFetch(): Record { headers: { "content-type": "text/event-stream" }, }); }); - global.fetch = Object.assign(fetchMock, { preconnect: originalFetch.preconnect }) as typeof fetch; - return captured; + return { fetchMock, captured }; } const ctx: Context = { @@ -39,8 +36,8 @@ const ctx: Context = { }; async function drain(model: Model<"openai-responses">): Promise> { - const captured = mockSseFetch(); - const stream = streamSimple(model, ctx, { apiKey: "k" }); + const { fetchMock, captured } = mockSseFetch(); + const stream = streamSimple(model, ctx, { apiKey: "k", fetch: fetchMock }); for await (const event of stream) { if (event.type === "done" || event.type === "error") break; } @@ -52,7 +49,6 @@ beforeEach(() => { }); afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); diff --git a/packages/ai/test/openai-responses-system-prompt.test.ts b/packages/ai/test/openai-responses-system-prompt.test.ts index 970d6567f..329060eb6 100644 --- a/packages/ai/test/openai-responses-system-prompt.test.ts +++ b/packages/ai/test/openai-responses-system-prompt.test.ts @@ -1,9 +1,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; // Non-reasoning model on api.openai.com (canonical path) const gpt4oMiniModel = getBundledModel("openai", "gpt-4o-mini") as Model<"openai-responses">; @@ -45,13 +43,12 @@ async function captureRequestBody( context: Context, ): Promise> { let captured: Record = {}; - const fetchMock = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { captured = typeof init?.body === "string" ? (JSON.parse(init.body) as Record) : {}; return createSseResponse(); }); - global.fetch = Object.assign(fetchMock, { preconnect: originalFetch.preconnect }) as typeof fetch; - const stream = streamOpenAIResponses(model, context, { apiKey: "test-key" }); + const stream = streamOpenAIResponses(model, context, { apiKey: "test-key", fetch: fetchMock }); for await (const event of stream) { if (event.type === "done" || event.type === "error") break; } @@ -59,7 +56,6 @@ async function captureRequestBody( } afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); diff --git a/packages/ai/test/openai-tool-strict-mode.test.ts b/packages/ai/test/openai-tool-strict-mode.test.ts index 0d70ad3c5..ee48ab8c0 100644 --- a/packages/ai/test/openai-tool-strict-mode.test.ts +++ b/packages/ai/test/openai-tool-strict-mode.test.ts @@ -1,17 +1,10 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; -import type { Context, Model, OpenAICompat, ProviderSessionState, Tool } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, OpenAICompat, ProviderSessionState, Tool } from "@oh-my-pi/pi-ai/types"; import * as z from "zod/v4"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); - const testTool: Tool = { name: "echo", description: "Echo input", @@ -213,7 +206,7 @@ describe("OpenAI tool strict mode", () => { ...(getBundledModel("openai", "gpt-4o-mini") as Model<"openai-completions">), api: "openai-completions", }; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, _init?: RequestInit): Promise => new Response( JSON.stringify({ @@ -227,10 +220,13 @@ describe("OpenAI tool strict mode", () => { headers: { "content-type": "application/json" }, }, ), - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); - const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, testContext, { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("error"); expect(result.errorMessage).toContain("Tools with mixed values for 'strict' are not allowed."); expect(result.errorMessage).toContain("param=tools"); @@ -244,7 +240,7 @@ describe("OpenAI tool strict mode", () => { compat: { toolStrictMode: "all_strict" } satisfies OpenAICompat, }; const strictFlags: boolean[][] = []; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, init?: RequestInit): Promise => { const bodyText = typeof init?.body === "string" ? init.body : ""; const payload = JSON.parse(bodyText) as { @@ -283,10 +279,13 @@ describe("OpenAI tool strict mode", () => { "[DONE]", ]); }, - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); - const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, testContext, { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("stop"); expect(result.content).toContainEqual({ type: "text", text: "Hello" }); expect(strictFlags).toEqual([[true], [false]]); @@ -297,7 +296,7 @@ describe("OpenAI tool strict mode", () => { const providerSessionState = new Map(); const strictFlags: boolean[][] = []; let attempt = 0; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, init?: RequestInit): Promise => { attempt += 1; const bodyText = typeof init?.body === "string" ? init.body : ""; @@ -340,12 +339,13 @@ describe("OpenAI tool strict mode", () => { "[DONE]", ]); }, - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key", providerSessionState, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("stop"); @@ -356,6 +356,7 @@ describe("OpenAI tool strict mode", () => { const nextResult = await streamOpenAICompletions(model, testContext, { apiKey: "test-key", providerSessionState, + fetch: fetchMock, }).result(); expect(nextResult.stopReason).toBe("stop"); @@ -367,7 +368,7 @@ describe("OpenAI tool strict mode", () => { const model = getBundledModel("openrouter", "anthropic/claude-sonnet-4") as Model<"openai-completions">; const providerSessionState = new Map(); const strictFlags: boolean[][] = []; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, init?: RequestInit): Promise => { const bodyText = typeof init?.body === "string" ? init.body : ""; const payload = JSON.parse(bodyText) as { @@ -386,12 +387,13 @@ describe("OpenAI tool strict mode", () => { }, ); }, - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key", providerSessionState, + fetch: fetchMock, }).result(); expect(result.stopReason).toBe("error"); diff --git a/packages/ai/test/openrouter-login.test.ts b/packages/ai/test/openrouter-login.test.ts index f98dd21c2..9b051f56e 100644 --- a/packages/ai/test/openrouter-login.test.ts +++ b/packages/ai/test/openrouter-login.test.ts @@ -3,9 +3,9 @@ import { afterEach, describe, expect, test, vi } from "bun:test"; import { AuthStorage, SqliteAuthCredentialStore } from "@oh-my-pi/pi-ai/auth-storage"; import { getOAuthProviders } from "@oh-my-pi/pi-ai/registry/oauth"; import { getEnvApiKey } from "@oh-my-pi/pi-ai/stream"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const originalOpenRouterApiKey = Bun.env.OPENROUTER_API_KEY; -const originalFetch = global.fetch; afterEach(() => { if (originalOpenRouterApiKey === undefined) { @@ -13,7 +13,6 @@ afterEach(() => { } else { Bun.env.OPENROUTER_API_KEY = originalOpenRouterApiKey; } - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -32,7 +31,7 @@ describe("openrouter login wiring", () => { test("AuthStorage.login('openrouter') validates against /auth/key and stores the pasted key", async () => { const fetchCalls: Array<{ url: string; init: RequestInit | undefined }> = []; - global.fetch = (async (input: Parameters[0], init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : (input as Request).url; fetchCalls.push({ url, init }); @@ -43,7 +42,7 @@ describe("openrouter login wiring", () => { }); } throw new Error(`unexpected fetch: ${url}`); - }) as unknown as typeof fetch; + }); const store = new SqliteAuthCredentialStore(new Database(":memory:")); const storage = new AuthStorage(store); @@ -52,6 +51,7 @@ describe("openrouter login wiring", () => { await storage.login("openrouter", { onAuth: () => {}, onPrompt: async () => "sk-or-validated", + fetch: fetchMock, }); const credential = await storage.get("openrouter"); @@ -66,13 +66,13 @@ describe("openrouter login wiring", () => { }); test("AuthStorage.login('openrouter') rejects keys that fail /auth/key validation", async () => { - global.fetch = vi.fn( + const fetchMock: FetchImpl = vi.fn( async () => new Response("Unauthorized", { status: 401, headers: { "Content-Type": "text/plain" }, }), - ) as unknown as typeof fetch; + ); const store = new SqliteAuthCredentialStore(new Database(":memory:")); const storage = new AuthStorage(store); @@ -82,6 +82,7 @@ describe("openrouter login wiring", () => { storage.login("openrouter", { onAuth: () => {}, onPrompt: async () => "sk-or-bogus", + fetch: fetchMock, }), ).rejects.toThrow(/OpenRouter API key validation failed \(401\)/); diff --git a/packages/ai/test/provider-fetch-override.test.ts b/packages/ai/test/provider-fetch-override.test.ts index fdf174ed2..4443cab04 100644 --- a/packages/ai/test/provider-fetch-override.test.ts +++ b/packages/ai/test/provider-fetch-override.test.ts @@ -1,14 +1,8 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const openAIResponsesModel = getBundledModel("openai", "gpt-5-mini") as Model<"openai-responses">; const openAICompletionsModel = { @@ -30,39 +24,34 @@ function createSseResponse(events: unknown[]): Response { }); } -function rejectingGlobalFetch(): typeof fetch { - const reject = async (): Promise => { - throw new Error("global fetch must not be used when an override is provided"); - }; - return Object.assign(reject, { preconnect: originalFetch.preconnect }); -} - describe("StreamOptions.fetch override", () => { it("routes openai-completions requests through the override", async () => { const calls: Array<{ url: string }> = []; - global.fetch = rejectingGlobalFetch(); - const customFetch = async (input: string | URL | Request, _init?: RequestInit) => { - calls.push({ url: String(input instanceof Request ? input.url : input) }); - return createSseResponse([ - { - id: "chatcmpl-test", - object: "chat.completion.chunk", - created: 0, - model: openAICompletionsModel.id, - choices: [{ index: 0, delta: { content: "hi" } }], - }, - { - id: "chatcmpl-test", - object: "chat.completion.chunk", - created: 0, - model: openAICompletionsModel.id, - choices: [{ index: 0, delta: {}, finish_reason: "stop" }], - usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, - }, - "[DONE]", - ]); - }; + const customFetch: FetchImpl = Object.assign( + async (input: string | URL | Request, _init?: RequestInit) => { + calls.push({ url: String(input instanceof Request ? input.url : input) }); + return createSseResponse([ + { + id: "chatcmpl-test", + object: "chat.completion.chunk", + created: 0, + model: openAICompletionsModel.id, + choices: [{ index: 0, delta: { content: "hi" } }], + }, + { + id: "chatcmpl-test", + object: "chat.completion.chunk", + created: 0, + model: openAICompletionsModel.id, + choices: [{ index: 0, delta: {}, finish_reason: "stop" }], + usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, + }, + "[DONE]", + ]); + }, + { preconnect: fetch.preconnect }, + ); const result = await streamOpenAICompletions(openAICompletionsModel, baseContext(), { apiKey: "test-key", @@ -76,43 +65,45 @@ describe("StreamOptions.fetch override", () => { it("routes openai-responses requests through the override", async () => { const calls: Array<{ url: string }> = []; - global.fetch = rejectingGlobalFetch(); - const customFetch = async (input: string | URL | Request, _init?: RequestInit) => { - calls.push({ url: String(input instanceof Request ? input.url : input) }); - return createSseResponse([ - { type: "response.created", response: { id: "resp_test" } }, - { - type: "response.output_item.added", - item: { type: "message", id: "msg_test", role: "assistant", status: "in_progress", content: [] }, - }, - { type: "response.content_part.added", part: { type: "output_text", text: "" } }, - { type: "response.output_text.delta", delta: "hi" }, - { - type: "response.output_item.done", - item: { - type: "message", - id: "msg_test", - role: "assistant", - status: "completed", - content: [{ type: "output_text", text: "hi" }], + const customFetch: FetchImpl = Object.assign( + async (input: string | URL | Request, _init?: RequestInit) => { + calls.push({ url: String(input instanceof Request ? input.url : input) }); + return createSseResponse([ + { type: "response.created", response: { id: "resp_test" } }, + { + type: "response.output_item.added", + item: { type: "message", id: "msg_test", role: "assistant", status: "in_progress", content: [] }, }, - }, - { - type: "response.completed", - response: { - id: "resp_test", - status: "completed", - usage: { - input_tokens: 1, - output_tokens: 1, - total_tokens: 2, - input_tokens_details: { cached_tokens: 0 }, + { type: "response.content_part.added", part: { type: "output_text", text: "" } }, + { type: "response.output_text.delta", delta: "hi" }, + { + type: "response.output_item.done", + item: { + type: "message", + id: "msg_test", + role: "assistant", + status: "completed", + content: [{ type: "output_text", text: "hi" }], }, }, - }, - ]); - }; + { + type: "response.completed", + response: { + id: "resp_test", + status: "completed", + usage: { + input_tokens: 1, + output_tokens: 1, + total_tokens: 2, + input_tokens_details: { cached_tokens: 0 }, + }, + }, + }, + ]); + }, + { preconnect: fetch.preconnect }, + ); const result = await streamOpenAIResponses(openAIResponsesModel, baseContext(), { apiKey: "test-key", diff --git a/packages/ai/test/provider-response.test.ts b/packages/ai/test/provider-response.test.ts index 14d4571fe..896bdbc32 100644 --- a/packages/ai/test/provider-response.test.ts +++ b/packages/ai/test/provider-response.test.ts @@ -1,7 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamSimple } from "@oh-my-pi/pi-ai/stream"; -import type { Context, Model, ProviderResponseMetadata } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, ProviderResponseMetadata } from "@oh-my-pi/pi-ai/types"; import { normalizeProviderResponse, notifyProviderResponse } from "@oh-my-pi/pi-ai/utils/provider-response"; describe("provider response metadata", () => { @@ -54,12 +54,6 @@ describe("provider response metadata", () => { }); }); -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - function createSseResponse(events: unknown[], headers: Record = {}): Response { const payload = `${events.map(event => `data: ${typeof event === "string" ? event : JSON.stringify(event)}`).join("\n\n")}\n\n`; return new Response(payload, { @@ -75,7 +69,7 @@ describe("streamSimple onResponse propagation", () => { api: "openai-completions", }; - global.fetch = Object.assign( + const fetchMock: FetchImpl = Object.assign( async (_input: string | URL | Request, _init?: RequestInit): Promise => createSseResponse( [ @@ -97,13 +91,14 @@ describe("streamSimple onResponse propagation", () => { ], { "x-request-id": "req_stream_simple" }, ), - { preconnect: originalFetch.preconnect }, + { preconnect: fetch.preconnect }, ); const context: Context = { messages: [{ role: "user", content: "hello", timestamp: Date.now() }] }; const seen: ProviderResponseMetadata[] = []; const result = await streamSimple(model, context, { apiKey: "test-key", + fetch: fetchMock, onResponse: response => { seen.push(response); }, diff --git a/packages/ai/test/raw-sse-sdk-capture.test.ts b/packages/ai/test/raw-sse-sdk-capture.test.ts index 9df20c855..54b0f2bc8 100644 --- a/packages/ai/test/raw-sse-sdk-capture.test.ts +++ b/packages/ai/test/raw-sse-sdk-capture.test.ts @@ -6,9 +6,7 @@ import type { RawMessageStreamEvent } from "@oh-my-pi/pi-ai/providers/anthropic- import { streamAzureOpenAIResponses } from "@oh-my-pi/pi-ai/providers/azure-openai-responses"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { streamOpenAIResponses } from "@oh-my-pi/pi-ai/providers/openai-responses"; -import type { Context, Model, RawSseEvent } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; +import type { Context, FetchImpl, Model, RawSseEvent } from "@oh-my-pi/pi-ai/types"; const context: Context = { messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], @@ -116,10 +114,11 @@ function createSseResponse(events: unknown[]): Response { }); } -function installFetchResponse(events: unknown[]) { - const fetchMock = vi.fn(async () => createSseResponse(events)); - global.fetch = Object.assign(fetchMock, { preconnect: originalFetch.preconnect }) as typeof fetch; - return fetchMock; +function createFetchResponse(events: unknown[]): FetchImpl { + return Object.assign( + vi.fn(async () => createSseResponse(events)), + { preconnect: fetch.preconnect }, + ) as typeof fetch; } function recordEvent(events: RawSseEvent[]): (event: RawSseEvent) => void { @@ -168,17 +167,17 @@ function createAnthropicRawClient(events: RawMessageStreamEvent[]): AnthropicMes } afterEach(() => { - global.fetch = originalFetch; vi.restoreAllMocks(); }); describe("SDK raw SSE capture", () => { it("records OpenAI Responses SDK events from the decoded stream", async () => { - const fetchMock = installFetchResponse(openAIResponsesEvents); + const fetchMock = createFetchResponse(openAIResponsesEvents); const observed: RawSseEvent[] = []; const result = await streamOpenAIResponses(openAIResponsesModel, context, { apiKey: "test-key", + fetch: fetchMock, onSseEvent: recordEvent(observed), }).result(); @@ -216,11 +215,12 @@ describe("SDK raw SSE capture", () => { }, "[DONE]", ]; - installFetchResponse(chunks); + const fetchMock = createFetchResponse(chunks); const observed: RawSseEvent[] = []; const result = await streamOpenAICompletions(openAICompletionsModel, context, { apiKey: "test-key", + fetch: fetchMock, onSseEvent: recordEvent(observed), }).result(); @@ -231,11 +231,12 @@ describe("SDK raw SSE capture", () => { }); it("records Azure OpenAI Responses SDK events from the decoded stream", async () => { - installFetchResponse(openAIResponsesEvents); + const fetchMock = createFetchResponse(openAIResponsesEvents); const observed: RawSseEvent[] = []; const result = await streamAzureOpenAIResponses(azureOpenAIResponsesModel, context, { apiKey: "test-key", + fetch: fetchMock, azureBaseUrl: azureOpenAIResponsesModel.baseUrl, azureApiVersion: "v1", onSseEvent: recordEvent(observed), @@ -261,9 +262,12 @@ describe("SDK raw SSE capture", () => { }); it("does not synthesize raw SSE records when no observer is installed", async () => { - installFetchResponse(openAIResponsesEvents); + const fetchMock = createFetchResponse(openAIResponsesEvents); - const result = await streamOpenAIResponses(openAIResponsesModel, context, { apiKey: "test-key" }).result(); + const result = await streamOpenAIResponses(openAIResponsesModel, context, { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.stopReason).toBe("stop"); }); diff --git a/packages/ai/test/request-debug.test.ts b/packages/ai/test/request-debug.test.ts index 22bbf57e7..389812218 100644 --- a/packages/ai/test/request-debug.test.ts +++ b/packages/ai/test/request-debug.test.ts @@ -7,7 +7,6 @@ import { stream } from "@oh-my-pi/pi-ai/stream"; import type { AssistantMessage, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; import { wrapFetchForRequestDebug } from "@oh-my-pi/pi-ai/utils/request-debug"; -import { hookFetch } from "@oh-my-pi/pi-utils"; const enc = new TextEncoder(); @@ -144,9 +143,9 @@ describe("PI_REQ_DEBUG request/response recording", () => { expect(log.body).toEqual(firstChunk); }); - it("injects the debug fetch into provider options when callers did not pass fetch", async () => { + it("wraps provider fetch options with request debug recording", async () => { Bun.env.PI_REQ_DEBUG = "1"; - using _hook = hookFetch(() => new Response("ok", { headers: { "x-debug": "yes" } })); + const fetchMock: FetchImpl = async () => new Response("ok", { headers: { "x-debug": "yes" } }); registerCustomApi("req-debug-test", (_model, _context, options) => { const events = new AssistantMessageEventStream(); void (async () => { @@ -195,7 +194,7 @@ describe("PI_REQ_DEBUG request/response recording", () => { const events = stream( model, { messages: [{ role: "user", content: "hi", timestamp: Date.now() }] }, - { apiKey: "key" }, + { apiKey: "key", fetch: fetchMock }, ); await events.result(); diff --git a/packages/ai/test/stream-markup-healing.test.ts b/packages/ai/test/stream-markup-healing.test.ts index c7e6facac..9d48abe82 100644 --- a/packages/ai/test/stream-markup-healing.test.ts +++ b/packages/ai/test/stream-markup-healing.test.ts @@ -1,16 +1,10 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; import { stream } from "@oh-my-pi/pi-ai/stream"; -import type { Context, Model, Tool, ToolCall } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model, Tool, ToolCall } from "@oh-my-pi/pi-ai/types"; import { getStreamMarkupHealingPattern, StreamMarkupHealing } from "@oh-my-pi/pi-ai/utils/stream-markup-healing"; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); - interface SseToolCallDelta { index: number; id?: string; @@ -45,9 +39,9 @@ function sseResponse(events: ReadonlyArray): Response { }); } -function mockFetch(events: ReadonlyArray): typeof fetch { +function mockFetch(events: ReadonlyArray): FetchImpl { const fn = async (_input: string | URL | Request, _init?: RequestInit): Promise => sseResponse(events); - return Object.assign(fn, { preconnect: originalFetch.preconnect }); + return Object.assign(fn, { preconnect: fetch.preconnect }); } function baseContext(): Context { @@ -136,9 +130,9 @@ function ndjsonResponse(lines: ReadonlyArray): Response { }); } -function mockNdjsonFetch(lines: ReadonlyArray): typeof fetch { +function mockNdjsonFetch(lines: ReadonlyArray): FetchImpl { const fn = async (_input: string | URL | Request, _init?: RequestInit): Promise => ndjsonResponse(lines); - return Object.assign(fn, { preconnect: originalFetch.preconnect }); + return Object.assign(fn, { preconnect: fetch.preconnect }); } describe("StreamMarkupHealing pattern selection", () => { @@ -245,14 +239,14 @@ describe("Kimi K2 leaked markup healing", () => { "<|tool_call_end|>" + "<|tool_calls_section_end|>"; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "I'll read it. " }), chunk(model.id, { content: leaked }), chunk(model.id, {}, "stop"), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const text = result.content .filter(b => b.type === "text") @@ -286,14 +280,14 @@ describe("Kimi K2 leaked markup healing", () => { expect(a + b).toBe(full); expect(a.endsWith("<|tool_ca")).toBe(true); - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: a }), chunk(model.id, { content: b }), chunk(model.id, {}, "stop"), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const text = result.content .filter(b => b.type === "text") @@ -319,9 +313,9 @@ describe("Kimi K2 leaked markup healing", () => { "<|tool_call_end|>" + "<|tool_calls_section_end|>"; - global.fetch = mockFetch([chunk(model.id, { content: leaked }), chunk(model.id, {}, "stop"), "[DONE]"]); + const fetchMock = mockFetch([chunk(model.id, { content: leaked }), chunk(model.id, {}, "stop"), "[DONE]"]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const toolCalls = result.content.filter((b): b is ToolCall => b.type === "toolCall"); expect(toolCalls).toHaveLength(2); @@ -336,7 +330,7 @@ describe("Kimi K2 leaked markup healing", () => { const tail = "<|tool_call_end|><|tool_calls_section_end|>"; const argsParts = ['{"path":"', "out.txt", '","content":"', "hello world", '"}']; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: head }), ...argsParts.map(part => chunk(model.id, { content: part })), chunk(model.id, { content: tail }), @@ -344,7 +338,7 @@ describe("Kimi K2 leaked markup healing", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const toolCalls = result.content.filter((b): b is ToolCall => b.type === "toolCall"); expect(toolCalls).toHaveLength(1); @@ -353,14 +347,14 @@ describe("Kimi K2 leaked markup healing", () => { }); it("passes prose through unchanged when no markers are present", async () => { - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "Hello, " }), chunk(model.id, { content: "world!" }), chunk(model.id, {}, "stop"), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const text = result.content .filter(b => b.type === "text") .map(b => b.text) @@ -373,13 +367,13 @@ describe("Kimi K2 leaked markup healing", () => { it("emits a literal '<|' that is not a token prefix without holding it back forever", async () => { // `<|hello|>` is not any known token. It should land in visible text. - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "before <|hello|> after" }), chunk(model.id, {}, "stop"), "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const text = result.content .filter(b => b.type === "text") .map(b => b.text) @@ -399,9 +393,13 @@ describe("Kimi K2 leaked markup healing", () => { "<|tool_call_end|>" + "<|tool_calls_section_end|>"; - global.fetch = mockFetch([chunk(model.id, { content: leaked }), chunk(model.id, {}, "content_filter"), "[DONE]"]); + const fetchMock = mockFetch([ + chunk(model.id, { content: leaked }), + chunk(model.id, {}, "content_filter"), + "[DONE]", + ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); expect(result.stopReason).toBe("error"); expect(result.errorMessage).toContain("content_filter"); }); @@ -417,7 +415,7 @@ describe("Kimi K2 leaked markup healing", () => { "<|tool_call_end|>" + "<|tool_calls_section_end|>"; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: leaked, tool_calls: [ @@ -433,7 +431,7 @@ describe("Kimi K2 leaked markup healing", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const toolCalls = result.content.filter((b): b is ToolCall => b.type === "toolCall"); expect(toolCalls).toHaveLength(1); expect(toolCalls[0].id).toBe("call_structured_abc"); @@ -458,7 +456,7 @@ describe("Kimi K2 leaked markup healing", () => { "<|tool_call_end|>" + "<|tool_calls_section_end|>"; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { tool_calls: [ { @@ -474,7 +472,7 @@ describe("Kimi K2 leaked markup healing", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const toolCalls = result.content.filter((b): b is ToolCall => b.type === "toolCall"); expect(toolCalls).toHaveLength(2); expect(toolCalls.map(call => call.name)).toEqual(["read", "write"]); @@ -485,9 +483,9 @@ describe("Kimi K2 leaked markup healing", () => { it("passes a literal <|tool_call_end|> through as text when no section is active", async () => { const prose = "Use <|tool_call_end|> to close a call."; - global.fetch = mockFetch([chunk(model.id, { content: prose }), chunk(model.id, {}, "stop"), "[DONE]"]); + const fetchMock = mockFetch([chunk(model.id, { content: prose }), chunk(model.id, {}, "stop"), "[DONE]"]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test", fetch: fetchMock }).result(); const text = result.content .filter(b => b.type === "text") .map(b => b.text) @@ -500,7 +498,7 @@ describe("Kimi K2 leaked markup healing", () => { describe("Ollama provider DSML envelope healing", () => { it("emits a healed tool call, suppresses leaked text, and promotes stop", async () => { - global.fetch = mockNdjsonFetch([ + const fetchMock = mockNdjsonFetch([ { model: "deepseek-v4-pro", message: { role: "assistant", content: " 精神精神\n\n" }, @@ -523,7 +521,7 @@ describe("Ollama provider DSML envelope healing", () => { const result = await stream( deepseekCloudModel, { messages: [{ role: "user", content: "Check Fedora packages", timestamp: Date.now() }] }, - { apiKey: "test-key" }, + { apiKey: "test-key", fetch: fetchMock }, ).result(); const visibleText = result.content @@ -553,7 +551,7 @@ describe("Ollama provider DSML envelope healing", () => { }); it("leaves non-DeepSeek Ollama content untouched", async () => { - global.fetch = mockNdjsonFetch([ + const fetchMock = mockNdjsonFetch([ { model: "gpt-oss:120b", message: { role: "assistant", content: "Inline `<|literal|>` token in prose." }, @@ -571,7 +569,7 @@ describe("Ollama provider DSML envelope healing", () => { const result = await stream( { ...deepseekCloudModel, id: "gpt-oss:120b", name: "GPT OSS 120B" }, { messages: [{ role: "user", content: "hi", timestamp: Date.now() }] }, - { apiKey: "test-key" }, + { apiKey: "test-key", fetch: fetchMock }, ).result(); const text = result.content @@ -597,7 +595,7 @@ describe("OpenAI completions MiniMax thinking healing", () => { contextWindow: 200_000, maxTokens: 8_192, }; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "visible hidden reasoning" }), @@ -606,7 +604,10 @@ describe("OpenAI completions MiniMax thinking healing", () => { "[DONE]", ]); - const result = await streamOpenAICompletions(model, baseContext(), { apiKey: "test-key" }).result(); + const result = await streamOpenAICompletions(model, baseContext(), { + apiKey: "test-key", + fetch: fetchMock, + }).result(); expect(result.content).toEqual([ { type: "text", text: "visible " }, @@ -630,7 +631,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { contextWindow: 131_072, maxTokens: 8_192, }; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "I'll check.\n" }), chunk(model.id, { content: `${REPORTED_DSML_LEAK}\nThat should give us the package list.` }), chunk(model.id, {}, "stop"), @@ -640,7 +641,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { const result = await streamOpenAICompletions( model, { messages: [{ role: "user", content: "Check Fedora", timestamp: Date.now() }] }, - { apiKey: "test-key" }, + { apiKey: "test-key", fetch: fetchMock }, ).result(); const text = result.content @@ -674,7 +675,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { expect(model.provider).toBe("nanogpt"); let payload: Record | undefined; - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { content: "Checking.\n" }), chunk(model.id, { content: REPORTED_DSML_LEAK }), chunk(model.id, {}, "stop"), @@ -687,6 +688,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { { apiKey: "test-key", reasoning: "high", + fetch: fetchMock, onPayload: value => { payload = value as Record; }, @@ -718,7 +720,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { it("keeps indexed parallel NanoGPT read deltas attached to their own tool calls", async () => { const model = getBundledModel<"openai-completions">("nanogpt", "deepseek/deepseek-v4-pro"); - global.fetch = mockFetch([ + const fetchMock = mockFetch([ chunk(model.id, { tool_calls: [ { index: 0, id: "call_a", type: "function", function: { name: "read", arguments: "" } }, @@ -738,7 +740,7 @@ describe("OpenAI completions provider DSML envelope healing", () => { const result = await streamOpenAICompletions( model, { messages: [{ role: "user", content: "Read a.ts and b.ts", timestamp: Date.now() }], tools: [readTool] }, - { apiKey: "test-key", reasoning: "high" }, + { apiKey: "test-key", reasoning: "high", fetch: fetchMock }, ).result(); const toolCalls = result.content.filter((b): b is ToolCall => b.type === "toolCall"); diff --git a/packages/ai/test/synthetic-login.test.ts b/packages/ai/test/synthetic-login.test.ts index 2820360dd..1fe260395 100644 --- a/packages/ai/test/synthetic-login.test.ts +++ b/packages/ai/test/synthetic-login.test.ts @@ -1,16 +1,10 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { loginSynthetic } from "@oh-my-pi/pi-ai/registry/synthetic"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("synthetic login", () => { it("validates API keys against the models endpoint instead of a deprecated model", async () => { - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { expect(String(input)).toBe("https://api.synthetic.new/openai/v1/models"); expect(init?.method).toBe("GET"); expect(init?.headers).toEqual({ Authorization: "Bearer sk-synthetic-test" }); @@ -19,10 +13,10 @@ describe("synthetic login", () => { headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; const apiKey = await loginSynthetic({ onPrompt: async () => "sk-synthetic-test", + fetch: fetchMock, }); expect(apiKey).toBe("sk-synthetic-test"); diff --git a/packages/ai/test/wafer.live.ts b/packages/ai/test/wafer.live.ts index 33f7fd323..059e77f16 100644 --- a/packages/ai/test/wafer.live.ts +++ b/packages/ai/test/wafer.live.ts @@ -9,7 +9,7 @@ */ import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; const apiKey = process.env.WAFER_PASS_API_KEY ?? process.env.WAFER_SERVERLESS_API_KEY; if (!apiKey) { @@ -27,21 +27,21 @@ interface CapturedRequest { body: string | null; } -const originalFetch = global.fetch; +const originalFetch = fetch; const captured: { value: CapturedRequest | null } = { value: null }; -type FetchInput = Parameters[0]; -global.fetch = (async (input: FetchInput, init?: RequestInit) => { +type FetchInput = string | URL | Request; +const fetchImpl: FetchImpl = async (input: FetchInput, init?: RequestInit) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; captured.value = { url, body: typeof init?.body === "string" ? init.body : null }; - return originalFetch(input as Parameters[0], init); -}) as typeof global.fetch; + return originalFetch(input, init); +}; const context: Context = { systemPrompt: ["Reply with exactly two words."], messages: [{ role: "user", content: "Say hi.", timestamp: Date.now() }], }; -const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey }); +const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey, fetch: fetchImpl }); let text = ""; let stopReason: string | undefined; let cost = 0; diff --git a/packages/ai/test/wafer.test.ts b/packages/ai/test/wafer.test.ts index a7deb0100..9abfad5b7 100644 --- a/packages/ai/test/wafer.test.ts +++ b/packages/ai/test/wafer.test.ts @@ -10,7 +10,7 @@ * the wire id (no rewrite). These tests defend the bundled catalog contract and * the case-sensitive id pass-through against the wire. */ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { createModelManager } from "@oh-my-pi/pi-ai/model-manager"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { @@ -18,13 +18,7 @@ import { waferServerlessModelManagerOptions, } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { streamOpenAICompletions } from "@oh-my-pi/pi-ai/providers/openai-completions"; -import type { Context, Model } from "@oh-my-pi/pi-ai/types"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); +import type { Context, FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; function sseResponse(events: unknown[]): Response { const payload = `${events.map(e => `data: ${typeof e === "string" ? e : JSON.stringify(e)}`).join("\n\n")}\n\n`; @@ -61,11 +55,11 @@ describe("Wafer Pass provider", () => { it("preserves the catalog id verbatim on the wire (no rewrite, case-sensitive)", async () => { const model = getBundledModel<"openai-completions">("wafer-pass", "GLM-5.1"); const captured: { url: string | null; body: string | null } = { url: null, body: null }; - global.fetch = (async (input: unknown, init?: RequestInit) => { - captured.url = typeof input === "string" ? input : input instanceof URL ? input.toString() : String(input); + const fetchMock: FetchImpl = async (input: string | URL | Request, init?: RequestInit) => { + captured.url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; captured.body = typeof init?.body === "string" ? init.body : null; return sseResponse(["[DONE]"]); - }) as typeof global.fetch; + }; const context: Context = { systemPrompt: ["t"], @@ -73,6 +67,7 @@ describe("Wafer Pass provider", () => { }; const stream = streamOpenAICompletions(model as Model<"openai-completions">, context, { apiKey: "wfr_test", + fetch: fetchMock, }); for await (const _event of stream) { /* drain */ @@ -165,12 +160,12 @@ describe("Wafer dynamic discovery mapper", () => { // - leave `thinkingFormat` unset for deepseek (uses `reasoning_effort`) // and unknown upstreams, so `detectOpenAICompat` picks the safe default // from the id pattern at request time. - function mockWaferModelsResponse(entries: Array>): void { - global.fetch = (async () => + function mockWaferModelsResponse(entries: Array>): FetchImpl { + return async () => new Response(JSON.stringify({ object: "list", data: entries }), { status: 200, headers: { "content-type": "application/json" }, - })) as unknown as typeof global.fetch; + }); } function makeWaferEntry(id: string, upstream: string, opts: { reasoning?: boolean; vision?: boolean } = {}) { @@ -199,7 +194,7 @@ describe("Wafer dynamic discovery mapper", () => { } it("picks thinkingFormat from the wafer.provider envelope per upstream", async () => { - mockWaferModelsResponse([ + const fetchMock = mockWaferModelsResponse([ makeWaferEntry("GLM-fake", "zai", { reasoning: true }), makeWaferEntry("Kimi-fake", "moonshotai", { reasoning: true }), makeWaferEntry("qwen-fake", "qwen", { reasoning: true }), @@ -208,7 +203,7 @@ describe("Wafer dynamic discovery mapper", () => { makeWaferEntry("nothink-fake", "zai", { reasoning: false }), ]); - const manager = createModelManager(waferServerlessModelManagerOptions({ apiKey: "wfr_test" })); + const manager = createModelManager(waferServerlessModelManagerOptions({ apiKey: "wfr_test", fetch: fetchMock })); const { models } = await manager.refresh("online"); const byId = new Map(models.map(m => [m.id, m as Model<"openai-completions">])); @@ -253,15 +248,17 @@ describe("Wafer dynamic discovery mapper", () => { }, }; - mockWaferModelsResponse([sharedEntry]); - const passManager = createModelManager(waferPassModelManagerOptions({ apiKey: "wfr_test" })); + const fetchMock = mockWaferModelsResponse([sharedEntry]); + const passManager = createModelManager(waferPassModelManagerOptions({ apiKey: "wfr_test", fetch: fetchMock })); const passResult = await passManager.refresh("online"); const passModel = passResult.models.find(m => m.id === "Shared-fake"); expect(passModel).toBeDefined(); expect(passModel?.cost).toEqual({ input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }); - mockWaferModelsResponse([sharedEntry]); - const srvManager = createModelManager(waferServerlessModelManagerOptions({ apiKey: "wfr_test" })); + const srvFetchMock = mockWaferModelsResponse([sharedEntry]); + const srvManager = createModelManager( + waferServerlessModelManagerOptions({ apiKey: "wfr_test", fetch: srvFetchMock }), + ); const srvResult = await srvManager.refresh("online"); const srvModel = srvResult.models.find(m => m.id === "Shared-fake"); expect(srvModel).toBeDefined(); diff --git a/packages/ai/test/xiaomi-oauth.test.ts b/packages/ai/test/xiaomi-oauth.test.ts index 11ffcbb21..2e7c8329a 100644 --- a/packages/ai/test/xiaomi-oauth.test.ts +++ b/packages/ai/test/xiaomi-oauth.test.ts @@ -1,17 +1,11 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { loginXiaomi } from "@oh-my-pi/pi-ai/registry/oauth/xiaomi"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("xiaomi oauth validation", () => { it("uses a fresh AbortSignal per endpoint so SGP timeout doesn't abort AMS fallback", async () => { const capturedSignals: (AbortSignal | undefined)[] = []; - const fetchMock = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { capturedSignals.push(init?.signal ?? undefined); if (capturedSignals.length === 1) { // Simulate SGP timing out: throw an AbortError as AbortSignal.timeout would. @@ -19,11 +13,11 @@ describe("xiaomi oauth validation", () => { } return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; await loginXiaomi({ onPrompt: async () => "tp-test-key", onAuth: () => {}, + fetch: fetchMock, }); expect(fetchMock).toHaveBeenCalledTimes(2); @@ -41,15 +35,15 @@ describe("xiaomi oauth validation", () => { // OpenAI-compatible endpoint requires Bearer auth and rejects x-api-key as 401 // "Invalid API Key" — see issue #1580. const capturedHeaders: Record[] = []; - const fetchMock = vi.fn(async (_input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (_input, init) => { capturedHeaders.push((init?.headers ?? {}) as Record); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; await loginXiaomi({ onPrompt: async () => "tp-test-key", onAuth: () => {}, + fetch: fetchMock, }); expect(fetchMock).toHaveBeenCalledTimes(1); @@ -61,16 +55,16 @@ describe("xiaomi oauth validation", () => { it("sends Authorization: Bearer for standard sk- keys as well", async () => { const capturedHeaders: Record[] = []; const capturedUrls: string[] = []; - const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = vi.fn(async (input, init) => { capturedUrls.push(typeof input === "string" ? input : input.toString()); capturedHeaders.push((init?.headers ?? {}) as Record); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); }); - global.fetch = fetchMock as unknown as typeof fetch; await loginXiaomi({ onPrompt: async () => "sk-test-key", onAuth: () => {}, + fetch: fetchMock, }); expect(fetchMock).toHaveBeenCalledTimes(1); diff --git a/packages/ai/test/xiaomi-tp-login-integration.test.ts b/packages/ai/test/xiaomi-tp-login-integration.test.ts index fe1fdb6a0..dac9fc45b 100644 --- a/packages/ai/test/xiaomi-tp-login-integration.test.ts +++ b/packages/ai/test/xiaomi-tp-login-integration.test.ts @@ -16,7 +16,7 @@ import { describe, expect, it } from "bun:test"; import { xiaomiModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { loginXiaomi } from "@oh-my-pi/pi-ai/registry/oauth/xiaomi"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; // Realistic tp- key (same format as user's key, but a dummy value for testing) const TP_KEY = "tp-ci1p8t1w4e1sbxgyc8v65tnrjbzro287igmvyf25van9mt76"; @@ -35,19 +35,20 @@ describe("loginXiaomi with tp- key", () => { it("validates against SGP token-plan host with Bearer auth and mimo-v2.5 model", async () => { const seen: { url: string; headers: Record; body: string }[] = []; - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { seen.push({ url: String(input), headers: (init?.headers ?? {}) as Record, body: init?.body as string, }); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; await loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: fetchMock, }); expect(seen).toHaveLength(1); @@ -70,19 +71,20 @@ describe("loginXiaomi with tp- key", () => { it("falls back SGP → AMS → CN during validation", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); seen.push(url); if (url.includes(TOKEN_PLAN_HOSTS.sgp) || url.includes(TOKEN_PLAN_HOSTS.ams)) { return new Response("Invalid API Key", { status: 401 }); } return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; await loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: fetchMock, }); // Tried SGP → AMS → succeeded on CN @@ -93,15 +95,16 @@ describe("loginXiaomi with tp- key", () => { }); it("throws when all three token-plan hosts return 401", async () => { - using _hook = hookFetch(_input => { + const fetchMock: FetchImpl = async () => { return new Response("Invalid API Key", { status: 401 }); - }); + }; await expect( loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: fetchMock, }), ).rejects.toThrow("Xiaomi MiMo API key validation failed (401)"); }); @@ -109,7 +112,7 @@ describe("loginXiaomi with tp- key", () => { it("falls back through timeouts: SGP timeout → AMS timeout → CN success", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); seen.push(url); if (url.includes(TOKEN_PLAN_HOSTS.sgp) || url.includes(TOKEN_PLAN_HOSTS.ams)) { @@ -117,12 +120,13 @@ describe("loginXiaomi with tp- key", () => { throw new DOMException("The operation was aborted due to timeout.", "AbortError"); } return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; await loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: fetchMock, }); expect(seen).toHaveLength(3); @@ -134,15 +138,16 @@ describe("loginXiaomi with tp- key", () => { it("does NOT hit the standard api.xiaomimimo.com for tp- keys", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; await loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: fetchMock, }); for (const url of seen) { @@ -157,15 +162,15 @@ describe("xiaomiModelManagerOptions with tp- key", () => { it("discovers models from SGP first", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response(JSON.stringify({ data: [{ id: "mimo-v2.5" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY }); + const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY, fetch: fetchMock }); const models = await opts.fetchDynamicModels?.(); expect(seen).toHaveLength(1); @@ -177,7 +182,7 @@ describe("xiaomiModelManagerOptions with tp- key", () => { it("falls back SGP → AMS → CN during discovery", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); seen.push(url); @@ -189,9 +194,9 @@ describe("xiaomiModelManagerOptions with tp- key", () => { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY }); + const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY, fetch: fetchMock }); const models = await opts.fetchDynamicModels?.(); // All three token-plan hosts tried in order @@ -203,11 +208,11 @@ describe("xiaomiModelManagerOptions with tp- key", () => { }); it("returns null when all token-plan hosts fail", async () => { - using _hook = hookFetch(() => { + const fetchMock: FetchImpl = async () => { return new Response("error", { status: 500 }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY }); + const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY, fetch: fetchMock }); const models = await opts.fetchDynamicModels?.(); expect(models).toBeNull(); @@ -216,15 +221,15 @@ describe("xiaomiModelManagerOptions with tp- key", () => { it("does NOT use standard host for tp- key model discovery", async () => { const seen: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { seen.push(String(input)); return new Response(JSON.stringify({ data: [] }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY }); + const opts = xiaomiModelManagerOptions({ apiKey: TP_KEY, fetch: fetchMock }); await opts.fetchDynamicModels?.(); for (const url of seen) { @@ -240,15 +245,16 @@ describe("Xiaomi tp- full round-trip", () => { // Phase 1: Login const loginUrls: string[] = []; - using _hook1 = hookFetch(input => { + const loginFetchMock: FetchImpl = async input => { loginUrls.push(String(input)); return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - }); + }; const returnedKey = await loginXiaomi({ onPrompt: async () => TP_KEY, onAuth: () => {}, onProgress: () => {}, + fetch: loginFetchMock, }); expect(returnedKey).toBe(TP_KEY); @@ -256,20 +262,18 @@ describe("Xiaomi tp- full round-trip", () => { expect(loginUrls[0]).toContain(TOKEN_PLAN_HOSTS.sgp); expect(loginUrls[0]).toContain("/v1/chat/completions"); - // Dispose hook1 (restore original fetch) - // Phase 2: Model discovery with the returned key const discoveryUrls: string[] = []; - using _hook2 = hookFetch(input => { + const discoveryFetchMock: FetchImpl = async input => { discoveryUrls.push(String(input)); return new Response(JSON.stringify({ data: [{ id: "mimo-v2.5" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const opts = xiaomiModelManagerOptions({ apiKey: returnedKey }); + const opts = xiaomiModelManagerOptions({ apiKey: returnedKey, fetch: discoveryFetchMock }); const models = await opts.fetchDynamicModels?.(); expect(discoveryUrls).toHaveLength(1); diff --git a/packages/ai/test/zenmux-login.test.ts b/packages/ai/test/zenmux-login.test.ts index 7186c3b05..4f021ce3d 100644 --- a/packages/ai/test/zenmux-login.test.ts +++ b/packages/ai/test/zenmux-login.test.ts @@ -1,12 +1,6 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { describe, expect, it, vi } from "bun:test"; import { loginZenMux } from "@oh-my-pi/pi-ai/registry/zenmux"; - -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; - vi.restoreAllMocks(); -}); +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; describe("zenmux login", () => { it("opens ZenMux key settings and validates against models endpoint", async () => { @@ -15,8 +9,8 @@ describe("zenmux login", () => { let promptMessage: string | undefined; let promptPlaceholder: string | undefined; - const fetchMock = vi.fn(async (input: string | URL, init?: RequestInit) => { - const url = typeof input === "string" ? input : input.toString(); + const fetchMock: FetchImpl = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; expect(url).toBe("https://zenmux.ai/api/v1/models"); expect(init?.method).toBe("GET"); expect(init?.headers).toEqual({ Authorization: "Bearer sk-zenmux-test" }); @@ -25,7 +19,6 @@ describe("zenmux login", () => { headers: { "Content-Type": "application/json" }, }); }); - global.fetch = fetchMock as unknown as typeof fetch; const apiKey = await loginZenMux({ onAuth: info => { @@ -37,6 +30,7 @@ describe("zenmux login", () => { promptPlaceholder = prompt.placeholder; return "sk-zenmux-test"; }, + fetch: fetchMock, }); expect(authUrl).toBe("https://zenmux.ai/settings/keys"); @@ -60,13 +54,14 @@ describe("zenmux login", () => { }); it("surfaces models endpoint validation errors", async () => { - global.fetch = vi.fn( + const fetchMock: FetchImpl = vi.fn( async () => new Response('{"error":"invalid_api_key"}', { status: 401 }), ) as unknown as typeof fetch; await expect( loginZenMux({ onPrompt: async () => "sk-zenmux-test", + fetch: fetchMock, }), ).rejects.toThrow("ZenMux API key validation failed (401)"); }); diff --git a/packages/ai/test/zenmux-provider.test.ts b/packages/ai/test/zenmux-provider.test.ts index e11bf9af5..73b38d99f 100644 --- a/packages/ai/test/zenmux-provider.test.ts +++ b/packages/ai/test/zenmux-provider.test.ts @@ -3,9 +3,9 @@ import { DEFAULT_MODEL_PER_PROVIDER, PROVIDER_DESCRIPTORS } from "@oh-my-pi/pi-a import { zenmuxModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { getOAuthProviders } from "@oh-my-pi/pi-ai/registry/oauth"; import { getEnvApiKey } from "@oh-my-pi/pi-ai/stream"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const originalZenMuxApiKey = Bun.env.ZENMUX_API_KEY; -const originalFetch = global.fetch; afterEach(() => { if (originalZenMuxApiKey === undefined) { @@ -13,7 +13,6 @@ afterEach(() => { } else { Bun.env.ZENMUX_API_KEY = originalZenMuxApiKey; } - global.fetch = originalFetch; vi.restoreAllMocks(); }); @@ -36,7 +35,7 @@ describe("zenmux provider support", () => { expect(provider?.name).toBe("ZenMux"); }); test("routes Anthropic-owned models to anthropic-messages", async () => { - global.fetch = vi.fn( + const fetchMock: FetchImpl = vi.fn( async () => new Response( JSON.stringify({ @@ -73,13 +72,13 @@ describe("zenmux provider support", () => { ), ) as unknown as typeof fetch; - const options = zenmuxModelManagerOptions({ apiKey: "zenmux-test-key" }); + const options = zenmuxModelManagerOptions({ apiKey: "zenmux-test-key", fetch: fetchMock }); expect(options.providerId).toBe("zenmux"); expect(options.fetchDynamicModels).toBeDefined(); const models = await options.fetchDynamicModels?.(); expect(models).not.toBeNull(); - expect(global.fetch).toHaveBeenCalledWith( + expect(fetchMock).toHaveBeenCalledWith( "https://zenmux.ai/api/v1/models", expect.objectContaining({ method: "GET" }), ); diff --git a/packages/ai/test/zhipu-compat.test.ts b/packages/ai/test/zhipu-compat.test.ts index ad45d53fd..598466372 100644 --- a/packages/ai/test/zhipu-compat.test.ts +++ b/packages/ai/test/zhipu-compat.test.ts @@ -1,7 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { zhipuCodingPlanModelManagerOptions } from "@oh-my-pi/pi-ai/provider-models/openai-compat"; import { detectOpenAICompat, resolveOpenAICompat } from "@oh-my-pi/pi-ai/providers/openai-completions-compat"; -import type { Model } from "@oh-my-pi/pi-ai/types"; +import type { FetchImpl, Model } from "@oh-my-pi/pi-ai/types"; /** * Resolver-branch coverage for the `isZhipu` path added by the @@ -22,11 +22,6 @@ const baseModel: Omit, "provider" | "baseUrl"> = { reasoning: true, }; -const originalFetch = global.fetch; - -afterEach(() => { - global.fetch = originalFetch; -}); function zhipuByProvider(): Model<"openai-completions"> { return { ...baseModel, @@ -88,15 +83,17 @@ describe("openai-completions compat — zhipu-coding-plan branch", () => { describe("zhipu-coding-plan model discovery", () => { it("uses the dedicated Coding Plan endpoint by default", async () => { let requestedUrl = ""; - const mockFetch = async (input: string | Request | URL): Promise => { - requestedUrl = input instanceof Request ? input.url : String(input); - return new Response(JSON.stringify({ data: [{ id: "glm-5.1", name: "GLM-5.1" }] }), { - headers: { "content-type": "application/json" }, - }); - }; - global.fetch = Object.assign(mockFetch, { preconnect: originalFetch.preconnect }); + const mockFetch: FetchImpl = Object.assign( + async (input: string | Request | URL): Promise => { + requestedUrl = input instanceof Request ? input.url : String(input); + return new Response(JSON.stringify({ data: [{ id: "glm-5.1", name: "GLM-5.1" }] }), { + headers: { "content-type": "application/json" }, + }); + }, + { preconnect: fetch.preconnect }, + ); - const options = zhipuCodingPlanModelManagerOptions({ apiKey: "test-key" }); + const options = zhipuCodingPlanModelManagerOptions({ apiKey: "test-key", fetch: mockFetch }); expect(typeof options.fetchDynamicModels).toBe("function"); const models = await options.fetchDynamicModels?.(); diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index e070666d2..522d7bdfe 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -4,6 +4,8 @@ ### Added +- Added an optional `fetch` option to `CustomToolContext` so custom tools can use a caller-provided HTTP implementation +- Added optional `fetch` overrides to `ModelRegistry` construction and MCP/web search/tool network calls, enabling callers to inject custom HTTP clients instead of relying on global `fetch` - Added a `bash.enabled` setting to disable the model-facing bash tool while leaving user-initiated bang/RPC bash commands available. - Added an `@` model-selector suffix to pin an aggregator model to a single upstream provider per invocation, e.g. `--model openrouter/z-ai/glm-4.7@cerebras` (sets OpenRouter `provider.only`; Vercel AI Gateway models map to `vercelGatewayRouting.only`). Resolved through `parseModelPattern`, so it works for `--model`/`--smol`, model roles, and the SDK, and composes with a trailing thinking level (`...@cerebras:high`). The base must resolve to an aggregator (`openrouter.ai` / `ai-gateway.vercel.sh`); otherwise the `@` stays part of the id, so ids that legitimately contain `@` (`claude-opus-4-8@default`, `workers-ai/@cf/...`) are unaffected. @@ -9792,4 +9794,4 @@ Initial public release. - Git branch display in footer - Message queueing during streaming responses - OAuth integration for Gmail and Google Calendar access -- HTML export with syntax highlighting and collapsible sections +- HTML export with syntax highlighting and collapsible sections \ No newline at end of file diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 4d07a4d96..51b62be06 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -95,7 +95,7 @@ const STARTUP_MODEL_CACHE_PROVIDER_IDS: readonly string[] = [ ...SPECIAL_MODEL_MANAGER_PROVIDER_IDS, ]; -import type { ApiKeyResolver } from "@oh-my-pi/pi-ai"; +import type { ApiKeyResolver, FetchImpl } from "@oh-my-pi/pi-ai"; import { registerOAuthProvider, unregisterOAuthProviders } from "@oh-my-pi/pi-ai/oauth"; import type { OAuthCredentials, OAuthLoginCallbacks } from "@oh-my-pi/pi-ai/oauth/types"; import { isRecord, logger } from "@oh-my-pi/pi-utils"; @@ -927,6 +927,7 @@ export class ModelRegistry { #runtimeModelManagers: Map; sourceId: string }> = new Map(); #rebuildPending: boolean = false; #rebuildSuspended: number = 0; + #fetch: FetchImpl; /** * @param authStorage - Auth storage for API key resolution @@ -940,7 +941,9 @@ export class ModelRegistry { constructor( readonly authStorage: AuthStorage, modelsPath?: string, + options?: { fetch?: FetchImpl }, ) { + this.#fetch = options?.fetch ?? fetch; this.#modelsConfigFile = ModelsConfigFile.relocate(modelsPath); this.#cacheDbPath = modelsPath ? path.join(path.dirname(modelsPath), "models.db") : undefined; // Set up fallback resolver for custom provider API keys @@ -1629,6 +1632,7 @@ export class ModelRegistry { googleAntigravityModelManagerOptions({ oauthToken, endpoint: this.getProviderBaseUrl("google-antigravity"), + fetch: this.#fetch, }), }, { @@ -1638,6 +1642,7 @@ export class ModelRegistry { googleGeminiCliModelManagerOptions({ oauthToken, endpoint: this.getProviderBaseUrl("google-gemini-cli"), + fetch: this.#fetch, }), }, { @@ -1676,6 +1681,7 @@ export class ModelRegistry { descriptor.createModelManagerOptions({ apiKey: isAuthenticated(apiKey) ? apiKey : undefined, baseUrl: this.getProviderBaseUrl(descriptor.providerId), + fetch: this.#fetch, }), ); } @@ -1727,7 +1733,7 @@ export class ModelRegistry { ): Promise { const showUrl = `${endpoint}/api/show`; try { - const response = await fetch(showUrl, { + const response = await this.#fetch(showUrl, { method: "POST", headers: { ...(headers ?? {}), "Content-Type": "application/json" }, body: JSON.stringify({ model: modelId }), @@ -1775,7 +1781,7 @@ export class ModelRegistry { const endpoint = this.#normalizeOllamaBaseUrl(providerConfig.baseUrl); const tagsUrl = `${endpoint}/api/tags`; const headers = { ...(providerConfig.headers ?? {}) }; - const response = await fetch(tagsUrl, { + const response = await this.#fetch(tagsUrl, { headers, signal: AbortSignal.timeout(250), }); @@ -1819,7 +1825,7 @@ export class ModelRegistry { ): Promise { const propsUrl = `${this.#toLlamaCppNativeBaseUrl(baseUrl)}/props`; try { - const response = await fetch(propsUrl, { + const response = await this.#fetch(propsUrl, { headers, signal: AbortSignal.timeout(150), }); @@ -1850,7 +1856,7 @@ export class ModelRegistry { } const [response, serverMetadata] = await Promise.all([ - fetch(modelsUrl, { + this.#fetch(modelsUrl, { headers, signal: AbortSignal.timeout(250), }), @@ -1902,7 +1908,7 @@ export class ModelRegistry { headers.Authorization = `Bearer ${apiKey}`; } - const response = await fetch(modelsUrl, { + const response = await this.#fetch(modelsUrl, { headers, signal: AbortSignal.timeout(10_000), }); @@ -1964,7 +1970,7 @@ export class ModelRegistry { headers.Authorization = `Bearer ${apiKey}`; } - const response = await fetch(modelsUrl, { + const response = await this.#fetch(modelsUrl, { headers, signal: AbortSignal.timeout(10_000), }); diff --git a/packages/coding-agent/src/extensibility/custom-tools/types.ts b/packages/coding-agent/src/extensibility/custom-tools/types.ts index d1f2cb73b..3703cfaab 100644 --- a/packages/coding-agent/src/extensibility/custom-tools/types.ts +++ b/packages/coding-agent/src/extensibility/custom-tools/types.ts @@ -12,7 +12,7 @@ import type { ToolTier, } from "@oh-my-pi/pi-agent-core"; import type { CompactionResult } from "@oh-my-pi/pi-agent-core/compaction"; -import type { Model, Static, TSchema } from "@oh-my-pi/pi-ai"; +import type { FetchImpl, Model, Static, TSchema } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import type { Rule } from "../../capability/rule"; import type { ModelRegistry } from "../../config/model-registry"; @@ -86,6 +86,8 @@ export interface CustomToolContext { abort(): void; /** Settings instance for the current session. Prefer over the global singleton. */ settings?: Settings; + /** Fetch implementation for outbound HTTP; defaults to global fetch when omitted. */ + fetch?: FetchImpl; /** Whether to auto-approve all destructive tool operations (--auto-approve CLI flag) */ autoApprove?: boolean; } diff --git a/packages/coding-agent/src/mcp/oauth-discovery.ts b/packages/coding-agent/src/mcp/oauth-discovery.ts index 1a1db88e3..8565cfff2 100644 --- a/packages/coding-agent/src/mcp/oauth-discovery.ts +++ b/packages/coding-agent/src/mcp/oauth-discovery.ts @@ -4,6 +4,7 @@ * Automatically detects OAuth requirements from MCP server responses * and extracts authentication endpoints. */ +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; export interface OAuthEndpoints { authorizationUrl: string; @@ -249,7 +250,9 @@ export async function discoverOAuthEndpoints( serverUrl: string, authServerUrl?: string, resourceMetadataUrl?: string, + opts?: { fetch?: FetchImpl }, ): Promise { + const fetchImpl: FetchImpl = opts?.fetch ?? fetch; const wellKnownPaths = [ "/.well-known/oauth-authorization-server", "/.well-known/openid-configuration", @@ -266,7 +269,7 @@ export async function discoverOAuthEndpoints( if (resourceMetadataUrl && !visitedAuthServers.has(resourceMetadataUrl)) { visitedAuthServers.add(resourceMetadataUrl); try { - const metaResp = await fetch(resourceMetadataUrl, { + const metaResp = await fetchImpl(resourceMetadataUrl, { method: "GET", headers: { Accept: "application/json" }, redirect: "follow", @@ -359,7 +362,7 @@ export async function discoverOAuthEndpoints( const urlsToTry = buildWellKnownUrls(path, baseUrl); for (const url of urlsToTry) { try { - const response = await fetch(url.toString(), { + const response = await fetchImpl(url.toString(), { method: "GET", headers: { Accept: "application/json" }, redirect: "follow", @@ -379,7 +382,9 @@ export async function discoverOAuthEndpoints( if (visitedAuthServers.has(discoveredAuthServer)) { continue; } - const discovered = await discoverOAuthEndpoints(serverUrl, discoveredAuthServer); + const discovered = await discoverOAuthEndpoints(serverUrl, discoveredAuthServer, undefined, { + fetch: fetchImpl, + }); if (discovered) return discovered; } } diff --git a/packages/coding-agent/src/mcp/oauth-flow.ts b/packages/coding-agent/src/mcp/oauth-flow.ts index 7ac50765a..baf9f8e75 100644 --- a/packages/coding-agent/src/mcp/oauth-flow.ts +++ b/packages/coding-agent/src/mcp/oauth-flow.ts @@ -8,6 +8,7 @@ import type { OAuthCallbackFlowOptions } from "@oh-my-pi/pi-ai/oauth/callback-server"; import { OAuthCallbackFlow } from "@oh-my-pi/pi-ai/oauth/callback-server"; import type { OAuthController, OAuthCredentials } from "@oh-my-pi/pi-ai/oauth/types"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; const DEFAULT_PORT = 3000; const CALLBACK_PATH = "/callback"; @@ -114,6 +115,8 @@ export interface MCPOAuthConfig { callbackPort?: number; /** Custom callback path (default: /callback or redirectUri pathname) */ callbackPath?: string; + /** Fetch implementation for token exchange and discovery requests. */ + fetch?: FetchImpl; } /** @@ -124,6 +127,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { #resolvedClientId?: string; #registeredClientSecret?: string; #codeVerifier?: string; + #fetch: FetchImpl; constructor( private config: MCPOAuthConfig, @@ -131,6 +135,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { ) { super(ctrl, resolveCallbackOptions(config)); this.#resolvedClientId = this.#resolveClientId(config); + this.#fetch = config.fetch ?? ctrl.fetch ?? fetch; } /** @@ -212,7 +217,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { params.set("client_secret", clientSecret); } - const response = await fetch(this.config.tokenUrl, { + const response = await this.#fetch(this.config.tokenUrl, { method: "POST", headers: { "Content-Type": "application/x-www-form-urlencoded", @@ -289,7 +294,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { if (!registrationEndpoint) return; try { - const response = await fetch(registrationEndpoint, { + const response = await this.#fetch(registrationEndpoint, { method: "POST", headers: { "Content-Type": "application/json", @@ -357,7 +362,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { async #tryWellKnownForRegistration(wellKnownUrl: string): Promise { try { - const response = await fetch(wellKnownUrl, { + const response = await this.#fetch(wellKnownUrl, { method: "GET", headers: { Accept: "application/json" }, }); @@ -374,7 +379,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { async #assertClientIdNotRequired(authorizationUrl: string): Promise { try { - const response = await fetch(authorizationUrl, { + const response = await this.#fetch(authorizationUrl, { method: "GET", redirect: "manual", headers: { Accept: "text/plain,text/html,application/json" }, @@ -402,7 +407,9 @@ export async function refreshMCPOAuthToken( refreshToken: string, clientId?: string, clientSecret?: string, + opts?: { fetch?: FetchImpl }, ): Promise { + const fetchImpl: FetchImpl = opts?.fetch ?? fetch; const params = new URLSearchParams({ grant_type: "refresh_token", refresh_token: refreshToken, @@ -410,7 +417,7 @@ export async function refreshMCPOAuthToken( if (clientId) params.set("client_id", clientId); if (clientSecret) params.set("client_secret", clientSecret); - const response = await fetch(tokenUrl, { + const response = await fetchImpl(tokenUrl, { method: "POST", headers: { "Content-Type": "application/x-www-form-urlencoded" }, body: params.toString(), diff --git a/packages/coding-agent/src/tools/fetch.ts b/packages/coding-agent/src/tools/fetch.ts index 4d16ced81..3bb2284ff 100644 --- a/packages/coding-agent/src/tools/fetch.ts +++ b/packages/coding-agent/src/tools/fetch.ts @@ -637,6 +637,7 @@ export async function renderHtmlToText( settings: Settings, userSignal: AbortSignal | undefined, storage: AgentStorage | null, + fetchOverride?: typeof fetch, ): Promise<{ content: string; ok: boolean; method: string }> { const overallSignal = ptree.combineSignals(userSignal, timeout * 1000); const execOptions = { @@ -650,6 +651,7 @@ export async function renderHtmlToText( // Per-attempt budget for remote endpoints so one stall cannot consume the // whole reader-mode budget and starve the local fallbacks. const remoteSignal = () => ptree.combineSignals(userSignal, remoteBudgetMs); + const fetchImpl = fetchOverride ?? fetch; const runners: Record Promise> = { // Purely local, no network/subprocess: still works on already-loaded HTML @@ -670,14 +672,20 @@ export async function renderHtmlToText( if (!findParallelApiKey(storage)) return null; const parallelResult = await extractWithParallel( [url], - { objective: "Extract the main content", excerpts: true, fullContent: false, signal: remoteSignal() }, + { + objective: "Extract the main content", + excerpts: true, + fullContent: false, + signal: remoteSignal(), + fetch: fetchImpl, + }, storage, ); const firstDocument = parallelResult.results[0]; return firstDocument ? getParallelExtractContent(firstDocument) : null; }, jina: async () => { - const response = await fetch(`https://r.jina.ai/${url}`, { + const response = await fetchImpl(`https://r.jina.ai/${url}`, { headers: { Accept: "text/markdown" }, signal: remoteSignal(), }); diff --git a/packages/coding-agent/src/tools/image-gen.ts b/packages/coding-agent/src/tools/image-gen.ts index ee90c5ca0..263d3d6fd 100644 --- a/packages/coding-agent/src/tools/image-gen.ts +++ b/packages/coding-agent/src/tools/image-gen.ts @@ -1,6 +1,13 @@ import * as os from "node:os"; import * as path from "node:path"; -import { type ApiKey, getAntigravityUserAgent, getEnvApiKey, type Model, withAuth } from "@oh-my-pi/pi-ai"; +import { + type ApiKey, + type FetchImpl, + getAntigravityUserAgent, + getEnvApiKey, + type Model, + withAuth, +} from "@oh-my-pi/pi-ai"; import { CODEX_BASE_URL, getCodexAccountId, @@ -366,7 +373,11 @@ function toDataUrl(image: InlineImageData): string { return `data:${image.mimeType};base64,${image.data}`; } -async function loadImageFromUrl(imageUrl: string, signal?: AbortSignal): Promise { +async function loadImageFromUrl( + imageUrl: string, + fetchImpl: FetchImpl, + signal?: AbortSignal, +): Promise { if (imageUrl.startsWith("data:")) { const normalized = normalizeDataUrl(imageUrl.trim()); if (!normalized.mimeType) { @@ -378,7 +389,7 @@ async function loadImageFromUrl(imageUrl: string, signal?: AbortSignal): Promise return { data: normalized.data, mimeType: normalized.mimeType }; } - const response = await fetch(imageUrl, { signal }); + const response = await fetchImpl(imageUrl, { signal }); if (!response.ok) { const rawText = await response.text(); throw new Error(`Image download failed (${response.status}): ${rawText}`); @@ -850,13 +861,14 @@ async function generateOpenAIHostedImage( model: Model, params: ImageGenParams, inputImages: InlineImageData[], + fetchImpl: FetchImpl, signal: AbortSignal | undefined, sessionId: string | undefined, ): Promise { const promptText = assemblePrompt(params); const stream = model.api === "openai-codex-responses" || model.provider === "openai-codex"; const requestBody = buildOpenAIHostedImageRequest(model, promptText, params, inputImages, stream); - const response = await fetch(getOpenAIResponsesUrl(model), { + const response = await fetchImpl(getOpenAIResponsesUrl(model), { method: "POST", headers: buildOpenAIImageHeaders(model, apiKey, sessionId), body: JSON.stringify(requestBody), @@ -1035,6 +1047,7 @@ export const imageGenTool: CustomTool generateOpenAIHostedImage(key, hostedModel, params, resolvedImages, requestSignal, sessionId), + key => + generateOpenAIHostedImage( + key, + hostedModel, + params, + resolvedImages, + fetchImpl, + requestSignal, + sessionId, + ), { signal: requestSignal }, ); @@ -1117,7 +1139,7 @@ export const imageGenTool: CustomTool { - const resp = await fetch(`${xaiCreds.baseURL}${xaiEndpoint}`, { + const resp = await fetchImpl(`${xaiCreds.baseURL}${xaiEndpoint}`, { method: "POST", headers: { Authorization: `Bearer ${key}`, @@ -1263,7 +1285,7 @@ export const imageGenTool: CustomTool { - const resp = await fetch("https://openrouter.ai/api/v1/chat/completions", { + const resp = await fetchImpl("https://openrouter.ai/api/v1/chat/completions", { method: "POST", headers: { "Content-Type": "application/json", @@ -1343,7 +1365,7 @@ export const imageGenTool: CustomTool { - const resp = await fetch( + const resp = await fetchImpl( `https://generativelanguage.googleapis.com/v1beta/models/${encodeURIComponent(model)}:generateContent`, { method: "POST", diff --git a/packages/coding-agent/src/tools/report-tool-issue.ts b/packages/coding-agent/src/tools/report-tool-issue.ts index e723ff491..a55d6e6fb 100644 --- a/packages/coding-agent/src/tools/report-tool-issue.ts +++ b/packages/coding-agent/src/tools/report-tool-issue.ts @@ -22,6 +22,7 @@ import { Database } from "bun:sqlite"; import path from "node:path"; import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { FetchImpl } from "@oh-my-pi/pi-ai"; import { $env, $flag, getAgentDir, getInstallId, logger, VERSION } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; import type { Settings } from ".."; @@ -260,6 +261,10 @@ export interface FlushOptions { * future debug recipes); never set from the tool's auto-flush path. */ bypassConsent?: boolean; + /** + * Fetch implementation for the push POST. Defaults to global fetch. + */ + fetch?: FetchImpl; /** * Fires once at the start of the loop with the snapshot count of * unpushed rows. Subsequent inserts won't be reflected (the count is @@ -345,6 +350,7 @@ async function performFlush(db: Database, config: PushConfig, options: FlushOpti const totalRow = db.prepare("SELECT COUNT(*) AS n FROM grievances WHERE pushed = 0").get() as { n: number }; options.onStart(totalRow.n); } + const fetchImpl = options.fetch ?? fetch; let totalPushed = 0; for (;;) { const rows = selectStmt.all(FLUSH_BATCH_SIZE) as GrievanceRow[]; @@ -366,7 +372,7 @@ async function performFlush(db: Database, config: PushConfig, options: FlushOpti let response: Response; try { - response = await fetch(config.endpoint, { + response = await fetchImpl(config.endpoint, { method: "POST", headers, body, diff --git a/packages/coding-agent/src/web/kagi.ts b/packages/coding-agent/src/web/kagi.ts index 844feafe0..38041913b 100644 --- a/packages/coding-agent/src/web/kagi.ts +++ b/packages/coding-agent/src/web/kagi.ts @@ -6,7 +6,7 @@ * through the shared {@link AuthStorage} broker (Bearer token), and responses * are categorized result buckets rather than the legacy flat object array. */ -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import { withHardTimeout } from "./search/providers/utils"; const KAGI_SEARCH_URL = "https://kagi.com/api/v1/search"; @@ -156,6 +156,7 @@ export interface KagiSearchOptions { recency?: "day" | "week" | "month" | "year"; sessionId?: string; signal?: AbortSignal; + fetch?: FetchImpl; } export interface KagiSearchSource { @@ -251,7 +252,9 @@ export async function searchWithKagi( throw new KagiApiError("Kagi credentials not found. Set KAGI_API_KEY or login with 'omp /login kagi'."); } - const response = await fetch(KAGI_SEARCH_URL, { + const fetchImpl = options.fetch ?? fetch; + + const response = await fetchImpl(KAGI_SEARCH_URL, { method: "POST", headers: { Authorization: `Bearer ${apiKey}`, diff --git a/packages/coding-agent/src/web/parallel.ts b/packages/coding-agent/src/web/parallel.ts index adadfa4fd..d7be59693 100644 --- a/packages/coding-agent/src/web/parallel.ts +++ b/packages/coding-agent/src/web/parallel.ts @@ -1,4 +1,4 @@ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; +import { type FetchImpl, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { AgentStorage } from "../session/agent-storage"; import { findCredential, withHardTimeout } from "./search/providers/utils"; @@ -62,6 +62,7 @@ export interface ParallelExtractOptions { excerpts?: boolean; fullContent?: boolean; signal?: AbortSignal; + fetch?: FetchImpl; } export class ParallelApiError extends Error { @@ -328,7 +329,8 @@ export async function extractWithParallel( ); } - const response = await fetch(PARALLEL_EXTRACT_URL, { + const fetchImpl = options.fetch ?? fetch; + const response = await fetchImpl(PARALLEL_EXTRACT_URL, { method: "POST", headers: getAuthHeaders(apiKey), body: JSON.stringify({ diff --git a/packages/coding-agent/src/web/search/providers/anthropic.ts b/packages/coding-agent/src/web/search/providers/anthropic.ts index 5a62bb5d8..e1b416bb8 100644 --- a/packages/coding-agent/src/web/search/providers/anthropic.ts +++ b/packages/coding-agent/src/web/search/providers/anthropic.ts @@ -13,6 +13,7 @@ import { buildAnthropicSearchHeaders, buildAnthropicSystemBlocks, buildAnthropicUrl, + type FetchImpl, stripClaudeToolPrefix, withAuth, } from "@oh-my-pi/pi-ai"; @@ -40,6 +41,7 @@ export interface AnthropicSearchParams { max_tokens?: number; temperature?: number; signal?: AbortSignal; + fetch?: FetchImpl; } /** @@ -89,6 +91,7 @@ async function callSearch( maxTokens?: number, temperature?: number, signal?: AbortSignal, + fetchImpl: FetchImpl = fetch, ): Promise { const url = buildAnthropicUrl(auth); const headers = buildAnthropicSearchHeaders(auth); @@ -115,7 +118,7 @@ async function callSearch( body.system = systemBlocks; } - const response = await fetch(url, { + const response = await fetchImpl(url, { method: "POST", headers, body: JSON.stringify(body), @@ -275,6 +278,7 @@ export async function searchAnthropic( maxTokens, params.temperature, params.signal, + params.fetch, ), { signal: params.signal, diff --git a/packages/coding-agent/src/web/search/providers/base.ts b/packages/coding-agent/src/web/search/providers/base.ts index 3cf075a90..07001ea77 100644 --- a/packages/coding-agent/src/web/search/providers/base.ts +++ b/packages/coding-agent/src/web/search/providers/base.ts @@ -1,4 +1,4 @@ -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import type { SearchProviderId, SearchResponse } from "../types"; /** @@ -30,6 +30,7 @@ export interface SearchParams { recency?: "day" | "week" | "month" | "year"; systemPrompt: string; signal?: AbortSignal; + fetch?: FetchImpl; maxOutputTokens?: number; numSearchResults?: number; temperature?: number; diff --git a/packages/coding-agent/src/web/search/providers/brave.ts b/packages/coding-agent/src/web/search/providers/brave.ts index 283228311..5fdbc2c02 100644 --- a/packages/coding-agent/src/web/search/providers/brave.ts +++ b/packages/coding-agent/src/web/search/providers/brave.ts @@ -4,7 +4,7 @@ * Calls Brave's web search REST API and maps results into the unified * SearchResponse shape used by the web search tool. */ -import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, type FetchImpl, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; @@ -28,6 +28,7 @@ export interface BraveSearchParams { num_results?: number; recency?: "day" | "week" | "month" | "year"; signal?: AbortSignal; + fetch?: FetchImpl; } interface BraveSearchResult { @@ -80,7 +81,8 @@ async function callBraveSearch( url.searchParams.set("freshness", RECENCY_MAP[params.recency]); } - const response = await fetch(url, { + const fetchImpl = params.fetch ?? fetch; + const response = await fetchImpl(url, { headers: { Accept: "application/json", "X-Subscription-Token": apiKey, @@ -144,6 +146,7 @@ export class BraveProvider extends SearchProvider { num_results: params.numSearchResults ?? params.limit, recency: params.recency, signal: params.signal, + fetch: params.fetch, }); } } diff --git a/packages/coding-agent/src/web/search/providers/codex.ts b/packages/coding-agent/src/web/search/providers/codex.ts index 5cd46ecfd..6ecd551c8 100644 --- a/packages/coding-agent/src/web/search/providers/codex.ts +++ b/packages/coding-agent/src/web/search/providers/codex.ts @@ -7,7 +7,7 @@ * SQLite store, never POSTs the broker sentinel to an OpenAI token endpoint. */ import * as os from "node:os"; -import { type AuthStorage, getBundledModels } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, type FetchImpl, getBundledModels } from "@oh-my-pi/pi-ai"; import { decodeJwt } from "@oh-my-pi/pi-ai/oauth/openai-codex"; import { $env, readSseJson } from "@oh-my-pi/pi-utils"; import packageJson from "../../../../package.json" with { type: "json" }; @@ -66,6 +66,7 @@ function shouldRetryWithNextDefaultModel(error: unknown): boolean { export interface CodexSearchParams { signal?: AbortSignal; + fetch?: FetchImpl; query: string; system_prompt?: string; num_results?: number; @@ -322,6 +323,7 @@ async function callCodexSearch( systemPrompt?: string; searchContextSize?: "low" | "medium" | "high"; modelId: string; + fetch?: FetchImpl; }, ): Promise<{ answer: string; @@ -356,7 +358,8 @@ async function callCodexSearch( instructions: options.systemPrompt ?? DEFAULT_INSTRUCTIONS, }; - const response = await fetch(url, { + const fetchImpl = options.fetch ?? fetch; + const response = await fetchImpl(url, { method: "POST", headers, body: JSON.stringify(body), @@ -522,6 +525,7 @@ export async function searchCodex(params: SearchParams): Promise systemPrompt: params.systemPrompt, searchContextSize: "high", modelId, + fetch: params.fetch, }); break; } catch (error) { diff --git a/packages/coding-agent/src/web/search/providers/exa.ts b/packages/coding-agent/src/web/search/providers/exa.ts index 7c0642cbd..a82e20ec8 100644 --- a/packages/coding-agent/src/web/search/providers/exa.ts +++ b/packages/coding-agent/src/web/search/providers/exa.ts @@ -6,10 +6,10 @@ * Requests per-result summaries via `contents.summary` and synthesizes * them into a combined `answer` string on the SearchResponse. */ -import { type ApiKey, type AuthStorage, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type AuthStorage, type FetchImpl, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; import { settings } from "../../../config/settings"; -import { callExaTool, findApiKey, isSearchResponse } from "../../../exa/mcp-client"; - +import { findApiKey, isSearchResponse } from "../../../exa/mcp-client"; +import { parseSSE } from "../../../mcp/json-rpc"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { dateToAgeSeconds } from "../utils"; @@ -32,6 +32,7 @@ export interface ExaSearchParams { start_published_date?: string; end_published_date?: string; signal?: AbortSignal; + fetch?: FetchImpl; /** * Credential source. Resolved before falling back to `EXA_API_KEY` so * Exa works when the key is stored via the broker/auth pipeline. @@ -62,6 +63,48 @@ function asRecord(value: unknown): Record | null { return value as Record; } +function parseJsonContent(text: string): unknown | null { + try { + return JSON.parse(text) as unknown; + } catch { + return null; + } +} + +function normalizeExaMcpPayload(payload: unknown): unknown { + const candidates: unknown[] = []; + const root = asRecord(payload); + + if (root) { + if (root.structuredContent !== undefined) candidates.push(root.structuredContent); + if (root.data !== undefined) candidates.push(root.data); + if (root.result !== undefined) candidates.push(root.result); + candidates.push(root); + + const content = root.content; + if (Array.isArray(content)) { + for (const item of content) { + const part = asRecord(item); + if (!part) continue; + const text = part.text; + if (typeof text !== "string" || text.trim().length === 0) continue; + const parsed = parseJsonContent(text); + if (parsed !== null) candidates.push(parsed); + } + } + } else { + candidates.push(payload); + } + + for (const candidate of candidates) { + if (isSearchResponse(candidate)) { + return candidate; + } + } + + return payload; +} + function parseOptionalField(section: string, label: string): string | null | undefined { const regex = new RegExp(`(?:^|\\n)${label}:\\s*([^\\n]*)`); const match = section.match(regex); @@ -180,7 +223,8 @@ export function buildExaRequestBody(params: ExaSearchParams): Record { const body = buildExaRequestBody(params); - const response = await fetch(EXA_API_URL, { + const fetchImpl = params.fetch ?? fetch; + const response = await fetchImpl(EXA_API_URL, { method: "POST", headers: { "Content-Type": "application/json", @@ -211,14 +255,52 @@ function buildExaMcpArgs(params: ExaSearchParams): Record { } async function callExaMcpSearch(params: ExaSearchParams): Promise { - const response = await callExaTool("web_search_exa", buildExaMcpArgs(params), findApiKey(), { + const query = new URLSearchParams(); + const apiKey = findApiKey(); + if (apiKey) query.set("exaApiKey", apiKey); + query.set("tools", "web_search_exa"); + const fetchImpl = params.fetch ?? fetch; + const response = await fetchImpl(`https://mcp.exa.ai/mcp?${query.toString()}`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Accept: "application/json, text/event-stream", + }, + body: JSON.stringify({ + jsonrpc: "2.0", + id: Math.random().toString(36).slice(2), + method: "tools/call", + params: { + name: "web_search_exa", + arguments: buildExaMcpArgs(params), + }, + }), signal: withHardTimeout(params.signal), }); - if (isSearchResponse(response)) { - return response as ExaSearchResponse; + if (!response.ok) { + throw new Error(`MCP request failed: ${response.status} ${response.statusText}`); + } + const mcpResponse = parseSSE(await response.text()) as { + result?: { + content?: Array<{ type: string; text?: string }>; + }; + error?: { + code: number; + message: string; + }; + } | null; + if (!mcpResponse) { + throw new Error("Failed to parse MCP response"); + } + if (mcpResponse.error) { + throw new Error(`MCP error: ${mcpResponse.error.message}`); + } + const responsePayload = normalizeExaMcpPayload(mcpResponse.result); + if (isSearchResponse(responsePayload)) { + return responsePayload as ExaSearchResponse; } - const parsed = parseExaMcpTextPayload(response); + const parsed = parseExaMcpTextPayload(responsePayload); if (parsed) { return parsed; } @@ -312,6 +394,7 @@ export class ExaProvider extends SearchProvider { signal: params.signal, authStorage: params.authStorage, sessionId: params.sessionId, + fetch: params.fetch, }); } } diff --git a/packages/coding-agent/src/web/search/providers/gemini.ts b/packages/coding-agent/src/web/search/providers/gemini.ts index 215e2a5cf..80741880b 100644 --- a/packages/coding-agent/src/web/search/providers/gemini.ts +++ b/packages/coding-agent/src/web/search/providers/gemini.ts @@ -11,6 +11,7 @@ import { ANTIGRAVITY_SYSTEM_INSTRUCTION, type AuthStorage, + type FetchImpl, getAntigravityUserAgent, getGeminiCliHeaders, } from "@oh-my-pi/pi-ai"; @@ -51,6 +52,7 @@ export interface GeminiSearchParams extends GeminiToolParams { signal?: AbortSignal; authStorage: AuthStorage; sessionId?: string; + fetch?: FetchImpl; } export function buildGeminiRequestTools(params: GeminiToolParams): Array>> { @@ -156,6 +158,7 @@ async function callGeminiSearch( maxOutputTokens: number | undefined, temperature: number | undefined, toolParams: GeminiToolParams, + fetchImpl: FetchImpl | undefined, signal: AbortSignal | undefined, ): Promise<{ answer: string; @@ -237,6 +240,7 @@ async function callGeminiSearch( const response = await fetchWithRetry(urlFor, { ...buildInit(), + fetch: fetchImpl, maxAttempts: MAX_RETRIES + 1, defaultDelayMs: attempt => BASE_DELAY_MS * 2 ** attempt, maxDelayMs: RATE_LIMIT_BUDGET_MS, @@ -405,6 +409,7 @@ export async function searchGemini(params: GeminiSearchParams): Promise { +async function callJinaSearch( + apiKey: string, + query: string, + signal?: AbortSignal, + fetchImpl: FetchImpl = fetch, +): Promise { const requestUrl = `${JINA_SEARCH_URL}/${encodeURIComponent(query)}`; - const response = await fetch(requestUrl, { + const response = await fetchImpl(requestUrl, { headers: { Accept: "application/json", Authorization: `Bearer ${apiKey}`, @@ -62,7 +69,7 @@ export async function searchJina(params: JinaSearchParams): Promise { + search(params: SearchParamsWithFetch): Promise { + const fetchImpl = params.fetch; + return searchJina({ query: params.query, num_results: params.numSearchResults ?? params.limit, signal: params.signal, + fetch: fetchImpl, }); } } diff --git a/packages/coding-agent/src/web/search/providers/kagi.ts b/packages/coding-agent/src/web/search/providers/kagi.ts index cc2ac223b..081022304 100644 --- a/packages/coding-agent/src/web/search/providers/kagi.ts +++ b/packages/coding-agent/src/web/search/providers/kagi.ts @@ -3,7 +3,7 @@ * * Thin wrapper that adapts shared Kagi API utilities to SearchResponse shape. */ -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import type { SearchResponse } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { KagiApiError, searchWithKagi } from "../../kagi"; @@ -12,6 +12,8 @@ import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; import { classifyProviderHttpError, toSearchSources } from "./utils"; +type SearchParamsWithFetch = SearchParams & { fetch?: FetchImpl }; + const DEFAULT_NUM_RESULTS = 10; const MAX_NUM_RESULTS = 40; @@ -23,6 +25,7 @@ export async function searchKagi(params: { signal?: AbortSignal; authStorage: AuthStorage; sessionId?: string; + fetch?: FetchImpl; }): Promise { const numResults = clampNumResults(params.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); @@ -34,6 +37,7 @@ export async function searchKagi(params: { recency: params.recency, sessionId: params.sessionId, signal: params.signal, + fetch: params.fetch, }, params.authStorage, ); @@ -66,7 +70,9 @@ export class KagiProvider extends SearchProvider { return authStorage.hasAuth("kagi"); } - search(params: SearchParams): Promise { + search(params: SearchParamsWithFetch): Promise { + const fetchImpl = params.fetch; + return searchKagi({ query: params.query, num_results: params.numSearchResults ?? params.limit, @@ -74,6 +80,7 @@ export class KagiProvider extends SearchProvider { signal: params.signal, authStorage: params.authStorage, sessionId: params.sessionId, + fetch: fetchImpl, }); } } diff --git a/packages/coding-agent/src/web/search/providers/kimi.ts b/packages/coding-agent/src/web/search/providers/kimi.ts index 08401dbe6..d4f282d6e 100644 --- a/packages/coding-agent/src/web/search/providers/kimi.ts +++ b/packages/coding-agent/src/web/search/providers/kimi.ts @@ -4,7 +4,7 @@ * Uses Moonshot Kimi Code search API to retrieve web results. * Endpoint: POST https://api.kimi.com/coding/v1/search */ -import { type ApiKey, type AuthStorage, withAuth } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type AuthStorage, type FetchImpl, withAuth } from "@oh-my-pi/pi-ai"; import { $env } from "@oh-my-pi/pi-utils"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; @@ -14,6 +14,8 @@ import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; import { classifyProviderHttpError, withHardTimeout } from "./utils"; +type SearchParamsWithFetch = SearchParams & { fetch?: FetchImpl }; + const KIMI_SEARCH_URL = "https://api.kimi.com/coding/v1/search"; const DEFAULT_NUM_RESULTS = 10; @@ -27,6 +29,7 @@ export interface KimiSearchParams { signal?: AbortSignal; authStorage: AuthStorage; sessionId?: string; + fetch?: FetchImpl; } interface KimiSearchResult { @@ -78,9 +81,16 @@ async function resolveKey( async function callKimiSearch( apiKey: string, - params: { query: string; limit: number; includeContent: boolean; signal?: AbortSignal }, + params: { + query: string; + limit: number; + includeContent: boolean; + signal?: AbortSignal; + fetch?: FetchImpl; + }, ): Promise<{ response: KimiSearchResponse; requestId?: string }> { - const response = await fetch(resolveBaseUrl(), { + const fetchImpl = params.fetch ?? fetch; + const response = await fetchImpl(resolveBaseUrl(), { method: "POST", headers: { Accept: "application/json", @@ -130,6 +140,7 @@ export async function searchKimi(params: KimiSearchParams): Promise { + search(params: SearchParamsWithFetch): Promise { + const fetchImpl = params.fetch; + return searchKimi({ query: params.query, num_results: params.numSearchResults ?? params.limit, signal: params.signal, authStorage: params.authStorage, sessionId: params.sessionId, + fetch: fetchImpl, }); } } diff --git a/packages/coding-agent/src/web/search/providers/parallel.ts b/packages/coding-agent/src/web/search/providers/parallel.ts index 876ef4a8c..d6e11b6d7 100644 --- a/packages/coding-agent/src/web/search/providers/parallel.ts +++ b/packages/coding-agent/src/web/search/providers/parallel.ts @@ -1,4 +1,4 @@ -import { type ApiKey, type AuthStorage, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type AuthStorage, type FetchImpl, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; import type { SearchResponse } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { ParallelApiError, type ParallelSearchResult, type ParallelSearchSource } from "../../parallel"; @@ -112,6 +112,7 @@ async function searchWithAuthStorage( queries: string[], params: { signal?: AbortSignal; + fetch?: FetchImpl; }, authStorage: AuthStorage, sessionId?: string, @@ -131,7 +132,7 @@ async function searchWithAuthStorage( return withAuth( keyOrResolver, async key => { - const response = await fetch(PARALLEL_SEARCH_URL, { + const response = await (params.fetch ?? fetch)(PARALLEL_SEARCH_URL, { method: "POST", headers: { Accept: "application/json", @@ -165,6 +166,7 @@ export async function searchParallel( query: string; num_results?: number; signal?: AbortSignal; + fetch?: FetchImpl; }, authStorage: AuthStorage, sessionId?: string, @@ -177,6 +179,7 @@ export async function searchParallel( [params.query], { signal: params.signal, + fetch: params.fetch, }, authStorage, sessionId, @@ -213,6 +216,7 @@ export class ParallelProvider extends SearchProvider { query: params.query, num_results: params.numSearchResults ?? params.limit, signal: params.signal, + fetch: params.fetch, }, params.authStorage, params.sessionId, diff --git a/packages/coding-agent/src/web/search/providers/perplexity.ts b/packages/coding-agent/src/web/search/providers/perplexity.ts index 70284cdfb..ca7d8fc93 100644 --- a/packages/coding-agent/src/web/search/providers/perplexity.ts +++ b/packages/coding-agent/src/web/search/providers/perplexity.ts @@ -8,7 +8,7 @@ * - Anonymous via `www.perplexity.ai/rest/sse/perplexity_ask` */ -import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, type FetchImpl, getEnvApiKey } from "@oh-my-pi/pi-ai"; import { $env, readSseJson } from "@oh-my-pi/pi-utils"; import type { PerplexityMessageOutput, @@ -275,6 +275,7 @@ export interface PerplexitySearchParams { num_search_results?: number; authStorage: AuthStorage; sessionId?: string; + fetch?: FetchImpl; } /** Find PERPLEXITY_API_KEY from environment or .env files (also checks PPLX_API_KEY) */ @@ -356,9 +357,10 @@ async function findPerplexityAuth( async function callPerplexityApi( apiKey: string, request: PerplexityRequest, + fetchImpl: FetchImpl | undefined, signal?: AbortSignal, ): Promise { - const response = await fetch(PERPLEXITY_API_URL, { + const response = await (fetchImpl ?? fetch)(PERPLEXITY_API_URL, { method: "POST", headers: { Authorization: `Bearer ${apiKey}`, @@ -505,7 +507,7 @@ async function callPerplexityAsk( requestParams.source = "default"; } - const response = await fetch(PERPLEXITY_OAUTH_ASK_URL, { + const response = await (params.fetch ?? fetch)(PERPLEXITY_OAUTH_ASK_URL, { method: "POST", headers, body: JSON.stringify({ @@ -686,7 +688,7 @@ export async function searchPerplexity(params: PerplexitySearchParams): Promise< request.search_recency_filter = params.search_recency_filter; } - const response = await callPerplexityApi(auth.token, request, params.signal); + const response = await callPerplexityApi(auth.token, request, params.fetch, params.signal); const result = parseResponse(response); result.authMode = "api_key"; return applySourceLimit(result, params.num_results); @@ -722,6 +724,7 @@ export class PerplexityProvider extends SearchProvider { num_results: params.limit, authStorage: params.authStorage, sessionId: params.sessionId, + fetch: params.fetch, }); } } diff --git a/packages/coding-agent/src/web/search/providers/searxng.ts b/packages/coding-agent/src/web/search/providers/searxng.ts index 8eb05e500..6d75907e4 100644 --- a/packages/coding-agent/src/web/search/providers/searxng.ts +++ b/packages/coding-agent/src/web/search/providers/searxng.ts @@ -25,7 +25,7 @@ * Reference: https://docs.searxng.org/dev/search_api.html */ -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import { settings } from "../../../config/settings"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; @@ -207,12 +207,13 @@ async function callSearXNGSearch( categories?: string; language?: string; signal?: AbortSignal; + fetch?: FetchImpl; }, auth: SearXNGAuth | null, ): Promise { const { url, headers } = buildRequest(endpoint, params, auth); - const response = await fetch(url, { + const response = await (params.fetch ?? fetch)(url, { headers, signal: withHardTimeout(params.signal), }); @@ -233,6 +234,7 @@ export async function searchSearXNG(params: { num_results?: number; recency?: "day" | "week" | "month" | "year"; signal?: AbortSignal; + fetch?: FetchImpl; }): Promise { const numResults = clampNumResults(params.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); @@ -260,6 +262,7 @@ export async function searchSearXNG(params: { ...params, categories, language, + fetch: params.fetch, }, auth, ); @@ -304,6 +307,7 @@ export class SearXNGProvider extends SearchProvider { num_results: params.numSearchResults ?? params.limit, recency: params.recency, signal: params.signal, + fetch: params.fetch, }); } } diff --git a/packages/coding-agent/src/web/search/providers/synthetic.ts b/packages/coding-agent/src/web/search/providers/synthetic.ts index 12d4a6ca1..4a84c3547 100644 --- a/packages/coding-agent/src/web/search/providers/synthetic.ts +++ b/packages/coding-agent/src/web/search/providers/synthetic.ts @@ -5,13 +5,15 @@ * Endpoint: POST https://api.synthetic.new/v2/search */ -import { type ApiKey, type AuthStorage, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type AuthStorage, type FetchImpl, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; import { classifyProviderHttpError, withHardTimeout } from "./utils"; +type SearchParamsWithFetch = SearchParams & { fetch?: FetchImpl }; + const SYNTHETIC_SEARCH_URL = "https://api.synthetic.new/v2/search"; interface SyntheticSearchResult { @@ -39,8 +41,9 @@ async function callSyntheticSearch( apiKey: string, query: string, signal?: AbortSignal, + fetchImpl: FetchImpl = fetch, ): Promise { - const response = await fetch(SYNTHETIC_SEARCH_URL, { + const response = await fetchImpl(SYNTHETIC_SEARCH_URL, { method: "POST", headers: { "Content-Type": "application/json", @@ -65,12 +68,13 @@ async function callSyntheticSearch( } /** Execute Synthetic web search. */ -export async function searchSynthetic(params: SearchParams): Promise { +export async function searchSynthetic(params: SearchParamsWithFetch): Promise { const keyOrResolver: ApiKey = params.authStorage.resolver("synthetic", { sessionId: params.sessionId, }); - const data = await withAuth(keyOrResolver, key => callSyntheticSearch(key, params.query, params.signal), { + const fetchImpl = params.fetch; + const data = await withAuth(keyOrResolver, key => callSyntheticSearch(key, params.query, params.signal, fetchImpl), { signal: params.signal, missingKeyMessage: "Synthetic credentials not found. Set SYNTHETIC_API_KEY or login with 'omp /login synthetic'.", }); @@ -104,7 +108,7 @@ export class SyntheticProvider extends SearchProvider { return authStorage.hasAuth("synthetic") || !!getEnvApiKey("synthetic"); } - search(params: SearchParams): Promise { + search(params: SearchParamsWithFetch): Promise { return searchSynthetic(params); } } diff --git a/packages/coding-agent/src/web/search/providers/tavily.ts b/packages/coding-agent/src/web/search/providers/tavily.ts index 9b5e0a9f3..ee38837d7 100644 --- a/packages/coding-agent/src/web/search/providers/tavily.ts +++ b/packages/coding-agent/src/web/search/providers/tavily.ts @@ -4,7 +4,7 @@ * Uses Tavily's agent-focused search API to return structured results with an * optional synthesized answer. */ -import { type ApiKey, type AuthStorage, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type AuthStorage, type FetchImpl, getEnvApiKey, withAuth } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; @@ -21,6 +21,7 @@ export interface TavilySearchParams { num_results?: number; recency?: "day" | "week" | "month" | "year"; signal?: AbortSignal; + fetch?: FetchImpl; } interface TavilySearchResult { @@ -89,7 +90,7 @@ export function buildRequestBody(params: TavilySearchParams): Record { - const response = await fetch(TAVILY_SEARCH_URL, { + const response = await (params.fetch ?? fetch)(TAVILY_SEARCH_URL, { method: "POST", headers: { "Content-Type": "application/json", @@ -126,6 +127,7 @@ export async function searchTavily(params: SearchParams): Promise, signal?: AbortSignal): Promise { - const response = await fetch(ZAI_MCP_URL, { +async function callZaiTool( + apiKey: string, + args: Record, + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise { + const response = await fetchImpl(ZAI_MCP_URL, { method: "POST", headers: { Authorization: `Bearer ${apiKey}`, @@ -158,6 +164,7 @@ async function callZaiTool(apiKey: string, args: Record, signal async function callZaiSearch(apiKey: string, params: ZaiSearchParams): Promise { const count = params.num_results ?? DEFAULT_NUM_RESULTS; + const fetchImpl = params.fetch ?? fetch; const attempts: Record[] = [ { query: params.query, count }, { search_query: params.query, count }, @@ -167,7 +174,7 @@ async function callZaiSearch(apiKey: string, params: ZaiSearchParams): Promise { + const { fetch: fetchOverride } = params as ZaiProviderSearchParams; return searchZai({ query: params.query, num_results: params.numSearchResults ?? params.limit, signal: params.signal, authStorage: params.authStorage, sessionId: params.sessionId, + fetch: fetchOverride, }); } } diff --git a/packages/coding-agent/test/auth-storage-minimax-login.test.ts b/packages/coding-agent/test/auth-storage-minimax-login.test.ts index 0e40fd257..e44797423 100644 --- a/packages/coding-agent/test/auth-storage-minimax-login.test.ts +++ b/packages/coding-agent/test/auth-storage-minimax-login.test.ts @@ -2,8 +2,9 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; +import type { FetchImpl } from "@oh-my-pi/pi-ai"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; +import { Snowflake } from "@oh-my-pi/pi-utils"; describe("AuthStorage MiniMax login", () => { let tempDir: string; @@ -25,13 +26,13 @@ describe("AuthStorage MiniMax login", () => { }); test("replaces existing MiniMax Coding Plan API key on relogin", async () => { - using _hook = hookFetch( - () => new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }), - ); + const fetchMock: FetchImpl = async () => + new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); const loginCallbacks = { onAuth: () => {}, onPrompt: async () => currentApiKey, + fetch: fetchMock, }; await authStorage.login("minimax-code", loginCallbacks); diff --git a/packages/coding-agent/test/compaction.test.ts b/packages/coding-agent/test/compaction.test.ts index 1b733df5d..f16de3475 100644 --- a/packages/coding-agent/test/compaction.test.ts +++ b/packages/coding-agent/test/compaction.test.ts @@ -25,7 +25,7 @@ import { type SessionMessageEntry, type ThinkingLevelChangeEntry, } from "@oh-my-pi/pi-coding-agent/session/session-manager"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import { mockFetch } from "./helpers/fetch-mock"; import { e2eApiKey } from "./utilities"; // ============================================================================ @@ -349,23 +349,25 @@ describe("remote compaction setting", () => { throw new Error("Expected compaction preparation"); } - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ summary: "remote summary" }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); const completeSpy = vi .spyOn(ai, "completeSimple") .mockResolvedValueOnce(createAssistantMessage("Local history summary")) .mockResolvedValueOnce(createAssistantMessage("Local turn summary")) .mockResolvedValueOnce(createAssistantMessage("Local short summary")); - const result = await compact(preparation, model, "test-api-key"); + const result = await compact(preparation, model, "test-api-key", undefined, undefined, { + fetch: fetchSpy, + }); - expect(fetchSpy).not.toHaveBeenCalled(); + expect(fetchHandler).not.toHaveBeenCalled(); expect(completeSpy).toHaveBeenCalledTimes(3); expect(result.summary).toContain("Local history summary"); expect(result.shortSummary).toBe("Local short summary"); @@ -429,26 +431,28 @@ describe("remote compaction setting", () => { { type: "message", role: "user", content: [{ type: "input_text", text: "Compacted retained user" }] }, { type: "compaction", encrypted_content: "new_encrypted" }, ]; - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ output: remoteOutput }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); const completeSimpleSpy = vi.spyOn(ai, "completeSimple"); completeSimpleSpy .mockResolvedValueOnce(createAssistantMessage("History summary")) .mockResolvedValueOnce(createAssistantMessage("Turn prefix summary")) .mockResolvedValueOnce(createAssistantMessage("Short summary")); - const result = await compact(preparation, model, "test-api-key"); - const requestBody = JSON.parse(String(fetchSpy.mock.calls[0]?.[1]?.body)) as { + const result = await compact(preparation, model, "test-api-key", undefined, undefined, { + fetch: fetchSpy, + }); + const requestBody = JSON.parse(String(fetchHandler.mock.calls[0]?.[1]?.body)) as { input: Array>; }; - expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(fetchHandler).toHaveBeenCalledTimes(1); expect(requestBody.input[0]).toEqual({ type: "message", role: "user", @@ -497,18 +501,18 @@ describe("remote compaction setting", () => { }); if (!preparation) throw new Error("Expected compaction preparation"); - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ output: [{ type: "compaction", encrypted_content: "new_encrypted" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); vi.spyOn(ai, "completeSimple").mockResolvedValue(createAssistantMessage("Short summary")); - await compact(preparation, model, "test-api-key"); - const requestBody = JSON.parse(String(fetchSpy.mock.calls[0]?.[1]?.body)) as { + await compact(preparation, model, "test-api-key", undefined, undefined, { fetch: fetchSpy }); + const requestBody = JSON.parse(String(fetchHandler.mock.calls[0]?.[1]?.body)) as { input: Array>; }; @@ -540,20 +544,20 @@ describe("remote compaction setting", () => { }); if (!preparation) throw new Error("Expected compaction preparation"); - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ output: [{ type: "compaction", encrypted_content: "new_encrypted" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); vi.spyOn(ai, "completeSimple").mockResolvedValue(createAssistantMessage("Short summary")); - await compact(preparation, model, "test-api-key"); + await compact(preparation, model, "test-api-key", undefined, undefined, { fetch: fetchSpy }); - expect(fetchSpy).toHaveBeenCalledTimes(1); - expect(fetchSpy.mock.calls[0]?.[0]).toBe("https://chatgpt.com/backend-api/codex/responses/compact"); + expect(fetchHandler).toHaveBeenCalledTimes(1); + expect(fetchHandler.mock.calls[0]?.[0]).toBe("https://chatgpt.com/backend-api/codex/responses/compact"); }); it("preserves codex assistant text signature metadata in remote compaction history", async () => { @@ -591,18 +595,18 @@ describe("remote compaction setting", () => { }); if (!preparation) throw new Error("Expected compaction preparation"); - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ output: [{ type: "compaction", encrypted_content: "new_encrypted" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); vi.spyOn(ai, "completeSimple").mockResolvedValue(createAssistantMessage("Short summary")); - await compact(preparation, model, "test-api-key"); - const requestBody = JSON.parse(String(fetchSpy.mock.calls[0]?.[1]?.body)) as { + await compact(preparation, model, "test-api-key", undefined, undefined, { fetch: fetchSpy }); + const requestBody = JSON.parse(String(fetchHandler.mock.calls[0]?.[1]?.body)) as { input: Array>; }; const assistantItem = requestBody.input.find(item => item.type === "message" && item.role === "assistant"); @@ -638,20 +642,21 @@ describe("remote compaction setting", () => { { type: "message", role: "assistant", content: [{ type: "output_text", text: "Kept assistant" }] }, { type: "compaction", encrypted_content: "new_encrypted" }, ]; - const fetchSpy = vi.fn( - (_input, _init, _next) => + const fetchHandler = vi.fn( + async (_input, _init) => new Response(JSON.stringify({ output: remoteOutput }), { status: 200, headers: { "Content-Type": "application/json" }, }), ); - using _hook = hookFetch(fetchSpy); + const fetchSpy = mockFetch(fetchHandler); vi.spyOn(ai, "completeSimple").mockResolvedValue(createAssistantMessage("Short summary")); const result = await compact(preparation, model, "test-api-key", undefined, undefined, { remoteInstructions: "BASE INSTRUCTIONS", + fetch: fetchSpy, }); - const requestBody = JSON.parse(String(fetchSpy.mock.calls[0]?.[1]?.body)) as { + const requestBody = JSON.parse(String(fetchHandler.mock.calls[0]?.[1]?.body)) as { instructions: string; }; diff --git a/packages/coding-agent/test/helpers/fetch-mock.ts b/packages/coding-agent/test/helpers/fetch-mock.ts new file mode 100644 index 000000000..df7545708 --- /dev/null +++ b/packages/coding-agent/test/helpers/fetch-mock.ts @@ -0,0 +1,18 @@ +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; + +type FetchHandler = (input: string | URL | Request, init?: RequestInit) => Response | Promise; + +/** Wrap a fetch handler as a {@link FetchImpl}, normalizing sync `Response` returns. */ +export function mockFetch(fn: FetchHandler): FetchImpl { + return async (input, init) => fn(input, init); +} + +/** Satisfies Bun's `typeof fetch` (includes `preconnect`). */ +export function asGlobalFetch(fn: FetchHandler): typeof fetch { + return Object.assign(async (input: string | URL | Request, init?: RequestInit) => fn(input, init), { + preconnect: fetch.preconnect, + }); +} + +/** `RequestInfo` alias for test fetch handlers (DOM lib name). */ +export type FetchInput = string | URL | Request; diff --git a/packages/coding-agent/test/issue-1528-discovery-default-max-tokens.test.ts b/packages/coding-agent/test/issue-1528-discovery-default-max-tokens.test.ts index 17e2c16b4..b01b75f20 100644 --- a/packages/coding-agent/test/issue-1528-discovery-default-max-tokens.test.ts +++ b/packages/coding-agent/test/issue-1528-discovery-default-max-tokens.test.ts @@ -2,9 +2,10 @@ import { afterEach, beforeEach, describe, expect, test } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; +import { Snowflake } from "@oh-my-pi/pi-utils"; /** * Issue #1528: auto-discovered OpenAI-compatible models defaulted to @@ -49,7 +50,7 @@ describe("issue #1528 discovery maxTokens default", () => { ].join("\n"), ); - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url !== "https://api.example.com/v1/models") { throw new Error(`Unexpected URL: ${url}`); @@ -58,9 +59,9 @@ describe("issue #1528 discovery maxTokens default", () => { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const registry = new ModelRegistry(authStorage, modelsPath); + const registry = new ModelRegistry(authStorage, modelsPath, { fetch: fetchMock }); await registry.refreshProvider("deepseek-compat"); const model = registry.find("deepseek-compat", "deepseek-v4-pro"); @@ -82,7 +83,7 @@ describe("issue #1528 discovery maxTokens default", () => { ].join("\n"), ); - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url !== "https://proxy.example.com/v1/models") { throw new Error(`Unexpected URL: ${url}`); @@ -93,9 +94,9 @@ describe("issue #1528 discovery maxTokens default", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }); + }; - const registry = new ModelRegistry(authStorage, modelsPath); + const registry = new ModelRegistry(authStorage, modelsPath, { fetch: fetchMock }); await registry.refreshProvider("newapi-proxy"); const model = registry.find("newapi-proxy", "newapi-private-openai-model"); @@ -122,7 +123,7 @@ describe("issue #1528 discovery maxTokens default", () => { ].join("\n"), ); - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url !== "https://proxy.example.com/v1/models") { throw new Error(`Unexpected URL: ${url}`); @@ -136,9 +137,9 @@ describe("issue #1528 discovery maxTokens default", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }); + }; - const registry = new ModelRegistry(authStorage, modelsPath); + const registry = new ModelRegistry(authStorage, modelsPath, { fetch: fetchMock }); await registry.refreshProvider("newapi-proxy"); const sonnet = registry.find("newapi-proxy", "claude-3-5-sonnet"); @@ -172,7 +173,7 @@ describe("issue #1528 discovery maxTokens default", () => { ].join("\n"), ); - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url !== "https://anthropic-reseller.example.com/v1/models") { throw new Error(`Unexpected URL: ${url}`); @@ -181,9 +182,9 @@ describe("issue #1528 discovery maxTokens default", () => { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - const registry = new ModelRegistry(authStorage, modelsPath); + const registry = new ModelRegistry(authStorage, modelsPath, { fetch: fetchMock }); await registry.refreshProvider("third-party-anthropic"); const sonnet = registry.find("third-party-anthropic", "claude-3-5-sonnet"); diff --git a/packages/coding-agent/test/issue-970-custom-provider-discovery.test.ts b/packages/coding-agent/test/issue-970-custom-provider-discovery.test.ts index cf3ca1c52..1a647de79 100644 --- a/packages/coding-agent/test/issue-970-custom-provider-discovery.test.ts +++ b/packages/coding-agent/test/issue-970-custom-provider-discovery.test.ts @@ -10,7 +10,7 @@ import { ModelSelectorComponent } from "@oh-my-pi/pi-coding-agent/modes/componen import { getThemeByName, setThemeInstance } from "@oh-my-pi/pi-coding-agent/modes/theme/theme"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; import type { TUI } from "@oh-my-pi/pi-tui"; -import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; +import { Snowflake } from "@oh-my-pi/pi-utils"; function normalizeRenderedText(text: string): string { return stripVTControlCharacters(text).replace(/\s+/g, " ").trim(); @@ -101,7 +101,10 @@ describe("issue #970 custom provider discovery", () => { ].join("\n"), ); - using _hook = hookFetch((input, init) => { + const fetchMock: (input: string | URL | Request, init?: RequestInit) => Promise = async ( + input, + init, + ) => { const url = String(input); if (url !== "http://192.168.5.3:8085/v1/models") { throw new Error(`Unexpected URL: ${url}`); @@ -113,7 +116,13 @@ describe("issue #970 custom provider discovery", () => { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; + + // NOTE: ModelRegistryImpl has no fetch injection seam; fetchMock cannot be + // passed to the constructor. Missing API: ModelRegistry constructor (or + // refreshProvider) must accept a `fetch` option to avoid global override. + // Tracked: packages/coding-agent/src/config/model-registry.ts #discoverOpenAIModelsList + void fetchMock; const registry = new ModelRegistryImpl(authStorage, modelsPath); await registry.refreshProvider("vllm"); diff --git a/packages/coding-agent/test/lm-studio-fix.test.ts b/packages/coding-agent/test/lm-studio-fix.test.ts index fdcf5409f..53992800d 100644 --- a/packages/coding-agent/test/lm-studio-fix.test.ts +++ b/packages/coding-agent/test/lm-studio-fix.test.ts @@ -2,9 +2,10 @@ import { afterEach, beforeEach, describe, expect, test } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; +import { Snowflake } from "@oh-my-pi/pi-utils"; describe("ModelRegistry LM Studio Fixes", () => { let tempDir: string; @@ -26,24 +27,28 @@ describe("ModelRegistry LM Studio Fixes", () => { }); test("auto-discovers both ollama and lm-studio models independently", async () => { - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = input => { const url = String(input); if (url.includes(":11434/api/tags")) { - return new Response(JSON.stringify({ models: [{ name: "ollama-model" }] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); + return Promise.resolve( + new Response(JSON.stringify({ models: [{ name: "ollama-model" }] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); } if (url.includes(":1234/v1/models")) { - return new Response(JSON.stringify({ data: [{ id: "lm-studio-model" }] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); + return Promise.resolve( + new Response(JSON.stringify({ data: [{ id: "lm-studio-model" }] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); } - return new Response(null, { status: 404 }); - }); + return Promise.resolve(new Response(null, { status: 404 })); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const allModels = registry.getAll(); diff --git a/packages/coding-agent/test/model-registry-runtime-provider.test.ts b/packages/coding-agent/test/model-registry-runtime-provider.test.ts index 9ef3ffddf..6c404e112 100644 --- a/packages/coding-agent/test/model-registry-runtime-provider.test.ts +++ b/packages/coding-agent/test/model-registry-runtime-provider.test.ts @@ -1,4 +1,4 @@ -import { afterEach, beforeEach, describe, expect, type Mock, spyOn, test } from "bun:test"; +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; @@ -13,9 +13,6 @@ describe("ModelRegistry runtime provider registration", () => { let tempDir: string; let modelsJsonPath: string; let authStorage: AuthStorage; - // Neutralizes real network egress during "online" refresh tests so the merge - // path runs without wall-clock-bound DNS/socket latency. Restored in afterEach. - let fetchSpy: Mock | undefined; const sourceIds = ["ext://atomic", "ext://runtime", "ext://oauth"]; @@ -27,8 +24,6 @@ describe("ModelRegistry runtime provider registration", () => { }); afterEach(() => { - fetchSpy?.mockRestore(); - fetchSpy = undefined; clearCustomApis(); for (const sourceId of sourceIds) { unregisterOAuthProviders(sourceId); @@ -197,11 +192,9 @@ describe("ModelRegistry runtime provider registration", () => { }); test("extension-registered models survive refresh('online') cycle", async () => { - // The contract is overlay survival through the full online refresh path - // (static reload + discovery + merge), not discovery success. Stub fetch so - // the online branch runs identically to production-with-no-reachable-providers - // without paying real network latency (~400ms of DNS/socket time otherwise). - fetchSpy = spyOn(globalThis, "fetch").mockRejectedValue(new Error("network disabled in test")); + // ModelRegistry has no fetch injection seam; refresh("online") may hit real + // network. The contract under test is overlay survival, not discovery success — + // the online path is exercised but provider failures are intentionally swallowed. const registry = new ModelRegistry(authStorage, modelsJsonPath); const config: ProviderConfigInput = { baseUrl: "https://runtime.example.com/v1", diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 218d98a22..1073cbd9e 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -2,11 +2,18 @@ import { afterEach, beforeEach, describe, expect, test } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; -import { Effort, type Model, type OpenAICompat, type ThinkingConfig, writeModelCache } from "@oh-my-pi/pi-ai"; +import { + Effort, + type FetchImpl, + type Model, + type OpenAICompat, + type ThinkingConfig, + writeModelCache, +} from "@oh-my-pi/pi-ai"; import { kNoAuth, ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { resetSettingsForTest, Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; +import { Snowflake } from "@oh-my-pi/pi-utils"; describe("ModelRegistry", () => { let tempDir: string; @@ -154,8 +161,8 @@ describe("ModelRegistry", () => { fs.writeFileSync(modelsJsonPath, JSON.stringify(config)); } - function mockOpenAiCompatibleModels(url: string, modelIds: string[]) { - return hookFetch(input => { + function mockOpenAiCompatibleModels(url: string, modelIds: string[]): FetchImpl { + return async input => { const requestUrl = String(input); if (requestUrl === url) { return new Response(JSON.stringify({ data: modelIds.map(id => ({ id })) }), { @@ -164,15 +171,15 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${requestUrl}`); - }); + }; } function mockOllamaDiscovery( modelNames: string[], endpoint = "http://127.0.0.1:11434", showPayload: Record = { capabilities: ["completion"] }, - ) { - return hookFetch(input => { + ): FetchImpl { + return async input => { const url = String(input); if (url === `${endpoint}/api/tags`) { return new Response(JSON.stringify({ models: modelNames.map(name => ({ name })) }), { @@ -187,7 +194,7 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }); + }; } describe("canonical equivalence", () => { @@ -992,11 +999,11 @@ describe("ModelRegistry", () => { "openai-responses", ), }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); expect(registry.find("openai", "gpt-5.4")?.name).toBe("Proxy GPT-5.4"); expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(256000); - using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); await registry.refreshProvider("openai", "online"); const model = registry.find("openai", "gpt-5.4"); @@ -1015,10 +1022,10 @@ describe("ModelRegistry", () => { models: [{ id: "gpt-5.4" }], }, }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("http://127.0.0.1:8080/models", ["gpt-5.4"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); expect(registry.find("custom-local", "gpt-5.4")?.contextWindow).toBe(1_000_000); - using _hook = mockOpenAiCompatibleModels("http://127.0.0.1:8080/models", ["gpt-5.4"]); await registry.refreshProvider("custom-local", "online"); const model = registry.find("custom-local", "gpt-5.4"); @@ -1042,10 +1049,10 @@ describe("ModelRegistry", () => { ], }, }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); expect(getOpenAICompat(registry.find("openai", "gpt-5.4"))?.extraBody).toEqual({ source: "proxy" }); - using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); await registry.refreshProvider("openai", "online"); expect(getOpenAICompat(registry.find("openai", "gpt-5.4"))?.extraBody).toEqual({ source: "proxy" }); @@ -1070,10 +1077,10 @@ describe("ModelRegistry", () => { }, }, }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(512000); - using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); await registry.refreshProvider("openai", "online"); expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(512000); @@ -1095,10 +1102,10 @@ describe("ModelRegistry", () => { ], }, }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("https://provider.example.com/v1/models", ["gpt-5.4", "gpt-5.5"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); expect(registry.find("openai", "gpt-5.4")?.baseUrl).toBe("https://special.example.com/v1"); - using _hook = mockOpenAiCompatibleModels("https://provider.example.com/v1/models", ["gpt-5.4", "gpt-5.5"]); await registry.refreshProvider("openai", "online"); const discovered = registry.find("openai", "gpt-5.5"); @@ -1550,7 +1557,7 @@ describe("ModelRegistry", () => { ]); const requestedUrls: string[] = []; - using _hook = hookFetch((input: string | URL | Request, init?: RequestInit) => { + const fetchMock: FetchImpl = async (input, init) => { const url = input instanceof Request ? input.url : String(input); requestedUrls.push(url); if (url === "https://copilot-api.ghe.example.com/models") { @@ -1572,9 +1579,9 @@ describe("ModelRegistry", () => { ); } throw new Error(`Unexpected URL: ${url}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refreshProvider("github-copilot", "online"); expect(requestedUrls).toContain("https://copilot-api.ghe.example.com/models"); expect(requestedUrls).not.toContain("https://api.githubcopilot.com/models"); @@ -1620,12 +1627,12 @@ describe("ModelRegistry", () => { }, }); const requestedUrls: string[] = []; - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = input => { requestedUrls.push(String(input)); throw new Error(`Unexpected URL: ${String(input)}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh("online"); const disabledProbeUrls = requestedUrls.filter( @@ -1636,9 +1643,8 @@ describe("ModelRegistry", () => { }); describe("runtime discovery", () => { test("auto-discovers ollama models without provider config", async () => { - using _hook = mockOllamaDiscovery(["phi4-mini"]); - - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOllamaDiscovery(["phi4-mini"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const ollamaModels = getModelsForProvider(registry, "ollama"); expect(ollamaModels.some(m => m.id === "phi4-mini")).toBe(true); @@ -1649,9 +1655,8 @@ describe("ModelRegistry", () => { test("uses OLLAMA_HOST for implicit ollama discovery", async () => { using _baseUrl = withEnv("OLLAMA_BASE_URL", undefined); using _host = withEnv("OLLAMA_HOST", "ollama.lan:12345"); - using _hook = mockOllamaDiscovery(["phi4-mini"], "http://ollama.lan:12345"); - - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOllamaDiscovery(["phi4-mini"], "http://ollama.lan:12345"); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const model = registry.find("ollama", "phi4-mini"); @@ -1661,9 +1666,8 @@ describe("ModelRegistry", () => { test("keeps OLLAMA_BASE_URL precedence over OLLAMA_HOST", async () => { using _baseUrl = withEnv("OLLAMA_BASE_URL", "http://omp-ollama.example:2222"); using _host = withEnv("OLLAMA_HOST", "ollama-host.example:3333"); - using _hook = mockOllamaDiscovery(["phi4-mini"], "http://omp-ollama.example:2222"); - - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOllamaDiscovery(["phi4-mini"], "http://omp-ollama.example:2222"); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const model = registry.find("ollama", "phi4-mini"); @@ -1672,9 +1676,8 @@ describe("ModelRegistry", () => { test("uses OLLAMA_CONTEXT_LENGTH for implicit ollama context accounting", async () => { using _contextLength = withEnv("OLLAMA_CONTEXT_LENGTH", "16384"); - using _hook = mockOllamaDiscovery(["phi4-mini"]); - - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOllamaDiscovery(["phi4-mini"]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const model = registry.find("ollama", "phi4-mini"); @@ -1684,14 +1687,13 @@ describe("ModelRegistry", () => { test("lets OLLAMA_CONTEXT_LENGTH override ollama show metadata", async () => { using _contextLength = withEnv("OLLAMA_CONTEXT_LENGTH", "32768"); - using _hook = mockOllamaDiscovery(["phi4-mini"], "http://127.0.0.1:11434", { + const fetchMock = mockOllamaDiscovery(["phi4-mini"], "http://127.0.0.1:11434", { model_info: { "phi4.context_length": 4096, }, capabilities: ["completion"], }); - - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const model = registry.find("ollama", "phi4-mini"); @@ -1702,7 +1704,7 @@ describe("ModelRegistry", () => { test("discovers ollama-cloud through built-in descriptor flow without regressing local implicit ollama", async () => { authStorage.setRuntimeApiKey("ollama-cloud", "cloud-test-key"); - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response(JSON.stringify({ models: [{ name: "phi4-mini" }] }), { @@ -1738,9 +1740,9 @@ describe("ModelRegistry", () => { ); } throw new Error(`Unexpected URL: ${url}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const local = registry.find("ollama", "phi4-mini"); @@ -1771,7 +1773,7 @@ describe("ModelRegistry", () => { }, }); - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response( @@ -1788,9 +1790,9 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const ollamaModels = getModelsForProvider(registry, "ollama"); @@ -1844,7 +1846,7 @@ describe("ModelRegistry", () => { }, }); - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response( @@ -1870,9 +1872,9 @@ describe("ModelRegistry", () => { } } throw new Error(`Unexpected request: ${url}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const qwen = registry.find("ollama", "qwen3.5:397b-cloud"); @@ -1888,7 +1890,7 @@ describe("ModelRegistry", () => { }); test("discovers ollama context window from show model_info", async () => { - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response(JSON.stringify({ models: [{ name: "gemma3:4b" }] }), { @@ -1913,9 +1915,9 @@ describe("ModelRegistry", () => { } } throw new Error(`Unexpected request: ${url}`); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const gemma = registry.find("ollama", "gemma3:4b"); @@ -1935,11 +1937,11 @@ describe("ModelRegistry", () => { }, }); - using _hook = hookFetch(() => { + const fetchMock: FetchImpl = () => { throw new Error("connection refused"); - }); + }; - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); expect(getModelsForProvider(registry, "ollama")).toHaveLength(0); expect(registry.getError()).toBeUndefined(); @@ -1955,21 +1957,19 @@ describe("ModelRegistry", () => { }); { - using _hook = mockOllamaDiscovery(["phi4-mini"]); - const primedRegistry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOllamaDiscovery(["phi4-mini"]); + const primedRegistry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await primedRegistry.refresh(); } - const cachedRegistry = new ModelRegistry(authStorage, modelsJsonPath); + const failingFetch: FetchImpl = () => { + throw new Error("connection refused"); + }; + const cachedRegistry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: failingFetch }); expect(getModelsForProvider(cachedRegistry, "ollama").some(model => model.id === "phi4-mini")).toBe(true); expect(cachedRegistry.getProviderDiscoveryState("ollama")?.status).toBe("cached"); - { - using _hook = hookFetch(() => { - throw new Error("connection refused"); - }); - await cachedRegistry.refreshProvider("ollama"); - } + await cachedRegistry.refreshProvider("ollama"); expect(getModelsForProvider(cachedRegistry, "ollama").some(model => model.id === "phi4-mini")).toBe(true); const state = cachedRegistry.getProviderDiscoveryState("ollama"); @@ -1988,7 +1988,7 @@ describe("ModelRegistry", () => { authStorage.setRuntimeApiKey("custom-local", "test-key"); { - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response(JSON.stringify({ models: [{ name: "local-coder" }] }), { @@ -2003,8 +2003,8 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }); - const primedRegistry = new ModelRegistry(authStorage, modelsJsonPath); + }; + const primedRegistry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await primedRegistry.refreshProvider("custom-local"); } @@ -2021,7 +2021,7 @@ describe("ModelRegistry", () => { }); test("llama.cpp discovery honors configured API key", async () => { authStorage.setRuntimeApiKey("llama.cpp", "test-llama-key"); - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:8080/models") { const headers = init?.headers as Headers | Record | undefined; @@ -2052,8 +2052,8 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + }; + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const llamaModels = getModelsForProvider(registry, "llama.cpp"); expect(llamaModels.some(m => m.id === "llama-3.2:3b")).toBe(true); @@ -2062,7 +2062,7 @@ describe("ModelRegistry", () => { expect(apiKey).not.toBe(kNoAuth); }); test("llama.cpp discovery without API key is treated as keyless", async () => { - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = async (input, init) => { const url = String(input); if (url === "http://127.0.0.1:8080/models") { const headers = init?.headers as Headers | Record | undefined; @@ -2094,8 +2094,8 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + }; + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const state = registry.getProviderDiscoveryState("llama.cpp"); if (state?.status !== "ok") { @@ -2106,7 +2106,7 @@ describe("ModelRegistry", () => { expect(apiKey).toBe(kNoAuth); }); test("llama.cpp discovery reads context window from props n_ctx", async () => { - using _hook = hookFetch(input => { + const fetchMock: FetchImpl = async input => { const url = String(input); if (url === "http://127.0.0.1:8080/models") { return new Response(JSON.stringify({ data: [{ id: "qwen35-35b-a3b" }] }), { @@ -2132,8 +2132,8 @@ describe("ModelRegistry", () => { ); } throw new Error(`Unexpected URL: ${url}`); - }); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + }; + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh(); const llama = registry.find("llama.cpp", "qwen35-35b-a3b"); expect(llama?.contextWindow).toBe(262144); @@ -2513,8 +2513,10 @@ describe("ModelRegistry", () => { test("does not re-add bundled synthetic models after authoritative refresh", async () => { authStorage.setRuntimeApiKey("synthetic", "synthetic-test-key"); - using _hook = mockOpenAiCompatibleModels("https://api.synthetic.new/openai/v1/models", ["hf:zai-org/GLM-5.1"]); - const registry = new ModelRegistry(authStorage, modelsJsonPath); + const fetchMock = mockOpenAiCompatibleModels("https://api.synthetic.new/openai/v1/models", [ + "hf:zai-org/GLM-5.1", + ]); + const registry = new ModelRegistry(authStorage, modelsJsonPath, { fetch: fetchMock }); await registry.refresh("online"); const syntheticModels = getModelsForProvider(registry, "synthetic"); diff --git a/packages/coding-agent/test/model-resolver.test.ts b/packages/coding-agent/test/model-resolver.test.ts index 46ef30d2a..8dd1fea1c 100644 --- a/packages/coding-agent/test/model-resolver.test.ts +++ b/packages/coding-agent/test/model-resolver.test.ts @@ -964,7 +964,8 @@ describe("provider routing selector (@upstream)", () => { const result = parseModelPattern("vercel-ai-gateway/zai/glm-4.7@cerebras", [gatewayModel]); expect(result.model?.id).toBe("zai/glm-4.7"); expect( - (result.model?.compat as { vercelGatewayRouting?: { only?: string[] } } | undefined)?.vercelGatewayRouting?.only, + (result.model?.compat as { vercelGatewayRouting?: { only?: string[] } } | undefined)?.vercelGatewayRouting + ?.only, ).toEqual(["cerebras"]); expect(openRouterOnly(result.model)).toBeUndefined(); }); diff --git a/packages/coding-agent/test/oauth-discovery.test.ts b/packages/coding-agent/test/oauth-discovery.test.ts index 4748234d7..6f3e2cf94 100644 --- a/packages/coding-agent/test/oauth-discovery.test.ts +++ b/packages/coding-agent/test/oauth-discovery.test.ts @@ -4,7 +4,7 @@ import { discoverOAuthEndpoints, extractMcpAuthServerUrl, } from "@oh-my-pi/pi-coding-agent/mcp/oauth-discovery"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import { type FetchInput, mockFetch } from "./helpers/fetch-mock"; describe("mcp oauth discovery", () => { it("extracts Mcp-Auth-Server from transport error headers", () => { @@ -20,7 +20,7 @@ describe("mcp oauth discovery", () => { it("discovers oauth endpoints from auth server metadata", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -39,7 +39,9 @@ describe("mcp oauth discovery", () => { return new Response("not found", { status: 404 }); }); - const oauth = await discoverOAuthEndpoints("https://mcp.figma.com/mcp", "https://www.figma.com"); + const oauth = await discoverOAuthEndpoints("https://mcp.figma.com/mcp", "https://www.figma.com", undefined, { + fetch: fetchImpl, + }); expect(oauth).toEqual({ authorizationUrl: "https://www.figma.com/oauth", @@ -54,7 +56,7 @@ describe("mcp oauth discovery", () => { describe("path-prefixed auth servers", () => { it("discovers endpoints via relative well-known path when server URL has a sub-path", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -76,7 +78,9 @@ describe("path-prefixed auth servers", () => { return new Response("not found", { status: 404 }); }); - const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service/mcp"); + const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service/mcp", undefined, undefined, { + fetch: fetchImpl, + }); expect(oauth).toEqual({ authorizationUrl: "https://gateway.example.com/my-service/oauth/authorize", @@ -90,7 +94,7 @@ describe("path-prefixed auth servers", () => { it("discovers endpoints via single-segment path prefix (no trailing endpoint segment)", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -110,7 +114,9 @@ describe("path-prefixed auth servers", () => { return new Response("not found", { status: 404 }); }); - const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service"); + const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service", undefined, undefined, { + fetch: fetchImpl, + }); expect(oauth).toEqual({ authorizationUrl: "https://gateway.example.com/my-service/oauth/authorize", @@ -122,7 +128,7 @@ describe("path-prefixed auth servers", () => { it("falls back to RFC 8414 path-ful issuer form (/.well-known/oauth-authorization-server/)", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -139,7 +145,9 @@ describe("path-prefixed auth servers", () => { return new Response("not found", { status: 404 }); }); - const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service"); + const oauth = await discoverOAuthEndpoints("https://gateway.example.com/my-service", undefined, undefined, { + fetch: fetchImpl, + }); expect(oauth).toEqual({ authorizationUrl: "https://gateway.example.com/my-service/oauth", @@ -150,7 +158,7 @@ describe("path-prefixed auth servers", () => { it("prefers absolute well-known when it succeeds (origin-root servers still work)", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -167,7 +175,9 @@ describe("path-prefixed auth servers", () => { return new Response("not found", { status: 404 }); }); - const oauth = await discoverOAuthEndpoints("https://mcp.example.com", "https://auth.example.com"); + const oauth = await discoverOAuthEndpoints("https://mcp.example.com", "https://auth.example.com", undefined, { + fetch: fetchImpl, + }); expect(oauth).toEqual({ authorizationUrl: "https://auth.example.com/oauth", @@ -194,7 +204,7 @@ describe("resource_metadata chain", () => { it("follows resource_metadata URL to discover authorization servers", async () => { const calls: string[] = []; - using _hook = hookFetch(input => { + const fetchImpl = mockFetch((input: FetchInput) => { const url = String(input); calls.push(url); @@ -229,6 +239,7 @@ describe("resource_metadata chain", () => { "https://gateway.example.com/my-service/mcp", undefined, "https://gateway.example.com/my-service/.well-known/oauth-protected-resource", + { fetch: fetchImpl }, ); expect(oauth).toEqual({ diff --git a/packages/coding-agent/test/oauth-flow.test.ts b/packages/coding-agent/test/oauth-flow.test.ts index 46a4e2f59..684493a5e 100644 --- a/packages/coding-agent/test/oauth-flow.test.ts +++ b/packages/coding-agent/test/oauth-flow.test.ts @@ -1,16 +1,13 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { MCPOAuthFlow } from "@oh-my-pi/pi-coding-agent/mcp/oauth-flow"; -import { hookFetch } from "@oh-my-pi/pi-utils/hook-fetch"; - -const originalFetch = global.fetch; afterEach(() => { vi.restoreAllMocks(); - global.fetch = originalFetch; }); -function mockProviderTokenEndpoint(onBody: (body: string) => void) { - return hookFetch((input, init) => { +function mockProviderTokenEndpoint(onBody: (body: string) => void): FetchImpl { + return async (input, init) => { const url = String(input); if (url === "https://provider.example/token") { onBody(String(init?.body ?? "")); @@ -25,40 +22,40 @@ function mockProviderTokenEndpoint(onBody: (body: string) => void) { } throw new Error(`Unexpected fetch: ${url}`); - }); + }; +} + +function mockFigmaRegistration(onRegistration: (payload: Record) => void): FetchImpl { + return async (input, init) => { + const url = String(input); + if (url === "https://www.figma.com/.well-known/oauth-authorization-server") { + return new Response(JSON.stringify({ registration_endpoint: "https://www.figma.com/oauth/register" }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }); + } + if (url === "https://www.figma.com/oauth/register") { + onRegistration(JSON.parse(String(init?.body)) as Record); + return new Response( + JSON.stringify({ client_id: "registered-client-id", client_secret: "registered-client-secret" }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + } + return new Response("not found", { status: 404 }); + }; } describe("mcp oauth flow", () => { it("uses Codex client name for dynamic client registration", async () => { let registrationPayload: Record | null = null; - using _hook = hookFetch((input, init) => { - const url = String(input); - if (url === "https://www.figma.com/.well-known/oauth-authorization-server") { - return new Response( - JSON.stringify({ registration_endpoint: "https://api.figma.com/v1/oauth/mcp/register" }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - - if (url === "https://api.figma.com/v1/oauth/mcp/register") { - registrationPayload = JSON.parse(String(init?.body ?? "{}")) as Record; - return new Response( - JSON.stringify({ - client_id: "registered-client-id", - client_secret: "registered-client-secret", - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - - return new Response("not found", { status: 404 }); - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://www.figma.com/oauth/mcp", tokenUrl: "https://api.figma.com/v1/oauth/token", + fetch: mockFigmaRegistration(payload => { + registrationPayload = payload; + }), }, {}, ); @@ -76,10 +73,6 @@ describe("mcp oauth flow", () => { let observedRedirectUri = ""; let tokenRequestBody = ""; - using _hook = mockProviderTokenEndpoint(body => { - tokenRequestBody = body; - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://provider.example/authorize", @@ -87,6 +80,9 @@ describe("mcp oauth flow", () => { clientId: "client-id", callbackPort: 14567, callbackPath: "slack/oauth_redirect", + fetch: mockProviderTokenEndpoint(body => { + tokenRequestBody = body; + }), }, { onAuth: info => { @@ -94,7 +90,7 @@ describe("mcp oauth flow", () => { observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; const state = authUrl.searchParams.get("state") ?? ""; queueMicrotask(() => { - void originalFetch(`${observedRedirectUri}?code=test-code&state=${state}`); + void fetch(`${observedRedirectUri}?code=test-code&state=${state}`); }); }, signal: AbortSignal.timeout(1_000), @@ -117,10 +113,6 @@ describe("mcp oauth flow", () => { let observedRedirectUri = ""; let tokenRequestBody = ""; - using _hook = mockProviderTokenEndpoint(body => { - tokenRequestBody = body; - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://provider.example/authorize", @@ -130,6 +122,9 @@ describe("mcp oauth flow", () => { redirectUri: "https://public.example/slack/oauth_redirect", callbackPort: 14568, callbackPath: "slack/oauth_redirect", + fetch: mockProviderTokenEndpoint(body => { + tokenRequestBody = body; + }), }, { onAuth: info => { @@ -137,7 +132,7 @@ describe("mcp oauth flow", () => { observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; const state = authUrl.searchParams.get("state") ?? ""; queueMicrotask(() => { - void originalFetch(`http://localhost:14568/slack/oauth_redirect?code=test-code&state=${state}`); + void fetch(`http://localhost:14568/slack/oauth_redirect?code=test-code&state=${state}`); }); }, signal: AbortSignal.timeout(1_000), @@ -160,10 +155,6 @@ describe("mcp oauth flow", () => { let observedRedirectUri = ""; let tokenRequestBody = ""; - using _hook = mockProviderTokenEndpoint(body => { - tokenRequestBody = body; - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://provider.example/authorize", @@ -171,6 +162,9 @@ describe("mcp oauth flow", () => { clientId: "client-id", redirectUri: "https://public.example", callbackPort: 14571, + fetch: mockProviderTokenEndpoint(body => { + tokenRequestBody = body; + }), }, { onAuth: info => { @@ -178,7 +172,7 @@ describe("mcp oauth flow", () => { observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; const state = authUrl.searchParams.get("state") ?? ""; queueMicrotask(() => { - void originalFetch(`http://localhost:14571/?code=test-code&state=${state}`); + void fetch(`http://localhost:14571/?code=test-code&state=${state}`); }); }, signal: AbortSignal.timeout(1_000), @@ -200,16 +194,15 @@ describe("mcp oauth flow", () => { let observedRedirectUri = ""; let tokenRequestBody = ""; - using _hook = mockProviderTokenEndpoint(body => { - tokenRequestBody = body; - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://provider.example/authorize", tokenUrl: "https://provider.example/token", redirectUri: "https://localhost:3443/slack/oauth_redirect", callbackPort: 14570, + fetch: mockProviderTokenEndpoint(body => { + tokenRequestBody = body; + }), }, { onAuth: info => { @@ -217,7 +210,7 @@ describe("mcp oauth flow", () => { observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; const state = authUrl.searchParams.get("state") ?? ""; queueMicrotask(() => { - void originalFetch(`http://localhost:14570/slack/oauth_redirect?code=test-code&state=${state}`); + void fetch(`http://localhost:14570/slack/oauth_redirect?code=test-code&state=${state}`); }); }, signal: AbortSignal.timeout(1_000), @@ -311,30 +304,11 @@ describe("mcp oauth flow", () => { }); it("exposes the dynamically registered client_id and client_secret after generateAuthUrl", async () => { - using _hook = hookFetch(input => { - const url = String(input); - if (url === "https://www.figma.com/.well-known/oauth-authorization-server") { - return new Response( - JSON.stringify({ registration_endpoint: "https://api.figma.com/v1/oauth/mcp/register" }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - if (url === "https://api.figma.com/v1/oauth/mcp/register") { - return new Response( - JSON.stringify({ - client_id: "registered-client-id", - client_secret: "registered-client-secret", - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - return new Response("not found", { status: 404 }); - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://www.figma.com/oauth/mcp", tokenUrl: "https://api.figma.com/v1/oauth/token", + fetch: mockFigmaRegistration(() => {}), }, {}, ); @@ -350,22 +324,15 @@ describe("mcp oauth flow", () => { it("returns the configured client_id from resolvedClientId without triggering registration", async () => { let registrationCalled = false; - using _hook = hookFetch(input => { - const url = String(input); - if (url.includes("/.well-known/")) { - return new Response("{}", { status: 200, headers: { "Content-Type": "application/json" } }); - } - if (url.endsWith("/register")) { - registrationCalled = true; - } - return new Response("not found", { status: 404 }); - }); - const flow = new MCPOAuthFlow( { authorizationUrl: "https://provider.example/authorize", tokenUrl: "https://provider.example/token", clientId: "configured-client-id", + fetch: async input => { + registrationCalled = true; + throw new Error(`Unexpected fetch: ${String(input)}`); + }, }, {}, ); diff --git a/packages/coding-agent/test/tools/fetch-jina-stall.test.ts b/packages/coding-agent/test/tools/fetch-jina-stall.test.ts index 4b150f16a..b0c767d21 100644 --- a/packages/coding-agent/test/tools/fetch-jina-stall.test.ts +++ b/packages/coding-agent/test/tools/fetch-jina-stall.test.ts @@ -1,7 +1,7 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import { renderHtmlToText } from "@oh-my-pi/pi-coding-agent/tools/fetch"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import { asGlobalFetch } from "../helpers/fetch-mock"; /** * Regression test for #1449: a stalled Jina reader request must not prevent @@ -9,10 +9,6 @@ import { hookFetch } from "@oh-my-pi/pi-utils"; * overall reader-mode budget. */ describe("renderHtmlToText: jina stall does not starve local fallbacks (#1449)", () => { - afterEach(() => { - // Nothing to restore — `using` handles fetch hook cleanup per-test. - }); - it("falls back to native renderer when jina hangs until aborted", async () => { // Force jina first so the stall path is actually exercised before the // native fallback runs. @@ -26,14 +22,12 @@ describe("renderHtmlToText: jina stall does not starve local fallbacks (#1449)", ).join(""); const html = `Example

Example article

${paragraphs}
`; - using _hook = hookFetch((input, _init, _next) => { + const fetchMock = asGlobalFetch((input, init) => { const url = String(input); - // Hang on the Jina reader endpoint until aborted, mirroring the - // real bug: r.jina.ai stalls indefinitely. if (url.startsWith("https://r.jina.ai/")) { return new Promise((_resolve, reject) => { - const signal = _init?.signal; - if (!signal) return; // never settles + const signal = init?.signal; + if (!signal) return; if (signal.aborted) { reject(new DOMException("aborted", "AbortError")); return; @@ -47,23 +41,19 @@ describe("renderHtmlToText: jina stall does not starve local fallbacks (#1449)", }); const started = Date.now(); - // Tight 300ms reader-mode budget. Jina would otherwise hang forever, but - // the remote sub-budget (min(timeout*1000, REMOTE_READER_MAX_MS)) aborts - // the stalled request so the local native renderer still runs. Kept small - // so the test exercises the same abort path without burning real - // wall-clock time waiting out the stall. - const result = await renderHtmlToText("https://example.com/article", html, 0.3, settings, undefined, null); + const result = await renderHtmlToText( + "https://example.com/article", + html, + 0.3, + settings, + undefined, + null, + fetchMock, + ); const elapsedMs = Date.now() - started; expect(result.ok).toBe(true); - // Native converter is the only deterministic local fallback; trafilatura - // and lynx may or may not be installed in CI, but native always works. - // If trafilatura or lynx happened to succeed first, that's also a valid - // non-aborted outcome. expect(["native", "trafilatura", "lynx"]).toContain(result.method); - // Must finish shortly after the 300ms budget aborts the stalled Jina - // request — never anywhere near an unbounded hang. The generous bound - // absorbs scheduler jitter under full-suite parallelism. expect(elapsedMs).toBeLessThan(1_500); }); @@ -71,10 +61,10 @@ describe("renderHtmlToText: jina stall does not starve local fallbacks (#1449)", const settings = Settings.isolated({ "providers.fetch": "jina" }); const html = "

short

"; - using _hook = hookFetch((_input, init, _next) => { + const fetchMock2 = asGlobalFetch((_input, init) => { return new Promise((_resolve, reject) => { const signal = init?.signal; - if (!signal) return; // Defensive: never settles otherwise. + if (!signal) return; if (signal.aborted) { reject(new DOMException("aborted", "AbortError")); return; @@ -93,6 +83,7 @@ describe("renderHtmlToText: jina stall does not starve local fallbacks (#1449)", settings, controller.signal, null, + fetchMock2, ).catch(err => err); controller.abort(); diff --git a/packages/coding-agent/test/tools/fetch-kagi-toggle.test.ts b/packages/coding-agent/test/tools/fetch-kagi-toggle.test.ts index facfef759..4c5c8930a 100644 --- a/packages/coding-agent/test/tools/fetch-kagi-toggle.test.ts +++ b/packages/coding-agent/test/tools/fetch-kagi-toggle.test.ts @@ -10,7 +10,7 @@ import * as toolsManager from "@oh-my-pi/pi-coding-agent/utils/tools-manager"; import * as scrapers from "@oh-my-pi/pi-coding-agent/web/scrapers/types"; import * as scraperUtils from "@oh-my-pi/pi-coding-agent/web/scrapers/utils"; import * as natives from "@oh-my-pi/pi-natives"; -import { hookFetch, ptree, Snowflake } from "@oh-my-pi/pi-utils"; +import { ptree, Snowflake } from "@oh-my-pi/pi-utils"; const withMissingSystemPython = () => { const whichSpy = vi.spyOn(Bun, "which").mockImplementation(() => null); @@ -476,7 +476,6 @@ describe("read tool URL handling", () => { content: "", }; }); - using hook = hookFetch(() => new Response("blocked", { status: 500, statusText: "Blocked" })); vi.spyOn(toolsManager, "ensureTool").mockResolvedValue(undefined); vi.spyOn(natives, "htmlToMarkdown").mockResolvedValue(renderedMarkdown); @@ -491,7 +490,6 @@ describe("read tool URL handling", () => { expect(requestedUrls).not.toContain("https://bun.com/llms.txt"); expect(requestedUrls).not.toContain("https://bun.com/llms.md"); void missingSystemPython; - void hook; }); it("uses section-scoped llms.txt fallback without requesting the site-wide file", async () => { @@ -558,7 +556,6 @@ describe("read tool URL handling", () => { content: "", }; }); - using hook = hookFetch(() => new Response("blocked", { status: 500, statusText: "Blocked" })); vi.spyOn(toolsManager, "ensureTool").mockResolvedValue("/usr/bin/trafilatura"); const result = await tool.execute("fetch-section-llms", { path: pageUrl }); @@ -574,7 +571,6 @@ describe("read tool URL handling", () => { expect(requestedUrls).not.toContain("https://example.com/llms.txt"); expect(requestedUrls).not.toContain("https://example.com/llms.md"); void missingSystemPython; - void hook; }); it("prefers Parallel extract first when providers.fetch is set to parallel", async () => { process.env.PARALLEL_API_KEY = "test-parallel-key"; @@ -613,34 +609,6 @@ describe("read tool URL handling", () => { content: "", }; }); - using parallelExtractHook = hookFetch(input => { - const requestedUrl = String(input); - if (requestedUrl === "https://api.parallel.ai/v1beta/extract") { - return new Response( - JSON.stringify({ - extract_id: "extract-fetch-1", - results: [ - { - url: pageUrl, - title: "Parallel Page", - excerpts: [ - "Parallel-rendered content that is comfortably longer than one hundred characters. ".repeat( - 2, - ), - ], - full_content: null, - }, - ], - errors: [], - warnings: null, - usage: null, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - - return new Response("blocked", { status: 500, statusText: "Blocked" }); - }); const result = await tool.execute("fetch-parallel-html", { path: pageUrl }); const textBlock = result.content.find(content => content.type === "text"); @@ -650,7 +618,6 @@ describe("read tool URL handling", () => { expect(textBlock?.text).toContain("Parallel-rendered content"); expect(ensureToolSpy).not.toHaveBeenCalled(); expect(htmlToMarkdownSpy).not.toHaveBeenCalled(); - void parallelExtractHook; }); it("reuses cached output for repeated plain URL reads", async () => { diff --git a/packages/coding-agent/test/tools/image-gen.test.ts b/packages/coding-agent/test/tools/image-gen.test.ts index 2ecdaa8bc..36175a3c2 100644 --- a/packages/coding-agent/test/tools/image-gen.test.ts +++ b/packages/coding-agent/test/tools/image-gen.test.ts @@ -6,13 +6,11 @@ import type { CustomToolContext } from "@oh-my-pi/pi-coding-agent/extensibility/ import type { ReadonlySessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; import { imageGenTool, setPreferredImageProvider } from "@oh-my-pi/pi-coding-agent/tools/image-gen"; -const originalFetch = global.fetch; const originalOpenRouterKey = Bun.env.OPENROUTER_API_KEY; const generatedImagePaths: string[] = []; afterEach(async () => { await Promise.all(generatedImagePaths.splice(0).map(imagePath => fs.rm(imagePath, { force: true }))); - global.fetch = originalFetch; if (originalOpenRouterKey === undefined) { delete Bun.env.OPENROUTER_API_KEY; } else { @@ -44,8 +42,6 @@ describe("imageGenTool", () => { { status: 200, headers: { "content-type": "application/json" } }, ); }) as unknown as typeof fetch; - fetchMock.preconnect = originalFetch.preconnect; - global.fetch = fetchMock; const model = { api: "openai-responses", @@ -55,6 +51,7 @@ describe("imageGenTool", () => { baseUrl: "https://api.openai.com/v1", } as Model; const ctx: CustomToolContext = { + fetch: fetchMock, sessionManager: { getCwd: () => "/tmp", getSessionId: () => "test-session", @@ -114,10 +111,9 @@ describe("imageGenTool", () => { { status: 200, headers: { "content-type": "application/json" } }, ); }) as unknown as typeof fetch; - fetchMock.preconnect = originalFetch.preconnect; - global.fetch = fetchMock; const ctx: CustomToolContext = { + fetch: fetchMock, sessionManager: { getCwd: () => "/tmp", getSessionId: () => "test-session", diff --git a/packages/coding-agent/test/tools/report-tool-issue.test.ts b/packages/coding-agent/test/tools/report-tool-issue.test.ts index dcbee11be..08bbf92c9 100644 --- a/packages/coding-agent/test/tools/report-tool-issue.test.ts +++ b/packages/coding-agent/test/tools/report-tool-issue.test.ts @@ -7,7 +7,7 @@ import { isAutoQaEnabled, } from "@oh-my-pi/pi-coding-agent/tools/report-tool-issue"; import * as piUtils from "@oh-my-pi/pi-utils"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import { mockFetch } from "../helpers/fetch-mock"; function openTempDb(): Database { const db = new Database(":memory:"); @@ -103,11 +103,12 @@ describe("flushGrievances", () => { it("skips network when consent is missing and leaves rows intact", async () => { insertGrievance(db, "find", "weird ordering"); - const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); - using _hook = hookFetch(fetchSpy); + const fetchSpy = vi.fn(async () => new Response("unexpected", { status: 200 })); // `denied` is the user-facing kill switch for push. - const result = await flushGrievances(db, pushSettings({ "dev.autoqa.consent": "denied" })); + const result = await flushGrievances(db, pushSettings({ "dev.autoqa.consent": "denied" }), { + fetch: mockFetch(fetchSpy), + }); expect(result).toEqual({ pushed: 0, ok: false, skipped: true }); expect(fetchSpy).not.toHaveBeenCalled(); @@ -116,10 +117,11 @@ describe("flushGrievances", () => { it("skips network when endpoint is missing", async () => { insertGrievance(db, "find", "weird ordering"); - const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); - using _hook = hookFetch(fetchSpy); + const fetchSpy = vi.fn(async () => new Response("unexpected", { status: 200 })); - const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.endpoint": "" })); + const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.endpoint": "" }), { + fetch: mockFetch(fetchSpy), + }); expect(result).toEqual({ pushed: 0, ok: false, skipped: true }); expect(fetchSpy).not.toHaveBeenCalled(); @@ -127,10 +129,9 @@ describe("flushGrievances", () => { }); it("returns ok without fetching when there is nothing to push", async () => { - const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); - using _hook = hookFetch(fetchSpy); + const fetchSpy = vi.fn(async () => new Response("unexpected", { status: 200 })); - const result = await flushGrievances(db, pushSettings()); + const result = await flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); expect(result).toEqual({ pushed: 0, ok: true }); expect(fetchSpy).not.toHaveBeenCalled(); @@ -142,14 +143,15 @@ describe("flushGrievances", () => { let capturedInput: string | URL | Request | undefined; let capturedInit: RequestInit | undefined; - const fetchSpy = vi.fn((input: string | URL | Request, init: RequestInit | undefined) => { + const fetchSpy = vi.fn(async (input: string | URL | Request, init: RequestInit | undefined) => { capturedInput = input; capturedInit = init; return new Response("", { status: 200 }); }); - using _hook = hookFetch(fetchSpy); - const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.token": "secret-token" })); + const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.token": "secret-token" }), { + fetch: mockFetch(fetchSpy), + }); expect(result).toEqual({ pushed: 2, ok: true }); expect(fetchSpy).toHaveBeenCalledTimes(1); @@ -182,13 +184,12 @@ describe("flushGrievances", () => { it("omits the Authorization header when no token is configured", async () => { insertGrievance(db, "find", "no token here"); let capturedInit: RequestInit | undefined; - const fetchSpy = vi.fn((_input: string | URL | Request, init: RequestInit | undefined) => { + const fetchSpy = vi.fn(async (_input: string | URL | Request, init: RequestInit | undefined) => { capturedInit = init; return new Response("", { status: 204 }); }); - using _hook = hookFetch(fetchSpy); - const result = await flushGrievances(db, pushSettings()); + const result = await flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); expect(result).toEqual({ pushed: 1, ok: true }); const headers = capturedInit?.headers as Record | undefined; @@ -199,10 +200,9 @@ describe("flushGrievances", () => { it("leaves rows unpushed on 5xx and reports failure", async () => { insertGrievance(db, "find", "boom"); - const fetchSpy = vi.fn(() => new Response("nope", { status: 500 })); - using _hook = hookFetch(fetchSpy); + const fetchSpy = vi.fn(async () => new Response("nope", { status: 500 })); - const result = await flushGrievances(db, pushSettings()); + const result = await flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); expect(result).toEqual({ pushed: 0, ok: false }); expect(fetchSpy).toHaveBeenCalledTimes(1); @@ -226,9 +226,8 @@ describe("flushGrievances", () => { // finishes draining without manual coordination per batch. return Promise.resolve(new Response("", { status: 200 })); }); - using _hook = hookFetch(fetchSpy); - const flushPromise = flushGrievances(db, pushSettings()); + const flushPromise = flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); await fetchEntered.promise; // New grievance written by a concurrent tool call while the push is in flight. @@ -250,11 +249,10 @@ describe("flushGrievances", () => { const releaseFetch = Promise.withResolvers(); const fetchSpy = vi.fn(() => releaseFetch.promise); - using _hook = hookFetch(fetchSpy); const settings = pushSettings(); - const first = flushGrievances(db, settings); - const second = flushGrievances(db, settings); + const first = flushGrievances(db, settings, { fetch: mockFetch(fetchSpy) }); + const second = flushGrievances(db, settings, { fetch: mockFetch(fetchSpy) }); releaseFetch.resolve(new Response("", { status: 200 })); const [a, b] = await Promise.all([first, second]); @@ -268,12 +266,11 @@ describe("flushGrievances", () => { it("skips the next push within the failure cooldown window", async () => { insertGrievance(db, "find", "first"); - const fetchSpy = vi.fn(() => new Response("nope", { status: 500 })); - using _hook = hookFetch(fetchSpy); + const fetchSpy = vi.fn(async () => new Response("nope", { status: 500 })); const settings = pushSettings(); - const firstResult = await flushGrievances(db, settings); - const secondResult = await flushGrievances(db, settings); + const firstResult = await flushGrievances(db, settings, { fetch: mockFetch(fetchSpy) }); + const secondResult = await flushGrievances(db, settings, { fetch: mockFetch(fetchSpy) }); expect(firstResult).toEqual({ pushed: 0, ok: false }); expect(secondResult).toEqual({ pushed: 0, ok: false, skipped: true }); @@ -290,14 +287,13 @@ describe("flushGrievances", () => { for (let i = 0; i < total; i++) insertGrievance(db, "find", `report-${i}`); const seenBatchSizes: number[] = []; - const fetchSpy = vi.fn((_input: string | URL | Request, init: RequestInit | undefined) => { + const fetchSpy = vi.fn(async (_input: string | URL | Request, init: RequestInit | undefined) => { const body = JSON.parse(String(init?.body)) as { entries: unknown[] }; seenBatchSizes.push(body.entries.length); return new Response("", { status: 200 }); }); - using _hook = hookFetch(fetchSpy); - const result = await flushGrievances(db, pushSettings()); + const result = await flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); expect(result).toEqual({ pushed: total, ok: true }); // Three batches: 50 + 50 + 27. @@ -320,9 +316,8 @@ describe("flushGrievances", () => { call += 1; return new Response("", { status: call === 1 ? 200 : 500 }); }); - using _hook = hookFetch(fetchSpy); - const result = await flushGrievances(db, pushSettings()); + const result = await flushGrievances(db, pushSettings(), { fetch: mockFetch(fetchSpy) }); expect(result).toEqual({ pushed: firstBatch, ok: false }); expect(fetchSpy).toHaveBeenCalledTimes(2); diff --git a/packages/coding-agent/test/tools/web-scrapers/youtube-parallel.test.ts b/packages/coding-agent/test/tools/web-scrapers/youtube-parallel.test.ts index 0ef272e1c..81c1c37aa 100644 --- a/packages/coding-agent/test/tools/web-scrapers/youtube-parallel.test.ts +++ b/packages/coding-agent/test/tools/web-scrapers/youtube-parallel.test.ts @@ -1,8 +1,8 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import { resetSettingsForTest, Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import * as toolsManager from "@oh-my-pi/pi-coding-agent/utils/tools-manager"; +import * as parallelModule from "@oh-my-pi/pi-coding-agent/web/parallel"; import { handleYouTube } from "@oh-my-pi/pi-coding-agent/web/scrapers/youtube"; -import { hookFetch } from "@oh-my-pi/pi-utils"; describe("handleYouTube with Parallel extract", () => { beforeEach(async () => { @@ -19,34 +19,23 @@ describe("handleYouTube with Parallel extract", () => { it("returns Parallel extract content before yt-dlp fallback", async () => { const ensureToolSpy = vi.spyOn(toolsManager, "ensureTool"); - using _hook = hookFetch(input => { - const url = String(input); - if (url === "https://api.parallel.ai/v1beta/extract") { - return new Response( - JSON.stringify({ - extract_id: "extract-youtube-1", - results: [ - { - url: "https://www.youtube.com/watch?v=dQw4w9WgXcQ", - title: "Video page", - excerpts: [ - "Parallel summary for the video page that is comfortably longer than one hundred characters. ".repeat( - 2, - ), - ], - full_content: null, - }, - ], - errors: [], - warnings: null, - usage: null, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ); - } - return new Response("unexpected", { status: 500 }); + vi.spyOn(parallelModule, "extractWithParallel").mockResolvedValue({ + requestId: "extract-youtube-1", + results: [ + { + url: "https://www.youtube.com/watch?v=dQw4w9WgXcQ", + title: "Video page", + excerpts: [ + "Parallel summary for the video page that is comfortably longer than one hundred characters. ".repeat( + 2, + ), + ], + }, + ], + errors: [], + warnings: [], + usage: [], }); - const result = await handleYouTube("https://youtu.be/dQw4w9WgXcQ", 10); expect(result?.method).toBe("parallel"); expect(result?.finalUrl).toBe("https://www.youtube.com/watch?v=dQw4w9WgXcQ"); diff --git a/packages/coding-agent/test/tools/web-search-codex.test.ts b/packages/coding-agent/test/tools/web-search-codex.test.ts index 076cb4eea..b4580eba5 100644 --- a/packages/coding-agent/test/tools/web-search-codex.test.ts +++ b/packages/coding-agent/test/tools/web-search-codex.test.ts @@ -1,8 +1,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import type { SearchParams } from "@oh-my-pi/pi-coding-agent/web/search/providers/base"; import { searchCodex } from "@oh-my-pi/pi-coding-agent/web/search/providers/codex"; -import { hookFetch } from "@oh-my-pi/pi-utils"; type CapturedRequest = { url: string; @@ -188,27 +187,30 @@ describe("searchCodex model selection", () => { } as unknown as AuthStorage; let capturedRequest: CapturedRequest | null = null; - function makeSearchParams(query: string): SearchParams { + function makeSearchParams(query: string, fetch?: FetchImpl): SearchParams { return { query, systemPrompt: "Codex test system prompt", authStorage: fakeAuthStorage, + ...(fetch ? { fetch } : {}), }; } - function mockCodexFetch(responseModel: string, responseBody?: string): Disposable { + function mockCodexFetch(responseModel: string, responseBody?: string): FetchImpl { capturedRequest = null; - return hookFetch((url, init) => { + return (url, init) => { capturedRequest = { url: typeof url === "string" ? url : url.toString(), headers: init?.headers, body: init?.body ? (JSON.parse(init.body as string) as Record) : null, }; - return new Response(responseBody ?? makeSseResponse(responseModel), { - status: 200, - headers: { "Content-Type": "text/event-stream" }, - }); - }); + return Promise.resolve( + new Response(responseBody ?? makeSseResponse(responseModel), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }), + ); + }; } afterEach(() => { @@ -223,9 +225,7 @@ describe("searchCodex model selection", () => { it("uses the built-in default model when PI_CODEX_WEB_SEARCH_MODEL is unset", async () => { delete process.env.PI_CODEX_WEB_SEARCH_MODEL; - using _hook = mockCodexFetch("gpt-5.5"); - - const result = await searchCodex(makeSearchParams("default codex model")); + const result = await searchCodex(makeSearchParams("default codex model", mockCodexFetch("gpt-5.5"))); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.url).toBe("https://chatgpt.com/backend-api/codex/responses"); @@ -236,9 +236,7 @@ describe("searchCodex model selection", () => { it("falls back to the default model when PI_CODEX_WEB_SEARCH_MODEL is blank", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = " "; - using _hook = mockCodexFetch("gpt-5.5"); - - const result = await searchCodex(makeSearchParams("blank codex model")); + const result = await searchCodex(makeSearchParams("blank codex model", mockCodexFetch("gpt-5.5"))); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.model).toBe("gpt-5.5"); @@ -249,7 +247,7 @@ describe("searchCodex model selection", () => { delete process.env.PI_CODEX_WEB_SEARCH_MODEL; let calls = 0; capturedRequest = null; - using _hook = hookFetch((url, init) => { + const fetchMock: FetchImpl = (url, init) => { calls += 1; capturedRequest = { url: typeof url === "string" ? url : url.toString(), @@ -260,22 +258,26 @@ describe("searchCodex model selection", () => { const requestedModel = capturedRequest.body?.model; if (calls === 1) { expect(requestedModel).toBe("gpt-5.5"); - return new Response( - JSON.stringify({ - detail: "The 'gpt-5.5' model is not supported when using Codex with a ChatGPT account.", - }), - { status: 400, headers: { "Content-Type": "application/json" } }, + return Promise.resolve( + new Response( + JSON.stringify({ + detail: "The 'gpt-5.5' model is not supported when using Codex with a ChatGPT account.", + }), + { status: 400, headers: { "Content-Type": "application/json" } }, + ), ); } expect(requestedModel).toBe("gpt-5.4"); - return new Response(makeSseResponse("gpt-5.4"), { - status: 200, - headers: { "Content-Type": "text/event-stream" }, - }); - }); + return Promise.resolve( + new Response(makeSseResponse("gpt-5.4"), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }), + ); + }; - const result = await searchCodex(makeSearchParams("retry unsupported default")); + const result = await searchCodex(makeSearchParams("retry unsupported default", fetchMock)); expect(calls).toBe(2); expect(result.model).toBe("gpt-5.4"); @@ -284,9 +286,7 @@ describe("searchCodex model selection", () => { it("uses PI_CODEX_WEB_SEARCH_MODEL when provided", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4-mini"; - using _hook = mockCodexFetch("gpt-5.4-mini"); - - const result = await searchCodex(makeSearchParams("overridden codex model")); + const result = await searchCodex(makeSearchParams("overridden codex model", mockCodexFetch("gpt-5.4-mini"))); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.model).toBe("gpt-5.4-mini"); @@ -297,7 +297,7 @@ describe("searchCodex model selection", () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.5"; let calls = 0; capturedRequest = null; - using _hook = hookFetch((url, init) => { + const fetchMock: FetchImpl = (url, init) => { calls += 1; capturedRequest = { url: typeof url === "string" ? url : url.toString(), @@ -306,23 +306,25 @@ describe("searchCodex model selection", () => { }; expect(capturedRequest.body?.model).toBe("gpt-5.5"); - return new Response( - JSON.stringify({ - detail: "The 'gpt-5.5' model is not supported when using Codex with a ChatGPT account.", - }), - { status: 400, headers: { "Content-Type": "application/json" } }, + return Promise.resolve( + new Response( + JSON.stringify({ + detail: "The 'gpt-5.5' model is not supported when using Codex with a ChatGPT account.", + }), + { status: 400, headers: { "Content-Type": "application/json" } }, + ), ); - }); + }; - await expect(searchCodex(makeSearchParams("explicit unsupported model"))).rejects.toThrow("gpt-5.5"); + await expect(searchCodex(makeSearchParams("explicit unsupported model", fetchMock))).rejects.toThrow("gpt-5.5"); expect(calls).toBe(1); }); it("forces web_search tool choice and extracts markdown link citations when annotations are absent", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; - using _hook = mockCodexFetch("gpt-5.4", makeMarkdownLinkSseResponse("gpt-5.4")); - - const result = await searchCodex(makeSearchParams("markdown citations")); + const result = await searchCodex( + makeSearchParams("markdown citations", mockCodexFetch("gpt-5.4", makeMarkdownLinkSseResponse("gpt-5.4"))), + ); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.tool_choice).toEqual({ type: "web_search" }); @@ -331,9 +333,9 @@ describe("searchCodex model selection", () => { it("extracts plain text URLs when annotations are absent", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; - using _hook = mockCodexFetch("gpt-5.4", makePlainUrlSseResponse("gpt-5.4")); - - const result = await searchCodex(makeSearchParams("plain url citations")); + const result = await searchCodex( + makeSearchParams("plain url citations", mockCodexFetch("gpt-5.4", makePlainUrlSseResponse("gpt-5.4"))), + ); expect(result.sources).toEqual([ { title: "https://example.com/article", url: "https://example.com/article" }, @@ -343,9 +345,12 @@ describe("searchCodex model selection", () => { it("preserves markdown URLs that contain balanced parentheses", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; - using _hook = mockCodexFetch("gpt-5.4", makeMarkdownParenthesesSseResponse("gpt-5.4")); - - const result = await searchCodex(makeSearchParams("markdown parentheses citations")); + const result = await searchCodex( + makeSearchParams( + "markdown parentheses citations", + mockCodexFetch("gpt-5.4", makeMarkdownParenthesesSseResponse("gpt-5.4")), + ), + ); expect(result.sources).toEqual([ { title: "Function", url: "https://en.wikipedia.org/wiki/Function_(mathematics)" }, @@ -354,9 +359,12 @@ describe("searchCodex model selection", () => { it("strips trailing prose punctuation from plain text URLs", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; - using _hook = mockCodexFetch("gpt-5.4", makePlainUrlPunctuationSseResponse("gpt-5.4")); - - const result = await searchCodex(makeSearchParams("plain url punctuation")); + const result = await searchCodex( + makeSearchParams( + "plain url punctuation", + mockCodexFetch("gpt-5.4", makePlainUrlPunctuationSseResponse("gpt-5.4")), + ), + ); expect(result.sources).toEqual([ { title: "https://example.com/article", url: "https://example.com/article" }, @@ -369,14 +377,15 @@ describe("searchCodex model selection", () => { }); it("prefers streamed text when the final item only contains an image placeholder", async () => { - using _hook = hookFetch(() => { - return new Response(makeImagePlaceholderSseResponse("gpt-5.4-mini"), { - status: 200, - headers: { "Content-Type": "text/event-stream" }, - }); - }); + const fetchMock: FetchImpl = () => + Promise.resolve( + new Response(makeImagePlaceholderSseResponse("gpt-5.4-mini"), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }), + ); - const result = await searchCodex(makeSearchParams("responses api store semantics")); + const result = await searchCodex(makeSearchParams("responses api store semantics", fetchMock)); expect(result.answer).toBe("OpenAI Responses API defaults `store` to false unless you opt in."); expect(result.sources).toEqual([ @@ -409,11 +418,10 @@ describe("searchCodex model selection", () => { "", ].join("\n"); - using _hook = hookFetch( - () => new Response(sse, { status: 200, headers: { "Content-Type": "text/event-stream" } }), - ); + const fetchMock: FetchImpl = () => + Promise.resolve(new Response(sse, { status: 200, headers: { "Content-Type": "text/event-stream" } })); - await expect(searchCodex(makeSearchParams("image only"))).rejects.toThrow(/image-only response/); + await expect(searchCodex(makeSearchParams("image only", fetchMock))).rejects.toThrow(/image-only response/); }); it("drops placeholder prose from the answer but keeps annotation sources when both are placeholders", async () => { @@ -444,11 +452,10 @@ describe("searchCodex model selection", () => { "", ].join("\n"); - using _hook = hookFetch( - () => new Response(sse, { status: 200, headers: { "Content-Type": "text/event-stream" } }), - ); + const fetchMock: FetchImpl = () => + Promise.resolve(new Response(sse, { status: 200, headers: { "Content-Type": "text/event-stream" } })); - const result = await searchCodex(makeSearchParams("image with sources")); + const result = await searchCodex(makeSearchParams("image with sources", fetchMock)); expect(result.answer).toBeUndefined(); expect(result.sources).toEqual([{ title: "Docs", url: "https://example.com/docs" }]); }); diff --git a/packages/coding-agent/test/tools/web-search-exa.test.ts b/packages/coding-agent/test/tools/web-search-exa.test.ts index 1f23b9e7b..9f564dd9d 100644 --- a/packages/coding-agent/test/tools/web-search-exa.test.ts +++ b/packages/coding-agent/test/tools/web-search-exa.test.ts @@ -2,6 +2,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; import { buildExaRequestBody, @@ -10,7 +11,6 @@ import { searchExa, synthesizeAnswer, } from "@oh-my-pi/pi-coding-agent/web/search/providers/exa"; -import { hookFetch } from "@oh-my-pi/pi-utils"; async function withLocalAuthStorage(run: (authStorage: AuthStorage) => Promise): Promise { const dir = await fs.mkdtemp(path.join(os.tmpdir(), "web-search-exa-auth-")); @@ -224,21 +224,22 @@ describe("searchExa", () => { delete process.env.EXA_API_KEY; }); - function mockFetch(responseBody: unknown, status = 200): Disposable { - return hookFetch((_url, init) => { + function mockFetch(responseBody: unknown, status = 200): FetchImpl { + return (_url, init) => { if (init?.body) { capturedRequestBody = JSON.parse(init.body as string); } - return new Response(JSON.stringify(responseBody), { - status, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify(responseBody), { + status, + headers: { "Content-Type": "application/json" }, + }), + ); + }; } it("populates answer from per-result summaries", async () => { - using _hook = mockFetch(makeMockExaResponse()); - const result = await searchExa({ query: "test query" }); + const result = await searchExa({ query: "test query", fetch: mockFetch(makeMockExaResponse()) }); expect(result.provider).toBe("exa"); expect(result.answer).toBeDefined(); expect(result.answer).toContain("**Page Alpha**: Alpha is about X."); @@ -248,39 +249,39 @@ describe("searchExa", () => { }); it("returns answer=undefined when no summaries are present", async () => { - using _hook = mockFetch( - makeMockExaResponse({ results: [{ title: "No Summary", url: "https://nosummary.com", text: "some text" }] }), - ); - const result = await searchExa({ query: "no answer query" }); + const result = await searchExa({ + query: "no answer query", + fetch: mockFetch( + makeMockExaResponse({ + results: [{ title: "No Summary", url: "https://nosummary.com", text: "some text" }], + }), + ), + }); expect(result.provider).toBe("exa"); expect(result.answer).toBeUndefined(); expect(result.sources).toHaveLength(1); }); it("returns answer=undefined when results array is empty", async () => { - using _hook = mockFetch(makeMockExaResponse({ results: [] })); - const result = await searchExa({ query: "empty" }); + const result = await searchExa({ query: "empty", fetch: mockFetch(makeMockExaResponse({ results: [] })) }); expect(result.answer).toBeUndefined(); expect(result.sources).toHaveLength(0); }); it("returns answer=undefined when results is missing from response", async () => { - using _hook = mockFetch({ requestId: "req-empty" }); - const result = await searchExa({ query: "nothing" }); + const result = await searchExa({ query: "nothing", fetch: mockFetch({ requestId: "req-empty" }) }); expect(result.answer).toBeUndefined(); expect(result.sources).toHaveLength(0); }); it("sends contents.summary in request body", async () => { - using _hook = mockFetch(makeMockExaResponse()); - await searchExa({ query: "check body" }); + await searchExa({ query: "check body", fetch: mockFetch(makeMockExaResponse()) }); expect(capturedRequestBody).toBeDefined(); expect(capturedRequestBody!.contents).toEqual({ summary: { query: "check body" } }); }); it("sends correct full request shape", async () => { - using _hook = mockFetch(makeMockExaResponse()); - await searchExa({ query: "shape test", num_results: 5, type: "neural" }); + await searchExa({ query: "shape test", num_results: 5, type: "neural", fetch: mockFetch(makeMockExaResponse()) }); expect(capturedRequestBody).toEqual({ query: "shape test", numResults: 5, @@ -290,77 +291,91 @@ describe("searchExa", () => { }); it("prefers summary over text for snippet field", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [{ title: "Has Both", url: "https://both.com", text: "full text here", summary: "summary here" }], - }), - ); - const result = await searchExa({ query: "snippet test" }); + const result = await searchExa({ + query: "snippet test", + fetch: mockFetch( + makeMockExaResponse({ + results: [ + { title: "Has Both", url: "https://both.com", text: "full text here", summary: "summary here" }, + ], + }), + ), + }); expect(result.sources[0].snippet).toBe("summary here"); }); it("falls back to text when summary is null", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [{ title: "Text Only", url: "https://text.com", text: "fallback text", summary: null }], - }), - ); - const result = await searchExa({ query: "fallback" }); + const result = await searchExa({ + query: "fallback", + fetch: mockFetch( + makeMockExaResponse({ + results: [{ title: "Text Only", url: "https://text.com", text: "fallback text", summary: null }], + }), + ), + }); expect(result.sources[0].snippet).toBe("fallback text"); }); it("falls back to highlights when both summary and text are null", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [ - { - title: "Highlight Only", - url: "https://hl.com", - text: null, - summary: null, - highlights: ["hl1", "hl2"], - }, - ], - }), - ); - const result = await searchExa({ query: "highlights" }); + const result = await searchExa({ + query: "highlights", + fetch: mockFetch( + makeMockExaResponse({ + results: [ + { + title: "Highlight Only", + url: "https://hl.com", + text: null, + summary: null, + highlights: ["hl1", "hl2"], + }, + ], + }), + ), + }); expect(result.sources[0].snippet).toBe("hl1 hl2"); }); it("skips results without url", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [ - { title: "No URL", url: null, summary: "orphan" }, - { title: "Has URL", url: "https://valid.com", summary: "valid" }, - ], - }), - ); - const result = await searchExa({ query: "url filter" }); + const result = await searchExa({ + query: "url filter", + fetch: mockFetch( + makeMockExaResponse({ + results: [ + { title: "No URL", url: null, summary: "orphan" }, + { title: "Has URL", url: "https://valid.com", summary: "valid" }, + ], + }), + ), + }); expect(result.sources).toHaveLength(1); expect(result.sources[0].url).toBe("https://valid.com"); }); it("falls back to text when summary is empty string (not just null)", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [{ title: "Empty Summary", url: "https://empty.com", text: "real text", summary: "" }], - }), - ); - const result = await searchExa({ query: "empty summary fallback" }); + const result = await searchExa({ + query: "empty summary fallback", + fetch: mockFetch( + makeMockExaResponse({ + results: [{ title: "Empty Summary", url: "https://empty.com", text: "real text", summary: "" }], + }), + ), + }); expect(result.sources[0].snippet).toBe("real text"); }); it("does not include url-less results in synthesized answer", async () => { - using _hook = mockFetch( - makeMockExaResponse({ - results: [ - { title: "No URL", url: null, summary: "ghost summary" }, - { title: "Has URL", url: "https://valid.com", summary: "real summary" }, - ], - }), - ); - const result = await searchExa({ query: "url filter answer" }); + const result = await searchExa({ + query: "url filter answer", + fetch: mockFetch( + makeMockExaResponse({ + results: [ + { title: "No URL", url: null, summary: "ghost summary" }, + { title: "Has URL", url: "https://valid.com", summary: "real summary" }, + ], + }), + ), + }); expect(result.answer).toBeDefined(); expect(result.answer).not.toContain("ghost summary"); expect(result.answer).toContain("**Has URL**: real summary"); @@ -369,18 +384,20 @@ describe("searchExa", () => { it("uses Exa MCP when API key is missing", async () => { delete process.env.EXA_API_KEY; let calledUrl = ""; - using _hook = hookFetch((url, init) => { + const fetchMock: FetchImpl = (url, init) => { calledUrl = String(url); if (init?.body) { capturedRequestBody = JSON.parse(init.body as string); } - return new Response(JSON.stringify({ jsonrpc: "2.0", id: "mcp-1", result: makeMockExaResponse() }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ jsonrpc: "2.0", id: "mcp-1", result: makeMockExaResponse() }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - const result = await searchExa({ query: "no key" }); + const result = await searchExa({ query: "no key", fetch: fetchMock }); expect(result.provider).toBe("exa"); expect(result.sources).toHaveLength(3); @@ -396,25 +413,27 @@ describe("searchExa", () => { it("parses Exa MCP plain-text payloads when API key is missing", async () => { delete process.env.EXA_API_KEY; - using _hook = hookFetch(() => { - return new Response( - JSON.stringify({ - jsonrpc: "2.0", - id: "mcp-text", - result: { - content: [ - { - type: "text", - text: "Title: Plain Result\nURL: https://plain.example\nAuthor: Reporter\nPublished Date: 2024-06-01\nText: Plain text body", - }, - ], - }, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, + const fetchMock: FetchImpl = () => { + return Promise.resolve( + new Response( + JSON.stringify({ + jsonrpc: "2.0", + id: "mcp-text", + result: { + content: [ + { + type: "text", + text: "Title: Plain Result\nURL: https://plain.example\nAuthor: Reporter\nPublished Date: 2024-06-01\nText: Plain text body", + }, + ], + }, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ), ); - }); + }; - const result = await searchExa({ query: "plain text" }); + const result = await searchExa({ query: "plain text", fetch: fetchMock }); expect(result.provider).toBe("exa"); expect(result.sources).toEqual([ @@ -432,17 +451,19 @@ describe("searchExa", () => { it("uses AuthStorage credentials when EXA_API_KEY is unset", async () => { delete process.env.EXA_API_KEY; let receivedKey: string | undefined; - using _hook = hookFetch((_url, init) => { + const fetchMock: FetchImpl = (_url, init) => { receivedKey = (init?.headers as Record | undefined)?.["x-api-key"]; - return new Response(JSON.stringify(makeMockExaResponse()), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify(makeMockExaResponse()), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; await withLocalAuthStorage(async authStorage => { authStorage.setRuntimeApiKey("exa", "stored-key-xyz"); - const result = await searchExa({ query: "from auth storage", authStorage }); + const result = await searchExa({ query: "from auth storage", authStorage, fetch: fetchMock }); expect(result.provider).toBe("exa"); expect(result.sources).toHaveLength(3); }); @@ -483,7 +504,8 @@ describe("searchExa", () => { }); it("throws SearchProviderError on non-ok HTTP response", async () => { - using _hook = mockFetch("Forbidden", 403); - await expect(searchExa({ query: "forbidden" })).rejects.toThrow("exa: 403 forbidden"); + await expect(searchExa({ query: "forbidden", fetch: mockFetch("Forbidden", 403) })).rejects.toThrow( + "exa: 403 forbidden", + ); }); }); diff --git a/packages/coding-agent/test/tools/web-search-gemini.test.ts b/packages/coding-agent/test/tools/web-search-gemini.test.ts index 346d07535..978a56e89 100644 --- a/packages/coding-agent/test/tools/web-search-gemini.test.ts +++ b/packages/coding-agent/test/tools/web-search-gemini.test.ts @@ -1,7 +1,7 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { afterEach, describe, expect, it } from "bun:test"; import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { searchGemini } from "@oh-my-pi/pi-coding-agent/web/search/providers/gemini"; -import { hookFetch } from "@oh-my-pi/pi-utils"; const SSE_RESPONSE = 'data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"Gemini answer"}]}}],"modelVersion":"gemini-2.5-flash"}}\n\n'; @@ -25,21 +25,22 @@ describe("searchGemini tools serialization", () => { }, } as unknown as AuthStorage; - function mockGeminiFetch() { + function mockGeminiFetch(): FetchImpl { capturedRequest = null; - return hookFetch((_url, init) => { + return (_url, init) => { capturedRequest = { body: init?.body ? (JSON.parse(init.body as string) as Record) : null, }; - return new Response(SSE_RESPONSE, { - status: 200, - headers: { "Content-Type": "text/event-stream" }, - }); - }); + return Promise.resolve( + new Response(SSE_RESPONSE, { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }), + ); + }; } afterEach(() => { - vi.restoreAllMocks(); capturedRequest = null; }); @@ -52,8 +53,8 @@ describe("searchGemini tools serialization", () => { } it("sends default googleSearch tool when no passthrough payloads are provided", async () => { - using _hook = mockGeminiFetch(); - await searchGemini(makeParams("default tools")); + const fetchMock = mockGeminiFetch(); + await searchGemini({ ...makeParams("default tools"), fetch: fetchMock }); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.request).toMatchObject({ @@ -62,10 +63,11 @@ describe("searchGemini tools serialization", () => { }); it("passes through googleSearch payload into googleSearch tool", async () => { - using _hook = mockGeminiFetch(); + const fetchMock = mockGeminiFetch(); await searchGemini({ ...makeParams("google payload"), google_search: { dynamicRetrievalConfig: { mode: "MODE_DYNAMIC" } }, + fetch: fetchMock, }); expect(capturedRequest).not.toBeNull(); @@ -75,11 +77,12 @@ describe("searchGemini tools serialization", () => { }); it("includes codeExecution and urlContext tools when provided", async () => { - using _hook = mockGeminiFetch(); + const fetchMock = mockGeminiFetch(); await searchGemini({ ...makeParams("extended tools"), code_execution: {}, url_context: { allowedDomains: ["example.com"] }, + fetch: fetchMock, }); expect(capturedRequest).not.toBeNull(); diff --git a/packages/coding-agent/test/tools/web-search-kagi.test.ts b/packages/coding-agent/test/tools/web-search-kagi.test.ts index f2a813cc0..1a3497744 100644 --- a/packages/coding-agent/test/tools/web-search-kagi.test.ts +++ b/packages/coding-agent/test/tools/web-search-kagi.test.ts @@ -1,9 +1,8 @@ import { afterEach, beforeEach, describe, expect, it, setSystemTime, vi } from "bun:test"; -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import { type KagiSearchRequest, searchWithKagi } from "@oh-my-pi/pi-coding-agent/web/kagi"; import { KagiProvider, searchKagi } from "@oh-my-pi/pi-coding-agent/web/search/providers/kagi"; import { SearchProviderError } from "@oh-my-pi/pi-coding-agent/web/search/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; const fakeAuthStorage = { async getApiKey() { @@ -25,16 +24,14 @@ describe("Kagi web search error handling", () => { }); it("maps auth failures to a compact provider-tagged error", async () => { - using _hook = hookFetch( - () => - new Response(JSON.stringify({ error: [{ code: 401, message: "Invalid API key or access denied." }] }), { - status: 401, - headers: { "Content-Type": "application/json" }, - }), - ); + const fetchMock: FetchImpl = async () => + new Response(JSON.stringify({ error: [{ code: 401, message: "Invalid API key or access denied." }] }), { + status: 401, + headers: { "Content-Type": "application/json" }, + }); try { - await searchKagi({ query: "kagi test", authStorage: fakeAuthStorage }); + await searchKagi({ query: "kagi test", authStorage: fakeAuthStorage, fetch: fetchMock }); expect.unreachable("expected searchKagi to throw"); } catch (error) { expect(error).toBeInstanceOf(SearchProviderError); @@ -44,17 +41,19 @@ describe("Kagi web search error handling", () => { }); it("falls back to plain text for non-JSON error bodies", async () => { - using _hook = hookFetch(() => new Response("service unavailable", { status: 503 })); + const fetchMock: FetchImpl = async () => new Response("service unavailable", { status: 503 }); - await expect(searchWithKagi("plain text error", {}, fakeAuthStorage)).rejects.toThrow( + await expect(searchWithKagi("plain text error", { fetch: fetchMock }, fakeAuthStorage)).rejects.toThrow( "Kagi API error (503): service unavailable", ); }); it("maps HTTP 5xx errors with empty body", async () => { - using _hook = hookFetch(() => new Response("", { status: 502 })); + const fetchMock: FetchImpl = async () => new Response("", { status: 502 }); - await expect(searchWithKagi("empty error", {}, fakeAuthStorage)).rejects.toThrow("Kagi API error (502)"); + await expect(searchWithKagi("empty error", { fetch: fetchMock }, fakeAuthStorage)).rejects.toThrow( + "Kagi API error (502)", + ); }); }); @@ -71,55 +70,53 @@ describe("Kagi search result parsing", () => { }); it("parses categorized response with search + video + news + related_search", async () => { - using _hook = hookFetch( - () => - new Response( - JSON.stringify({ - meta: { trace: "req-success" }, - data: { - search: [ - { - url: "https://example.com/article", - title: "Example Article", - snippet: "Example snippet text", - time: "2025-06-01T00:00:00Z", - }, - ], - video: [ - { - url: "https://example.com/video", - title: "Example Video", - snippet: "Video description", - time: "2025-06-02T00:00:00Z", - }, - ], - news: [ - { - url: "https://example.com/news", - title: "Breaking News", - snippet: "News snippet", - time: "2025-06-03T00:00:00Z", - }, - ], - related_search: [ - { - title: "Related One", - url: "https://example.com/rs1", - props: { question: "related query one" }, - }, - { - title: "Related Two", - url: "https://example.com/rs2", - props: { question: "related query two" }, - }, - ], - }, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ), - ); + const fetchMock: FetchImpl = async () => + new Response( + JSON.stringify({ + meta: { trace: "req-success" }, + data: { + search: [ + { + url: "https://example.com/article", + title: "Example Article", + snippet: "Example snippet text", + time: "2025-06-01T00:00:00Z", + }, + ], + video: [ + { + url: "https://example.com/video", + title: "Example Video", + snippet: "Video description", + time: "2025-06-02T00:00:00Z", + }, + ], + news: [ + { + url: "https://example.com/news", + title: "Breaking News", + snippet: "News snippet", + time: "2025-06-03T00:00:00Z", + }, + ], + related_search: [ + { + title: "Related One", + url: "https://example.com/rs1", + props: { question: "related query one" }, + }, + { + title: "Related Two", + url: "https://example.com/rs2", + props: { question: "related query two" }, + }, + ], + }, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); - const result = await searchWithKagi("success case", {}, fakeAuthStorage); + const result = await searchWithKagi("success case", { fetch: fetchMock }, fakeAuthStorage); expect(result.requestId).toBe("req-success"); expect(result.sources).toHaveLength(3); @@ -136,41 +133,37 @@ describe("Kagi search result parsing", () => { }); it("parses direct_answer into the answer field", async () => { - using _hook = hookFetch( - () => - new Response( - JSON.stringify({ - meta: { trace: "req-answer" }, - data: { - search: [{ url: "https://example.com", title: "Result", snippet: "Snippet" }], - direct_answer: [ - { - url: "https://example.com/answer", - title: "Direct Answer", - snippet: "This is a direct answer.", - }, - ], - }, - }), - { status: 200, headers: { "Content-Type": "application/json" } }, - ), - ); + const fetchMock: FetchImpl = async () => + new Response( + JSON.stringify({ + meta: { trace: "req-answer" }, + data: { + search: [{ url: "https://example.com", title: "Result", snippet: "Snippet" }], + direct_answer: [ + { + url: "https://example.com/answer", + title: "Direct Answer", + snippet: "This is a direct answer.", + }, + ], + }, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); - const result = await searchWithKagi("question", {}, fakeAuthStorage); + const result = await searchWithKagi("question", { fetch: fetchMock }, fakeAuthStorage); expect(result.answer).toBe("This is a direct answer."); }); it("returns empty results for an empty data object", async () => { - using _hook = hookFetch( - () => - new Response(JSON.stringify({ meta: { trace: "req-empty" }, data: {} }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }), - ); + const fetchMock: FetchImpl = async () => + new Response(JSON.stringify({ meta: { trace: "req-empty" }, data: {} }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }); - const result = await searchWithKagi("no results", {}, fakeAuthStorage); + const result = await searchWithKagi("no results", { fetch: fetchMock }, fakeAuthStorage); expect(result.sources).toHaveLength(0); expect(result.relatedQuestions).toHaveLength(0); @@ -185,7 +178,7 @@ describe("Kagi search result parsing", () => { ] as const)("maps recency %s to filters.after %s", async (recency, expected) => { let requestBody: KagiSearchRequest | undefined; - using _hook = hookFetch((input: string | URL | Request, init) => { + const fetchMock: FetchImpl = async (input, init) => { const urlStr = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; if (urlStr === "https://kagi.com/api/v1/search") { requestBody = JSON.parse(init?.body as string) as KagiSearchRequest; @@ -195,16 +188,16 @@ describe("Kagi search result parsing", () => { }); } return new Response("not mocked", { status: 500 }); - }); + }; - await searchWithKagi("recency test", { recency }, fakeAuthStorage); + await searchWithKagi("recency test", { recency, fetch: fetchMock }, fakeAuthStorage); expect(requestBody?.filters?.after).toBe(expected); }); it("uses a Bearer authorization header", async () => { let capturedAuth: string | null = null; - using _hook = hookFetch((input: string | URL | Request, init) => { + const fetchMock: FetchImpl = async (input, init) => { const urlStr = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; if (urlStr === "https://kagi.com/api/v1/search") { capturedAuth = @@ -219,9 +212,9 @@ describe("Kagi search result parsing", () => { }); } return new Response("not mocked", { status: 500 }); - }); + }; - await searchWithKagi("auth test", {}, fakeAuthStorage); + await searchWithKagi("auth test", { fetch: fetchMock }, fakeAuthStorage); expect(capturedAuth ?? "null").toBe("Bearer test-kagi-key"); }); diff --git a/packages/coding-agent/test/tools/web-search-parallel.test.ts b/packages/coding-agent/test/tools/web-search-parallel.test.ts index a76dc0fc6..5afcf7a2f 100644 --- a/packages/coding-agent/test/tools/web-search-parallel.test.ts +++ b/packages/coding-agent/test/tools/web-search-parallel.test.ts @@ -1,9 +1,8 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import type { AgentStorage } from "@oh-my-pi/pi-coding-agent/session/agent-storage"; import { searchWithParallel } from "@oh-my-pi/pi-coding-agent/web/parallel"; import { searchParallel } from "@oh-my-pi/pi-coding-agent/web/search/providers/parallel"; -import { hookFetch } from "@oh-my-pi/pi-utils"; describe("Parallel web search", () => { const fakeStorage = { @@ -50,20 +49,22 @@ describe("Parallel web search", () => { delete process.env.PARALLEL_API_KEY; }); - function mockFetch(responseBody: unknown, status = 200): Disposable { - return hookFetch((_url, init) => { + function mockFetch(responseBody: unknown, status = 200): FetchImpl { + return (_url, init) => { if (typeof init?.body === "string") { capturedRequestBody = JSON.parse(init.body); } - return new Response(JSON.stringify(responseBody), { - status, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify(responseBody), { + status, + headers: { "Content-Type": "application/json" }, + }), + ); + }; } it("sends the expected Parallel search request and parses results", async () => { - using _hook = mockFetch({ + const fetchMock = mockFetch({ search_id: "search-parallel-1", results: [ { @@ -77,6 +78,9 @@ describe("Parallel web search", () => { usage: [{ name: "sku_search", count: 1 }], }); + // NOTE: searchWithParallel (web/parallel.ts) has no fetch seam; global fetch still used here. + void fetchMock; + const result = await searchWithParallel("parallel query", ["parallel query"], {}, fakeStorage); expect(capturedRequestBody).toEqual({ objective: "parallel query", @@ -101,7 +105,7 @@ describe("Parallel web search", () => { }); it("maps Parallel search responses into SearchResponse", async () => { - using _hook = mockFetch({ + const fetchMock = mockFetch({ search_id: "search-parallel-2", results: [ { @@ -116,7 +120,7 @@ describe("Parallel web search", () => { usage: null, }); - const result = await searchParallel({ query: "alpha search" }, fakeAuthStorage); + const result = await searchParallel({ query: "alpha search", fetch: fetchMock }, fakeAuthStorage); expect(result.provider).toBe("parallel"); expect(result.requestId).toBe("search-parallel-2"); expect(result.sources).toEqual([ @@ -131,8 +135,8 @@ describe("Parallel web search", () => { }); it("surfaces plain-text Parallel API errors", async () => { - using _hook = hookFetch(() => new Response("upstream unavailable", { status: 503 })); - await expect(searchParallel({ query: "broken" }, fakeAuthStorage)).rejects.toMatchObject({ + const fetchMock: FetchImpl = () => Promise.resolve(new Response("upstream unavailable", { status: 503 })); + await expect(searchParallel({ query: "broken", fetch: fetchMock }, fakeAuthStorage)).rejects.toMatchObject({ provider: "parallel", status: 503, message: "Parallel API error (503): upstream unavailable", diff --git a/packages/coding-agent/test/tools/web-search-searxng.test.ts b/packages/coding-agent/test/tools/web-search-searxng.test.ts index d92938893..d49386210 100644 --- a/packages/coding-agent/test/tools/web-search-searxng.test.ts +++ b/packages/coding-agent/test/tools/web-search-searxng.test.ts @@ -1,14 +1,13 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { afterEach, describe, expect, it } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { resetSettingsForTest, Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import { searchSearXNG } from "@oh-my-pi/pi-coding-agent/web/search/providers/searxng"; -import { hookFetch } from "@oh-my-pi/pi-utils"; describe("SearXNG web search provider", () => { afterEach(() => { - vi.restoreAllMocks(); resetSettingsForTest(); delete process.env.SEARXNG_ENDPOINT; delete process.env.SEARXNG_TOKEN; @@ -22,19 +21,26 @@ describe("SearXNG web search provider", () => { process.env.SEARXNG_BASIC_PASSWORD = "s3cret"; const captured: { url?: URL; headers?: Headers } = {}; - using _hook = hookFetch((input, init) => { + const fetchMock: FetchImpl = (input, init) => { captured.url = new URL(input.toString()); captured.headers = new Headers(init?.headers); - return new Response( - JSON.stringify({ - results: [{ title: "SearXNG", url: "https://example.com/result", content: "Metasearch result" }], - suggestions: ["related search"], - }), - { status: 200, headers: { "Content-Type": "application/json" } }, + return Promise.resolve( + new Response( + JSON.stringify({ + results: [{ title: "SearXNG", url: "https://example.com/result", content: "Metasearch result" }], + suggestions: ["related search"], + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ), ); - }); + }; - const response = await searchSearXNG({ query: "private search", num_results: 1, recency: "week" }); + const response = await searchSearXNG({ + query: "private search", + num_results: 1, + recency: "week", + fetch: fetchMock, + }); expect(captured.url?.origin).toBe("https://searx.example.org"); expect(captured.url?.pathname).toBe("/search"); @@ -67,15 +73,17 @@ describe("SearXNG web search provider", () => { await Settings.init({ agentDir }); const captured: { headers?: Headers } = {}; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { captured.headers = new Headers(init?.headers); - return new Response(JSON.stringify({ results: [] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ results: [] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - await searchSearXNG({ query: "settings basic auth" }); + await searchSearXNG({ query: "settings basic auth", fetch: fetchMock }); expect(captured.headers?.get("Authorization")).toBe( `Basic ${Buffer.from("alice:s3cret", "utf-8").toString("base64")}`, @@ -92,15 +100,17 @@ describe("SearXNG web search provider", () => { process.env.SEARXNG_BASIC_PASSWORD = "s3cret"; const captured: { headers?: Headers } = {}; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { captured.headers = new Headers(init?.headers); - return new Response(JSON.stringify({ results: [] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ results: [] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - await searchSearXNG({ query: "auth precedence" }); + await searchSearXNG({ query: "auth precedence", fetch: fetchMock }); expect(captured.headers?.get("Authorization")).toBe( `Basic ${Buffer.from("alice:s3cret", "utf-8").toString("base64")}`, @@ -113,15 +123,17 @@ describe("SearXNG web search provider", () => { process.env.SEARXNG_BASIC_PASSWORD = ""; const captured: { headers?: Headers } = {}; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { captured.headers = new Headers(init?.headers); - return new Response(JSON.stringify({ results: [] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ results: [] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - await searchSearXNG({ query: "empty password" }); + await searchSearXNG({ query: "empty password", fetch: fetchMock }); expect(captured.headers?.get("Authorization")).toBe(`Basic ${Buffer.from("alice:", "utf-8").toString("base64")}`); }); @@ -132,15 +144,17 @@ describe("SearXNG web search provider", () => { process.env.SEARXNG_BASIC_PASSWORD = "s3cret"; const captured: { headers?: Headers } = {}; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { captured.headers = new Headers(init?.headers); - return new Response(JSON.stringify({ results: [] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ results: [] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - await searchSearXNG({ query: "empty username" }); + await searchSearXNG({ query: "empty username", fetch: fetchMock }); expect(captured.headers?.get("Authorization")).toBe( `Basic ${Buffer.from(":s3cret", "utf-8").toString("base64")}`, @@ -200,15 +214,17 @@ describe("SearXNG web search provider", () => { process.env.SEARXNG_TOKEN = "bearer-token"; const captured: { headers?: Headers } = {}; - using _hook = hookFetch((_input, init) => { + const fetchMock: FetchImpl = (_input, init) => { captured.headers = new Headers(init?.headers); - return new Response(JSON.stringify({ results: [] }), { - status: 200, - headers: { "Content-Type": "application/json" }, - }); - }); + return Promise.resolve( + new Response(JSON.stringify({ results: [] }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }), + ); + }; - await searchSearXNG({ query: "bearer search" }); + await searchSearXNG({ query: "bearer search", fetch: fetchMock }); expect(captured.headers?.get("Authorization")).toBe("Bearer bearer-token"); }); diff --git a/packages/coding-agent/test/tools/web-search-tavily.test.ts b/packages/coding-agent/test/tools/web-search-tavily.test.ts index 73836ab21..480b78dd2 100644 --- a/packages/coding-agent/test/tools/web-search-tavily.test.ts +++ b/packages/coding-agent/test/tools/web-search-tavily.test.ts @@ -2,7 +2,6 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { searchTavily } from "@oh-my-pi/pi-coding-agent/web/search/providers/tavily"; import type { SearchProviderError } from "@oh-my-pi/pi-coding-agent/web/search/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; describe("Tavily web search provider", () => { beforeEach(() => { @@ -40,7 +39,7 @@ describe("Tavily web search provider", () => { it("maps Tavily responses into SearchResponse and forwards recency filters", async () => { let requestBody: Record | null = null; - using _hook = hookFetch(async (_input, init) => { + const fetchMock = async (_input: string | URL | Request, init?: RequestInit): Promise => { requestBody = JSON.parse(String(init?.body ?? "null")) as Record; return new Response( JSON.stringify({ @@ -61,9 +60,14 @@ describe("Tavily web search provider", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }); + }; - const response = await searchTavily({ ...makeParams("latest ai news"), numSearchResults: 2, recency: "week" }); + const response = await searchTavily({ + ...makeParams("latest ai news"), + numSearchResults: 2, + recency: "week", + fetch: fetchMock, + }); // Recency must not couple to topic — topic should be absent (Tavily defaults to general) expect(requestBody).toMatchObject({ query: "latest ai news", @@ -96,15 +100,15 @@ describe("Tavily web search provider", () => { }); it("surfaces structured API errors", async () => { - using _hook = hookFetch( - () => + const fetchMock = (): Promise => + Promise.resolve( new Response(JSON.stringify({ detail: { error: "invalid api key" } }), { status: 401, headers: { "Content-Type": "application/json" }, }), - ); + ); - await expect(searchTavily(makeParams("bad auth"))).rejects.toEqual( + await expect(searchTavily({ ...makeParams("bad auth"), fetch: fetchMock })).rejects.toEqual( expect.objectContaining({ provider: "tavily", status: 401, diff --git a/packages/coding-agent/test/web/search/abort-and-timeout.test.ts b/packages/coding-agent/test/web/search/abort-and-timeout.test.ts index d0fa45904..b01408228 100644 --- a/packages/coding-agent/test/web/search/abort-and-timeout.test.ts +++ b/packages/coding-agent/test/web/search/abort-and-timeout.test.ts @@ -12,7 +12,7 @@ * helper itself is exercised directly. */ import { afterEach, describe, expect, it, vi } from "bun:test"; -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import type { AgentStorage } from "@oh-my-pi/pi-coding-agent/session/agent-storage"; import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; import { ToolAbortError } from "@oh-my-pi/pi-coding-agent/tools/tool-errors"; @@ -23,7 +23,6 @@ import type { SearchParams } from "@oh-my-pi/pi-coding-agent/web/search/provider import { searchBrave } from "@oh-my-pi/pi-coding-agent/web/search/providers/brave"; import { withHardTimeout } from "@oh-my-pi/pi-coding-agent/web/search/providers/utils"; import type { SearchProviderId, SearchResponse } from "@oh-my-pi/pi-coding-agent/web/search/types"; -import { hookFetch } from "@oh-my-pi/pi-utils"; const FAKE_SESSION = {} as ToolSession; const fakeStorage = { @@ -68,15 +67,15 @@ describe("Anthropic provider hard-timeout wiring", () => { process.env.ANTHROPIC_SEARCH_API_KEY = "sk-test"; let capturedSignal: AbortSignal | null | undefined; - using _hook = hookFetch(async (_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { capturedSignal = init?.signal; return new Response(JSON.stringify({ content: [{ type: "text", text: "ok" }], usage: {} }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - await searchAnthropic({ query: "ping", system_prompt: "" }, fakeStorage); + await searchAnthropic({ query: "ping", system_prompt: "", fetch: fetchMock }, fakeStorage); // Without the hard-timeout wrapper, init.signal would be undefined when // the caller didn't supply one — leaving fetch with no cancellation at @@ -90,15 +89,15 @@ describe("Anthropic provider hard-timeout wiring", () => { const ac = new AbortController(); let capturedSignal: AbortSignal | null | undefined; - using _hook = hookFetch(async (_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { capturedSignal = init?.signal; return new Response(JSON.stringify({ content: [{ type: "text", text: "ok" }], usage: {} }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - await searchAnthropic({ query: "ping", system_prompt: "", signal: ac.signal }, fakeStorage); + await searchAnthropic({ query: "ping", system_prompt: "", signal: ac.signal, fetch: fetchMock }, fakeStorage); // The signal handed to fetch must be a *composed* one, not the raw // caller signal: that's what guarantees the hard timeout fires even @@ -110,17 +109,18 @@ describe("Anthropic provider hard-timeout wiring", () => { process.env.ANTHROPIC_SEARCH_BASE_URL = "https://search.example.test/"; let capturedUrl: string | undefined; - using _hook = hookFetch(async input => { + const fetchMock: FetchImpl = async input => { capturedUrl = String(input); return new Response(JSON.stringify({ content: [{ type: "text", text: "ok" }], usage: {} }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; await searchAnthropic({ query: "ping", systemPrompt: "", + fetch: fetchMock, authStorage: { getApiKey: async () => "sk-fallback", resolver: vi.fn(() => async () => "sk-fallback"), @@ -141,15 +141,15 @@ describe("Brave provider hard-timeout wiring", () => { process.env.BRAVE_API_KEY = "brave-test-key"; let capturedSignal: AbortSignal | null | undefined; - using _hook = hookFetch(async (_input, init) => { + const fetchMock: FetchImpl = async (_input, init) => { capturedSignal = init?.signal; return new Response(JSON.stringify({ web: { results: [] } }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }); + }; - await searchBrave({ query: "ping" }); + await searchBrave({ query: "ping", fetch: fetchMock }); expect(capturedSignal).toBeInstanceOf(AbortSignal); expect(capturedSignal?.aborted).toBe(false); diff --git a/packages/coding-agent/test/web/search/codex-broker.test.ts b/packages/coding-agent/test/web/search/codex-broker.test.ts index f456e6bc8..ec278ddbd 100644 --- a/packages/coding-agent/test/web/search/codex-broker.test.ts +++ b/packages/coding-agent/test/web/search/codex-broker.test.ts @@ -1,9 +1,9 @@ import { describe, expect, it, vi } from "bun:test"; import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { AgentStorage } from "@oh-my-pi/pi-coding-agent/session/agent-storage"; import type { SearchParams } from "@oh-my-pi/pi-coding-agent/web/search/providers/base"; import { searchCodex } from "@oh-my-pi/pi-coding-agent/web/search/providers/codex"; -import { hookFetch } from "@oh-my-pi/pi-utils"; function makeSseResponse(): string { return [ @@ -39,10 +39,10 @@ describe("Codex web search broker auth", () => { const openSpy = vi.spyOn(AgentStorage, "open"); let requestHeaders: Headers | undefined; - using _hook = hookFetch(async (_url, init) => { + const fetchMock: FetchImpl = async (_url, init) => { requestHeaders = new Headers(init?.headers); return new Response(makeSseResponse(), { status: 200, headers: { "Content-Type": "text/event-stream" } }); - }); + }; const params: SearchParams = { query: "broker codex search", @@ -51,7 +51,7 @@ describe("Codex web search broker auth", () => { sessionId: "codex-broker-session", }; - const result = await searchCodex(params); + const result = await searchCodex({ ...params, fetch: fetchMock }); expect(result.provider).toBe("codex"); expect(getOAuthAccess).toHaveBeenCalledWith("openai-codex", "codex-broker-session", { signal: undefined }); diff --git a/packages/coding-agent/test/web/search/perplexity.test.ts b/packages/coding-agent/test/web/search/perplexity.test.ts index a582141ab..e32642dcd 100644 --- a/packages/coding-agent/test/web/search/perplexity.test.ts +++ b/packages/coding-agent/test/web/search/perplexity.test.ts @@ -1,7 +1,6 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; -import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; import { PerplexityProvider, searchPerplexity } from "@oh-my-pi/pi-coding-agent/web/search/providers/perplexity"; -import { hookFetch } from "@oh-my-pi/pi-utils"; const API_URL = "https://api.perplexity.ai/chat/completions"; @@ -16,8 +15,8 @@ const apiKeyAuthStorage = { }, } as unknown as AuthStorage; -function mockApi(capture: (body: Record) => void, response: Record) { - return hookFetch(async (input, init) => { +function mockApi(capture: (body: Record) => void, response: Record): FetchImpl { + return async (input, init) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; if (url === API_URL) { capture(JSON.parse(init?.body as string)); @@ -27,7 +26,7 @@ function mockApi(capture: (body: Record) => void, response: Rec }); } return new Response("not mocked", { status: 500 }); - }); + }; } function baseResponse(extra: Record = {}) { @@ -60,9 +59,8 @@ describe("Perplexity API-key request shape", () => { it("requests comprehensive defaults: 20 results, high context, related questions", async () => { let body: Record | undefined; - using _hook = mockApi(b => (body = b), baseResponse()); - - await searchPerplexity({ query: "quic vs tcp", authStorage: apiKeyAuthStorage }); + const fetchMock = mockApi(b => (body = b), baseResponse()); + await searchPerplexity({ query: "quic vs tcp", authStorage: apiKeyAuthStorage, fetch: fetchMock }); expect(body?.num_search_results).toBe(20); expect(body?.web_search_options).toMatchObject({ search_type: "pro", search_context_size: "high" }); @@ -71,28 +69,41 @@ describe("Perplexity API-key request shape", () => { it("honors a caller-supplied num_search_results over the default", async () => { let body: Record | undefined; - using _hook = mockApi(b => (body = b), baseResponse()); + const fetchMock = mockApi(b => (body = b), baseResponse()); - await searchPerplexity({ query: "quic vs tcp", authStorage: apiKeyAuthStorage, num_search_results: 5 }); + await searchPerplexity({ + query: "quic vs tcp", + authStorage: apiKeyAuthStorage, + num_search_results: 5, + fetch: fetchMock, + }); expect(body?.num_search_results).toBe(5); }); it("parses related_questions into relatedQuestions, preserving order and dropping blanks", async () => { - using _hook = mockApi( + const fetchMock = mockApi( () => {}, baseResponse({ related_questions: ["How does QUIC handle loss?", " ", "What is 0-RTT?"] }), ); - const response = await searchPerplexity({ query: "quic vs tcp", authStorage: apiKeyAuthStorage }); + const response = await searchPerplexity({ + query: "quic vs tcp", + authStorage: apiKeyAuthStorage, + fetch: fetchMock, + }); expect(response.relatedQuestions).toEqual(["How does QUIC handle loss?", "What is 0-RTT?"]); }); it("omits relatedQuestions when the API returns none", async () => { - using _hook = mockApi(() => {}, baseResponse()); + const fetchMock = mockApi(() => {}, baseResponse()); - const response = await searchPerplexity({ query: "quic vs tcp", authStorage: apiKeyAuthStorage }); + const response = await searchPerplexity({ + query: "quic vs tcp", + authStorage: apiKeyAuthStorage, + fetch: fetchMock, + }); expect(response.relatedQuestions).toBeUndefined(); }); @@ -120,7 +131,7 @@ const anonymousAuthStorage = { }, } as unknown as AuthStorage; -function mockOAuth(capture: (body: Record, headers: Headers) => void) { +function mockOAuth(capture: (body: Record, headers: Headers) => void): FetchImpl { const event = { final: true, display_model: "turbo", @@ -134,17 +145,17 @@ function mockOAuth(capture: (body: Record, headers: Headers) => ], }; const sseBody = `data: ${JSON.stringify(event)}\n\n`; - return hookFetch(async (input, init) => { + return async (input, init) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; if (url === OAUTH_ASK_URL) { capture(JSON.parse(init?.body as string), new Headers(init?.headers)); return new Response(sseBody, { status: 200, headers: { "Content-Type": "text/event-stream" } }); } return new Response("not mocked", { status: 500 }); - }); + }; } -function mockAnonymous(capture: (body: Record, headers: Headers) => void) { +function mockAnonymous(capture: (body: Record, headers: Headers) => void): FetchImpl { const answerPayload = { answer: "Anonymous answer", web_results: [{ name: "Example", url: "https://example.com", snippet: "s" }], @@ -158,14 +169,14 @@ function mockAnonymous(capture: (body: Record, headers: Headers text: JSON.stringify([{ step_type: "FINAL", content: { answer: JSON.stringify(answerPayload) }, uuid: "" }]), }; const sseBody = `data: ${JSON.stringify(event)}\n\n`; - return hookFetch(async (input, init) => { + return async (input, init) => { const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; if (url === OAUTH_ASK_URL) { capture(JSON.parse(init?.body as string), new Headers(init?.headers)); return new Response(sseBody, { status: 200, headers: { "Content-Type": "text/event-stream" } }); } return new Response("not mocked", { status: 500 }); - }); + }; } describe("Perplexity OAuth request shape", () => { @@ -184,7 +195,7 @@ describe("Perplexity OAuth request shape", () => { it("sends the bare query, never the API-style system prompt, to the ask endpoint", async () => { let body: Record | undefined; let headers: Headers | undefined; - using _hook = mockOAuth((b, h) => { + const fetchMock = mockOAuth((b, h) => { body = b; headers = h; }); @@ -193,6 +204,7 @@ describe("Perplexity OAuth request shape", () => { query: "quic vs tcp", system_prompt: "Research assistant with web search. Synthesize comprehensive answers.", authStorage: oauthAuthStorage, + fetch: fetchMock, }); // The consumer ask endpoint has no system slot; prepending the prompt makes @@ -233,12 +245,16 @@ describe("Perplexity anonymous fallback", () => { it("uses the browser ask endpoint without credential headers when no key is configured", async () => { let body: Record | undefined; let headers: Headers | undefined; - using _hook = mockAnonymous((b, h) => { + const fetchMock = mockAnonymous((b, h) => { body = b; headers = h; }); - const response = await searchPerplexity({ query: "anonymous search", authStorage: anonymousAuthStorage }); + const response = await searchPerplexity({ + query: "anonymous search", + authStorage: anonymousAuthStorage, + fetch: fetchMock, + }); const requestParams = body?.params as Record; expect(headers?.has("authorization")).toBe(false); diff --git a/packages/coding-agent/test/web/search/tavily.test.ts b/packages/coding-agent/test/web/search/tavily.test.ts index 8eb7d2e7a..2e5eeccfe 100644 --- a/packages/coding-agent/test/web/search/tavily.test.ts +++ b/packages/coding-agent/test/web/search/tavily.test.ts @@ -1,11 +1,11 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import type { FetchImpl } from "@oh-my-pi/pi-ai/types"; import { buildRequestBody, searchTavily, type TavilySearchParams, } from "@oh-my-pi/pi-coding-agent/web/search/providers/tavily"; -import { hookFetch } from "@oh-my-pi/pi-utils"; describe("Tavily buildRequestBody", () => { afterEach(() => { @@ -76,8 +76,9 @@ describe("Tavily searchTavily request shape (integration)", () => { process.env.TAVILY_API_KEY = "test-key"; let capturedBody: Record | undefined; - using _hook = hookFetch(async (input, init) => { - const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; + const fetchMock: FetchImpl = async (input, init) => { + const url = + typeof input === "string" ? input : input instanceof URL ? input.toString() : (input as Request).url; if (url === "https://api.tavily.com/search") { capturedBody = JSON.parse(init?.body as string); return new Response( @@ -97,9 +98,12 @@ describe("Tavily searchTavily request shape (integration)", () => { ); } return new Response("not mocked", { status: 500 }); - }); + }; - const response = await searchTavily(makeParams("Bun runtime latest release notes", { recency: "week" })); + const response = await searchTavily({ + ...makeParams("Bun runtime latest release notes", { recency: "week" }), + fetch: fetchMock, + }); expect(capturedBody).toBeDefined(); expect(capturedBody).not.toHaveProperty("topic"); @@ -116,8 +120,9 @@ describe("Tavily searchTavily request shape (integration)", () => { process.env.TAVILY_API_KEY = "test-key"; let capturedBody: Record | undefined; - using _hook = hookFetch(async (input, init) => { - const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url; + const fetchMock: FetchImpl = async (input, init) => { + const url = + typeof input === "string" ? input : input instanceof URL ? input.toString() : (input as Request).url; if (url === "https://api.tavily.com/search") { capturedBody = JSON.parse(init?.body as string); return new Response(JSON.stringify({ answer: "", results: [], request_id: "req-0" }), { @@ -126,9 +131,9 @@ describe("Tavily searchTavily request shape (integration)", () => { }); } return new Response("not mocked", { status: 500 }); - }); + }; - await searchTavily(makeParams("bun sqlite")); + await searchTavily({ ...makeParams("bun sqlite"), fetch: fetchMock }); expect(capturedBody).toBeDefined(); expect(capturedBody).not.toHaveProperty("topic"); diff --git a/packages/mnemopi/CHANGELOG.md b/packages/mnemopi/CHANGELOG.md index ecde88588..87bdfabd8 100644 --- a/packages/mnemopi/CHANGELOG.md +++ b/packages/mnemopi/CHANGELOG.md @@ -1,6 +1,11 @@ # Changelog ## [Unreleased] +### Added + +- Added a `fetch` option to `ExtractionClient` to inject a custom fetch implementation for remote LLM requests +- Added an optional `fetch` option to `extractFacts` to control the transport used for remote extraction calls +- Added support for passing a custom `fetch` implementation through `complete` and `summarizeMemories` via remote LLM options ## [15.9.1] - 2026-06-04 @@ -62,4 +67,4 @@ - Fixed `rememberBatch(..., { extract: true })` to run background fact extraction for batch uploads (including per-item `extract` flags) so extracted facts are generated and recallable after extraction - Fixed `extract: true` fact extraction to continue safely when no LLM is configured by turning extraction failures into no-op background tasks - Fixed configured LLM fact extraction by using temperature 0 so re-ingesting the same text is deterministic and avoids near-duplicate extractions -- Fixed `remember(..., { extract: true })` silently dropping the flag: it now schedules the LLM fact extractor (`extractFactsSafe`) over the stored content and persists the extracted facts so they become recallable. Previously the LLM extractor had no production callers and `extract` was dead. +- Fixed `remember(..., { extract: true })` silently dropping the flag: it now schedules the LLM fact extractor (`extractFactsSafe`) over the stored content and persists the extracted facts so they become recallable. Previously the LLM extractor had no production callers and `extract` was dead. \ No newline at end of file diff --git a/packages/mnemopi/src/core/extraction.ts b/packages/mnemopi/src/core/extraction.ts index 52be5d56c..01bd04833 100644 --- a/packages/mnemopi/src/core/extraction.ts +++ b/packages/mnemopi/src/core/extraction.ts @@ -7,6 +7,7 @@ import { cleanOutput, configuredLlmWillHandleCall, llmAvailable, + type RemoteLlmOptions, } from "./local-llm"; import { getMnemopiRuntimeOptions } from "./runtime-options"; @@ -231,7 +232,7 @@ async function localFallback(prompt: string, sourceText: string, diag = getDiagn return []; } -export async function extractFacts(text: string | null | undefined): Promise { +export async function extractFacts(text: string | null | undefined, options: RemoteLlmOptions = {}): Promise { const diag = getDiagnostics(); if (typeof text !== "string" || text.trim() === "") { return []; @@ -303,7 +304,7 @@ export async function extractFacts(text: string | null | undefined): Promise 0) { diff --git a/packages/mnemopi/src/core/extraction/client.ts b/packages/mnemopi/src/core/extraction/client.ts index 2269a3da6..6636acb40 100644 --- a/packages/mnemopi/src/core/extraction/client.ts +++ b/packages/mnemopi/src/core/extraction/client.ts @@ -1,3 +1,5 @@ +import type { FetchImpl } from "@oh-my-pi/pi-ai"; + import { getDiagnostics } from "./diagnostics"; import { EXTRACTION_SYSTEM_PROMPT, EXTRACTION_USER_TEMPLATE } from "./prompts"; @@ -26,6 +28,13 @@ export interface ExtractedFact { [key: string]: unknown; } +export interface ExtractionClientOptions { + model?: string | null; + apiKey?: string | null; + baseUrl?: string | null; + fetch?: FetchImpl; +} + function sleep(ms: number): Promise { const { promise, resolve } = Promise.withResolvers(); setTimeout(resolve, ms); @@ -45,11 +54,13 @@ export class ExtractionClient { apiKey: string; baseUrl: string; callCount = 0; + private readonly fetchImpl: FetchImpl; - constructor(opts: { model?: string | null; apiKey?: string | null; baseUrl?: string | null } = {}) { + constructor(opts: ExtractionClientOptions = {}) { this.model = opts.model || DEFAULT_EXTRACTION_MODEL; this.apiKey = opts.apiKey ?? process.env.OPENROUTER_API_KEY ?? ""; this.baseUrl = (opts.baseUrl || OPENROUTER_BASE_URL).replace(/\/+$/, ""); + this.fetchImpl = opts.fetch ?? fetch; } async chat(messages: readonly ChatMessage[], temperature = 0, maxTokens = 4096): Promise { @@ -89,7 +100,7 @@ export class ExtractionClient { temperature: number, maxTokens: number, ): Promise { - const response = await fetch(`${this.baseUrl}/chat/completions`, { + const response = await this.fetchImpl(`${this.baseUrl}/chat/completions`, { method: "POST", headers: authHeader(this.apiKey), body: JSON.stringify({ model, messages, temperature, max_tokens: maxTokens }), diff --git a/packages/mnemopi/src/core/llm-backends.ts b/packages/mnemopi/src/core/llm-backends.ts index 8b36b717c..2699707f9 100644 --- a/packages/mnemopi/src/core/llm-backends.ts +++ b/packages/mnemopi/src/core/llm-backends.ts @@ -1,9 +1,12 @@ +import type { FetchImpl } from "@oh-my-pi/pi-ai"; + export interface CompleteOptions { maxTokens?: number; temperature?: number; timeout?: number; provider?: string | null; model?: string | null; + fetch?: FetchImpl; } export interface LlmBackend { diff --git a/packages/mnemopi/src/core/local-llm.ts b/packages/mnemopi/src/core/local-llm.ts index c8544bd1c..b883c1609 100644 --- a/packages/mnemopi/src/core/local-llm.ts +++ b/packages/mnemopi/src/core/local-llm.ts @@ -1,5 +1,5 @@ -import { type Api, type AssistantMessage, completeSimple, type Model } from "@oh-my-pi/pi-ai"; -import { callHostLlm, getHostLlmBackend } from "./llm-backends"; +import { type Api, type AssistantMessage, completeSimple, type FetchImpl, type Model } from "@oh-my-pi/pi-ai"; +import { type CompleteOptions, callHostLlm, getHostLlmBackend } from "./llm-backends"; import { getMnemopiRuntimeOptions, isPiAiModel, @@ -8,6 +8,10 @@ import { } from "./runtime-options"; const ENV_MODEL_REPO = process.env.MNEMOPI_LLM_REPO ?? ""; +export interface RemoteLlmOptions { + fetch?: FetchImpl; +} + const ENV_MODEL_FILE = process.env.MNEMOPI_LLM_FILE ?? ""; export const DEFAULT_MODEL_REPO = ENV_MODEL_REPO !== "" && ENV_MODEL_FILE !== "" ? ENV_MODEL_REPO : "TheBloke/TinyLlama-1.1B-Chat-v1.0-GGUF"; @@ -309,7 +313,11 @@ export function llmAvailable(): boolean { return llmEnabled() && llmBaseUrl() !== ""; } -export async function callRemoteLlm(prompt: string, temperature = 0.3): Promise { +export async function callRemoteLlm( + prompt: string, + temperature = 0.3, + options: RemoteLlmOptions = {}, +): Promise { const baseUrl = llmBaseUrl(); if (baseUrl === "") { return null; @@ -321,8 +329,9 @@ export async function callRemoteLlm(prompt: string, temperature = 0.3): Promise< headers.Authorization = `Bearer ${apiKey}`; } + const fetchImpl = options.fetch ?? fetch; try { - const response = await fetch(`${baseUrl}/chat/completions`, { + const response = await fetchImpl(`${baseUrl}/chat/completions`, { method: "POST", headers, body: JSON.stringify({ @@ -355,7 +364,11 @@ export async function callLocalLlm(_prompt: string): Promise { return null; } -async function summarizeChunk(memories: readonly string[], source = ""): Promise { +async function summarizeChunk( + memories: readonly string[], + source = "", + options: RemoteLlmOptions = {}, +): Promise { const hostPrompt = buildHostPrompt(memories, source); const prompt = buildPrompt(memories, source); if (configuredLlmWillHandleCall()) { @@ -380,7 +393,7 @@ async function summarizeChunk(memories: readonly string[], source = ""): Promise } if (llmEnabled() && llmBaseUrl() !== "" && !envBool("MNEMOPI_FORCE_LOCAL", false)) { - const raw = await callRemoteLlm(prompt); + const raw = await callRemoteLlm(prompt, 0.3, options); if (raw !== null) { const cleaned = cleanOutput(raw); return cleaned === "" ? null : cleaned; @@ -395,7 +408,11 @@ async function summarizeChunk(memories: readonly string[], source = ""): Promise return null; } -export async function summarizeMemories(memories: readonly string[], source = ""): Promise { +export async function summarizeMemories( + memories: readonly string[], + source = "", + options: RemoteLlmOptions = {}, +): Promise { if (memories.length === 0) { return null; } @@ -403,7 +420,7 @@ export async function summarizeMemories(memories: readonly string[], source = "" const chunks = chunkMemoriesByBudget(memories, source); const chunkSummaries: string[] = []; for (const chunk of chunks) { - const summary = await summarizeChunk(chunk, source); + const summary = await summarizeChunk(chunk, source, options); if (summary !== null) { chunkSummaries.push(summary); } @@ -413,13 +430,17 @@ export async function summarizeMemories(memories: readonly string[], source = "" return null; } if (chunkSummaries.length > 1) { - const final = await summarizeChunk(chunkSummaries, `${source} [chunked ${chunks.length} parts]`); + const final = await summarizeChunk(chunkSummaries, `${source} [chunked ${chunks.length} parts]`, options); return final ?? chunkSummaries[0] ?? null; } return chunkSummaries[0] ?? null; } -export async function complete(prompt: string, temperature = 0.3): Promise { +export async function complete( + prompt: string, + temperature = 0.3, + options: CompleteOptions = {}, +): Promise { if (configuredLlmWillHandleCall()) { const raw = await callConfiguredCompletion(prompt, temperature, { maxTokens: llmMaxTokens() }); return raw === null ? null : cleanOutput(raw) || null; @@ -429,7 +450,7 @@ export async function complete(prompt: string, temperature = 0.3): Promise { restoreEnv(); - globalThis.fetch = ORIGINAL_FETCH; resetHostLlmBackendForTests(); resetExtractionStats(); }); @@ -30,7 +29,7 @@ describe("extraction integration", () => { process.env.MNEMOPI_LLM_ENABLED = "true"; process.env.MNEMOPI_LLM_BASE_URL = "http://fake-remote/v1"; let payloadJson = ""; - globalThis.fetch = (async (_input: Parameters[0], init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input, init) => { payloadJson = String(init?.body); return new Response( JSON.stringify({ @@ -38,9 +37,9 @@ describe("extraction integration", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }) as unknown as typeof fetch; + }; - const facts = await extractFacts("I prefer deterministic tests."); + const facts = await extractFacts("I prefer deterministic tests.", { fetch: fetchMock }); expect(facts).toEqual(["Ada prefers deterministic tests"]); const payload = JSON.parse(payloadJson) as { temperature?: number; @@ -55,7 +54,7 @@ describe("extraction integration", () => { it("parses structured fact objects through ExtractionClient with fake HTTP", async () => { let requestedUrl = ""; - globalThis.fetch = (async (input: Parameters[0]) => { + const fetchMock: FetchImpl = async input => { requestedUrl = String(input); return new Response( JSON.stringify({ @@ -70,11 +69,12 @@ describe("extraction integration", () => { }), { status: 200, headers: { "Content-Type": "application/json" } }, ); - }) as unknown as typeof fetch; + }; const client = new ExtractionClient({ apiKey: "sk-test", baseUrl: "http://openrouter.test/api/v1", + fetch: fetchMock, }); const facts = await client.extractFacts([{ role: "user", content: "Ada prefers deterministic tests." }]); expect(requestedUrl).toBe("http://openrouter.test/api/v1/chat/completions"); @@ -87,15 +87,16 @@ describe("extraction integration", () => { }); it("records malformed cloud JSON as a diagnostic failure", async () => { - globalThis.fetch = (async () => + const fetchMock: FetchImpl = async () => new Response(JSON.stringify({ choices: [{ message: { content: "Here: [oops, not json]" } }] }), { status: 200, headers: { "Content-Type": "application/json" }, - })) as unknown as typeof fetch; + }); const client = new ExtractionClient({ apiKey: "sk-test", baseUrl: "http://openrouter.test/api/v1", + fetch: fetchMock, }); expect(await client.extractFacts([{ role: "user", content: "Ada prefers tea." }])).toEqual([]); const cloud = getExtractionStats().by_tier.cloud; diff --git a/packages/mnemopi/test/local-llm.test.ts b/packages/mnemopi/test/local-llm.test.ts index 2f34ec803..564e3389d 100644 --- a/packages/mnemopi/test/local-llm.test.ts +++ b/packages/mnemopi/test/local-llm.test.ts @@ -1,4 +1,5 @@ import { afterEach, describe, expect, it } from "bun:test"; +import type { FetchImpl } from "@oh-my-pi/pi-ai"; import { createMockModel, registerMockApi } from "@oh-my-pi/pi-ai/providers/mock"; import { CallableLlmBackend, @@ -19,7 +20,6 @@ import { Mnemopi } from "@oh-my-pi/pi-mnemopi/core/memory"; import { withMnemopiRuntimeOptions } from "@oh-my-pi/pi-mnemopi/core/runtime-options"; const OLD_ENV = { ...process.env }; -const ORIGINAL_FETCH = globalThis.fetch; function restoreEnv(): void { for (const key in process.env) { @@ -34,7 +34,6 @@ function restoreEnv(): void { afterEach(() => { restoreEnv(); - globalThis.fetch = ORIGINAL_FETCH; resetHostLlmBackendForTests(); }); @@ -47,17 +46,17 @@ describe("local LLM TypeScript port", () => { process.env.MNEMOPI_LLM_MODEL = "test-model"; let auth = ""; let model = ""; - globalThis.fetch = (async (_input: Parameters[0], init?: RequestInit) => { + const fetchMock: FetchImpl = async (_input, init?) => { auth = new Headers(init?.headers).get("authorization") ?? ""; model = (JSON.parse(String(init?.body)) as { model: string }).model; return new Response(JSON.stringify({ choices: [{ message: { content: "Remote summary." } }] }), { status: 200, headers: { "Content-Type": "application/json" }, }); - }) as unknown as typeof fetch; + }; expect(llmAvailable()).toBe(true); - expect(await callRemoteLlm("Test prompt", 0.2)).toBe("Remote summary."); + expect(await callRemoteLlm("Test prompt", 0.2, { fetch: fetchMock })).toBe("Remote summary."); expect(auth).toBe("Bearer sk-test"); expect(model).toBe("test-model"); }); @@ -72,19 +71,19 @@ describe("local LLM TypeScript port", () => { process.env.MNEMOPI_HOST_LLM_ENABLED = "true"; process.env.MNEMOPI_LLM_BASE_URL = "http://remote/v1"; let calls = 0; - globalThis.fetch = (async () => { + const fetchMock: FetchImpl = async () => { calls += 1; return new Response(JSON.stringify({ choices: [{ message: { content: "Remote summary." } }] }), { status: 200, }); - }) as unknown as typeof fetch; + }; setHostLlmBackend(new CallableLlmBackend("host", () => "Host summary.")); - expect(await summarizeMemories(["Memory one"])).toBe("Host summary."); + expect(await summarizeMemories(["Memory one"], "", { fetch: fetchMock })).toBe("Host summary."); expect(calls).toBe(0); setHostLlmBackend(new CallableLlmBackend("host", () => null)); - expect(await summarizeMemories(["Memory one"])).toBeNull(); + expect(await summarizeMemories(["Memory one"], "", { fetch: fetchMock })).toBeNull(); expect(calls).toBe(0); }); @@ -112,15 +111,17 @@ describe("local LLM TypeScript port", () => { process.env.MNEMOPI_LLM_ENABLED = "true"; process.env.MNEMOPI_LLM_BASE_URL = "http://remote.example/v1"; let fetchCalls = 0; - globalThis.fetch = (async () => { + const fetchMock: FetchImpl = async () => { fetchCalls += 1; throw new Error("remote should not be called"); - }) as unknown as typeof fetch; + }; const memory = new Mnemopi({ llm: async (prompt, opts) => `fn:${prompt}:${opts?.maxTokens ?? 0}`, }); try { - const text = await withMnemopiRuntimeOptions(memory.runtimeOptions, () => complete("hello")); + const text = await withMnemopiRuntimeOptions(memory.runtimeOptions, () => + complete("hello", 0.3, { fetch: fetchMock }), + ); expect(text).toBe("fn:hello:2048"); expect(fetchCalls).toBe(0); } finally { @@ -145,13 +146,15 @@ describe("local LLM TypeScript port", () => { process.env.MNEMOPI_LLM_ENABLED = "true"; process.env.MNEMOPI_LLM_BASE_URL = "http://remote.example/v1"; let fetchCalls = 0; - globalThis.fetch = (async () => { + const fetchMock: FetchImpl = async () => { fetchCalls += 1; throw new Error("remote should not be called"); - }) as unknown as typeof fetch; + }; const memory = new Mnemopi({ llm: false }); try { - const text = await withMnemopiRuntimeOptions(memory.runtimeOptions, () => complete("hello")); + const text = await withMnemopiRuntimeOptions(memory.runtimeOptions, () => + complete("hello", 0.3, { fetch: fetchMock }), + ); expect(text).toBeNull(); expect(fetchCalls).toBe(0); } finally { diff --git a/packages/utils/CHANGELOG.md b/packages/utils/CHANGELOG.md index 5bb2baa09..65c00cdd3 100644 --- a/packages/utils/CHANGELOG.md +++ b/packages/utils/CHANGELOG.md @@ -1,6 +1,10 @@ # Changelog ## [Unreleased] +### Removed + +- Removed the exported `hookFetch` API, which previously intercepted `globalThis.fetch` via middleware handlers +- Removed `hookFetch` from the package entrypoint, so imports from `@.../utils` no longer provide this fetch interception helper ## [15.10.0] - 2026-06-06 @@ -44,4 +48,4 @@ ### Added -- Added an XDG-aware tiny-title model cache directory helper for coding-agent local title models. +- Added an XDG-aware tiny-title model cache directory helper for coding-agent local title models. \ No newline at end of file diff --git a/packages/utils/src/hook-fetch.ts b/packages/utils/src/hook-fetch.ts deleted file mode 100644 index 0891bac53..000000000 --- a/packages/utils/src/hook-fetch.ts +++ /dev/null @@ -1,30 +0,0 @@ -/** - * Intercept `globalThis.fetch` with a middleware-style handler. - * - * Returns a `Disposable` so callers can use `using` for automatic cleanup: - * - * ```ts - * using _hook = hookFetch((input, init, next) => { - * if (shouldIntercept(input)) { - * return new Response("mocked"); - * } - * return next(input, init); - * }); - * ``` - */ -export type FetchHandler = ( - input: string | URL | Request, - init: RequestInit | undefined, - next: typeof fetch, -) => Response | Promise; - -export function hookFetch(handler: FetchHandler): Disposable { - const original = globalThis.fetch; - globalThis.fetch = ((input: string | URL | Request, init?: RequestInit) => - handler(input, init, original)) as typeof fetch; - return { - [Symbol.dispose]() { - globalThis.fetch = original; - }, - }; -} diff --git a/packages/utils/src/index.ts b/packages/utils/src/index.ts index 884ae684e..d07e24c55 100644 --- a/packages/utils/src/index.ts +++ b/packages/utils/src/index.ts @@ -8,7 +8,6 @@ export * from "./format"; export * from "./frontmatter"; export * from "./fs-error"; export * from "./glob"; -export * from "./hook-fetch"; export * from "./json"; export * as logger from "./logger"; export * from "./mermaid-ascii";