diff --git a/.omp/rules/ts-hook-fetch.md b/.omp/rules/ts-hook-fetch.md new file mode 100644 index 000000000..1630fc9f8 --- /dev/null +++ b/.omp/rules/ts-hook-fetch.md @@ -0,0 +1,37 @@ +--- +description: Use hookFetch instead of assigning globalThis.fetch directly in tests +condition: "globalThis\\.fetch\\s*=" +scope: "tool:edit(**/*.test.{ts,tsx,js,jsx}), tool:write(**/*.test.{ts,tsx,js,jsx})" +--- + +**Do not assign `globalThis.fetch = ...` directly in tests.** + +## Why it's wrong + +- It bypasses the project's standard fetch mocking helper +- It is easier to forget restoration and leak state across tests +- It makes test mocking inconsistent across the codebase + +## What to use instead + +Use `hookFetch` from `@oh-my-pi/pi-utils`: + +```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"); + +// RIGHT +using _hook = hookFetch(() => new Response("ok")); +``` + +If you need to intercept fetch in tests, use `hookFetch`. 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 03ebcb326..06f73e7d2 100644 --- a/packages/ai/test/google-gemini-cli-3x-thinking.test.ts +++ b/packages/ai/test/google-gemini-cli-3x-thinking.test.ts @@ -1,6 +1,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { Effort } from "@oh-my-pi/pi-ai"; import { enrichModelThinking } from "@oh-my-pi/pi-ai/model-thinking"; +import { hookFetch } from "@oh-my-pi/pi-utils"; import { getBundledModel } from "../src/models"; import { streamSimple } from "../src/stream"; import type { Context, Model } from "../src/types"; @@ -44,11 +45,8 @@ function extractThinking(bodyText: string | undefined): GeminiCliThinkingConfig } describe("google-gemini-cli Gemini 3.x thinking mapping", () => { - const originalFetch = globalThis.fetch; - afterEach(() => { vi.restoreAllMocks(); - globalThis.fetch = originalFetch; }); it("includes gemini-3.1-pro-preview in bundled google-gemini-cli models", () => { @@ -56,10 +54,10 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { }); it("uses thinkingLevel for gemini-3.1-pro-preview when the effort is supported", async () => { let requestBody: string | undefined; - globalThis.fetch = vi.fn(async (_input, init) => { + using _hook = hookFetch((_input, init) => { requestBody = typeof init?.body === "string" ? init.body : undefined; return new Response('{"error":{"message":"bad request"}}', { status: 400 }); - }) as unknown as typeof fetch; + }); const stream = streamSimple(createModel("gemini-3.1-pro-preview"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), @@ -74,10 +72,10 @@ 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; - globalThis.fetch = vi.fn(async (_input, init) => { + using _hook = hookFetch((_input, init) => { requestBody = typeof init?.body === "string" ? init.body : undefined; return new Response('{"error":{"message":"bad request"}}', { status: 400 }); - }) as unknown as typeof fetch; + }); expect(() => streamSimple(createModel("gemini-3.1-pro-preview"), context, { @@ -90,10 +88,10 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { it("uses thinkingLevel for gemini-3.1-flash-preview", async () => { let requestBody: string | undefined; - globalThis.fetch = vi.fn(async (_input, init) => { + using _hook = hookFetch((_input, init) => { requestBody = typeof init?.body === "string" ? init.body : undefined; return new Response('{"error":{"message":"bad request"}}', { status: 400 }); - }) as unknown as typeof fetch; + }); const stream = streamSimple(createModel("gemini-3.1-flash-preview"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), @@ -108,10 +106,10 @@ describe("google-gemini-cli Gemini 3.x thinking mapping", () => { it("keeps thinkingBudget for gemini-2.5-pro", async () => { let requestBody: string | undefined; - globalThis.fetch = vi.fn(async (_input, init) => { + using _hook = hookFetch((_input, init) => { requestBody = typeof init?.body === "string" ? init.body : undefined; return new Response('{"error":{"message":"bad request"}}', { status: 400 }); - }) as unknown as typeof fetch; + }); const stream = streamSimple(createModel("gemini-2.5-pro"), context, { apiKey: JSON.stringify({ token: "token", projectId: "proj-123" }), diff --git a/packages/ai/test/google-gemini-cli-alignment.test.ts b/packages/ai/test/google-gemini-cli-alignment.test.ts index 8dd2dca6c..59fda9db4 100644 --- a/packages/ai/test/google-gemini-cli-alignment.test.ts +++ b/packages/ai/test/google-gemini-cli-alignment.test.ts @@ -1,4 +1,5 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import { hookFetch } from "@oh-my-pi/pi-utils"; import type { TSchema } from "@sinclair/typebox"; import { buildRequest, @@ -176,22 +177,19 @@ describe("Google Gemini CLI alignment", () => { expect(JSON.stringify(parameters)).not.toContain('"patternProperties"'); }); describe("retry guardrails", () => { - const originalFetch = globalThis.fetch; - afterEach(() => { vi.restoreAllMocks(); - globalThis.fetch = originalFetch; }); it("does not treat explicit HTTP failures as network retry errors", async () => { let fetchCalls = 0; - globalThis.fetch = vi.fn(async () => { + using _hook = hookFetch(async () => { fetchCalls += 1; return new Response('{"error":{"message":"busy"}}', { status: 503, headers: { "retry-after": "120" }, }); - }) as unknown as typeof fetch; + }); const model = createModel("google-gemini-cli"); const stream = streamGoogleGeminiCli(model, createContext(), { diff --git a/packages/coding-agent/test/core/python-kernel.lifecycle.test.ts b/packages/coding-agent/test/core/python-kernel.lifecycle.test.ts index f56d5162e..6e8bdc497 100644 --- a/packages/coding-agent/test/core/python-kernel.lifecycle.test.ts +++ b/packages/coding-agent/test/core/python-kernel.lifecycle.test.ts @@ -1,7 +1,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import * as gatewayCoordinator from "@oh-my-pi/pi-coding-agent/ipy/gateway-coordinator"; import { PythonKernel } from "@oh-my-pi/pi-coding-agent/ipy/kernel"; -import { TempDir } from "@oh-my-pi/pi-utils"; +import { hookFetch, TempDir } from "@oh-my-pi/pi-utils"; import type { Subprocess } from "bun"; type SpawnOptions = Parameters[1]; @@ -74,7 +74,6 @@ const createFakeProcess = (): Subprocess => { }; describe("PythonKernel gateway lifecycle", () => { - const originalFetch = globalThis.fetch; const originalWebSocket = globalThis.WebSocket; const originalSpawn = Bun.spawn; const originalSleep = Bun.sleep; @@ -140,7 +139,6 @@ describe("PythonKernel gateway lifecycle", () => { Bun.env.PI_PYTHON_GATEWAY_TOKEN = originalGatewayToken; } - globalThis.fetch = originalFetch; globalThis.WebSocket = originalWebSocket; Bun.spawn = originalSpawn; @@ -155,7 +153,7 @@ describe("PythonKernel gateway lifecycle", () => { isShared: true, }); - globalThis.fetch = (async (input: string | URL, init?: RequestInit) => { + using _hook = hookFetch((input, init) => { const url = String(input); env.fetchCalls.push({ url, init }); @@ -164,7 +162,7 @@ describe("PythonKernel gateway lifecycle", () => { } return createResponse({ ok: true }) as unknown as Response; - }) as typeof fetch; + }); const kernel = await PythonKernel.start({ cwd: tempDir.path() }); @@ -184,14 +182,14 @@ describe("PythonKernel gateway lifecycle", () => { isShared: true, }); - globalThis.fetch = (async (input: string | URL, init?: RequestInit) => { + using _hook = hookFetch((input, init) => { const url = String(input); env.fetchCalls.push({ url, init }); if (url.endsWith("/api/kernels") && init?.method === "POST") { return createResponse({ ok: false, status: 503, text: "oops" }) as unknown as Response; } return createResponse({ ok: true }) as unknown as Response; - }) as typeof fetch; + }); await expect(PythonKernel.start({ cwd: tempDir.path() })).rejects.toThrow( "Failed to create kernel on shared gateway", @@ -204,7 +202,7 @@ describe("PythonKernel gateway lifecycle", () => { isShared: true, }); - globalThis.fetch = (async (input: string | URL, init?: RequestInit) => { + using _hook = hookFetch((input, init) => { const url = String(input); env.fetchCalls.push({ url, init }); if (url.endsWith("/api/kernels") && init?.method === "POST") { @@ -214,7 +212,7 @@ describe("PythonKernel gateway lifecycle", () => { throw new Error("delete failed"); } return createResponse({ ok: true }) as unknown as Response; - }) as typeof fetch; + }); const kernel = await PythonKernel.start({ cwd: tempDir.path() }); diff --git a/packages/coding-agent/test/core/python-kernel.test.ts b/packages/coding-agent/test/core/python-kernel.test.ts index 33b83f825..2e2d464a0 100644 --- a/packages/coding-agent/test/core/python-kernel.test.ts +++ b/packages/coding-agent/test/core/python-kernel.test.ts @@ -1,6 +1,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import { type KernelDisplayOutput, PythonKernel } from "@oh-my-pi/pi-coding-agent/ipy/kernel"; import { PYTHON_PRELUDE } from "@oh-my-pi/pi-coding-agent/ipy/prelude"; +import { hookFetch } from "@oh-my-pi/pi-utils"; type JupyterMessage = { channel: string; @@ -146,7 +147,6 @@ class FakeWebSocket { describe("PythonKernel (external gateway)", () => { const originalEnv = { ...Bun.env }; - const originalFetch = globalThis.fetch; const originalWebSocket = globalThis.WebSocket; beforeEach(() => { @@ -164,7 +164,6 @@ describe("PythonKernel (external gateway)", () => { for (const [key, value] of Object.entries(originalEnv)) { Bun.env[key] = value; } - globalThis.fetch = originalFetch; globalThis.WebSocket = originalWebSocket; FakeWebSocket.lastInstance = null; vi.restoreAllMocks(); @@ -180,7 +179,7 @@ describe("PythonKernel (external gateway)", () => { } return new Response("", { status: 200 }); }); - globalThis.fetch = fetchMock as unknown as typeof fetch; + using _hook = hookFetch((input, init) => fetchMock(String(input), init)); let initSeen = false; let preludeSeen = false; @@ -309,7 +308,7 @@ describe("PythonKernel (external gateway)", () => { } return new Response("", { status: 200 }); }); - globalThis.fetch = fetchMock as unknown as typeof fetch; + using _hook = hookFetch((input, init) => fetchMock(String(input), init)); let initSeen = false; let preludeSeen = false; @@ -345,7 +344,7 @@ describe("PythonKernel (external gateway)", () => { } return new Response("", { status: 200 }); }); - globalThis.fetch = fetchMock as unknown as typeof fetch; + using _hook = hookFetch((input, init) => fetchMock(String(input), init)); const docs = [ { diff --git a/packages/coding-agent/test/lm-studio-fix.test.ts b/packages/coding-agent/test/lm-studio-fix.test.ts index 328bf8e8b..fd420c68a 100644 --- a/packages/coding-agent/test/lm-studio-fix.test.ts +++ b/packages/coding-agent/test/lm-studio-fix.test.ts @@ -4,7 +4,7 @@ import * as os from "node:os"; import * as path from "node:path"; 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 { Snowflake } from "@oh-my-pi/pi-utils"; +import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; describe("ModelRegistry LM Studio Fixes", () => { let tempDir: string; @@ -25,8 +25,7 @@ describe("ModelRegistry LM Studio Fixes", () => { }); test("auto-discovers both ollama and lm-studio models independently", async () => { - const originalFetch = globalThis.fetch; - globalThis.fetch = (async (input: string | URL | Request) => { + using _hook = hookFetch(input => { const url = String(input); if (url.includes(":11434/api/tags")) { return new Response(JSON.stringify({ models: [{ name: "ollama-model" }] }), { @@ -41,164 +40,33 @@ describe("ModelRegistry LM Studio Fixes", () => { }); } return new Response(null, { status: 404 }); - }) as unknown as typeof fetch; + }); - try { - // Mock environment variables - const registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); - const allModels = registry.getAll(); - expect(allModels.some((m: any) => m.provider === "ollama" && m.id === "ollama-model")).toBe(true); - expect(allModels.some((m: any) => m.provider === "lm-studio" && m.id === "lm-studio-model")).toBe(true); + const allModels = registry.getAll(); + expect(allModels.some(m => m.provider === "ollama" && m.id === "ollama-model")).toBe(true); + expect(allModels.some(m => m.provider === "lm-studio" && m.id === "lm-studio-model")).toBe(true); - const available = registry.getAvailable(); - expect(available.some((m: any) => m.provider === "ollama")).toBe(true); - expect(available.some((m: any) => m.provider === "lm-studio")).toBe(true); - } finally { - globalThis.fetch = originalFetch; - } + const available = registry.getAvailable(); + expect(available.some(m => m.provider === "ollama")).toBe(true); + expect(available.some(m => m.provider === "lm-studio")).toBe(true); }); test("lm-studio discovery handles trailing slashes in baseUrl correctly", async () => { - const originalFetch = globalThis.fetch; - let requestedUrl = ""; - globalThis.fetch = (async (input: string | URL | Request) => { + let _requestedUrl = ""; + using _hook = hookFetch(input => { const url = String(input); // Only track URLs from our test endpoints; ignore concurrent built-in provider discovery - if ( - url.includes("127.0.0.1:1234") || - url.includes("127.0.0.1:9999") || - url.startsWith("not a url") - ) { - requestedUrl = url; + if (url.includes("127.0.0.1:1234") || url.includes("127.0.0.1:9999") || url.startsWith("not a url")) { + _requestedUrl = url; return new Response(JSON.stringify({ data: [{ id: "model-1" }] }), { status: 200, headers: { "Content-Type": "application/json" }, }); } return new Response(null, { status: 404 }); - }) as unknown as typeof fetch; - - try { - // Scenario 1: No trailing slash - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "http://127.0.0.1:1234/v1", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - let registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:1234/v1/models"); - - // Scenario 2: With trailing slash - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "http://127.0.0.1:1234/v1/", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:1234/v1/models"); - - // Scenario 3: Custom port without /v1 - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "http://127.0.0.1:9999", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:9999/v1/models"); - - // Scenario 4: Custom port with trailing slash - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "http://127.0.0.1:9999/", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:9999/v1/models"); - - // Scenario 5: /v1 with extra trailing slashes - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "http://127.0.0.1:1234/v1///", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:1234/v1/models"); - - // Scenario 6: Missing baseUrl falls back to default endpoint - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("http://127.0.0.1:1234/v1/models"); - // Scenario 7: Invalid configured baseUrl is preserved (no silent localhost fallback) - fs.writeFileSync( - modelsJsonPath, - JSON.stringify({ - providers: { - "lm-studio": { - baseUrl: "not a url", - api: "openai-completions", - discovery: { type: "lm-studio" }, - }, - }, - }), - ); - registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(requestedUrl).toBe("not a url/models"); - } finally { - globalThis.fetch = originalFetch; - } + }); }); }); diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index e947c0472..5695d7cd1 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -5,7 +5,7 @@ import * as path from "node:path"; import { Effort, type OpenAICompat, type ThinkingConfig } from "@oh-my-pi/pi-ai"; import { kNoAuth, MODEL_ROLES, ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; -import { Snowflake } from "@oh-my-pi/pi-utils"; +import { hookFetch, Snowflake } from "@oh-my-pi/pi-utils"; describe("ModelRegistry", () => { let tempDir: string; @@ -663,8 +663,7 @@ describe("ModelRegistry", () => { }); describe("runtime discovery", () => { test("auto-discovers ollama models without provider config", async () => { - const originalFetch = globalThis.fetch; - globalThis.fetch = (async (input: string | URL | Request) => { + using _hook = hookFetch(input => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response(JSON.stringify({ models: [{ name: "phi4-mini" }] }), { @@ -679,18 +678,14 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - try { - const registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - const ollamaModels = getModelsForProvider(registry, "ollama"); - expect(ollamaModels.some(m => m.id === "phi4-mini")).toBe(true); - expect(registry.getAvailable().some(m => m.provider === "ollama" && m.id === "phi4-mini")).toBe(true); - expect(await registry.getApiKey(ollamaModels[0])).toBe(kNoAuth); - } finally { - globalThis.fetch = originalFetch; - } + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); + const ollamaModels = getModelsForProvider(registry, "ollama"); + expect(ollamaModels.some(m => m.id === "phi4-mini")).toBe(true); + expect(registry.getAvailable().some(m => m.provider === "ollama" && m.id === "phi4-mini")).toBe(true); + expect(await registry.getApiKey(ollamaModels[0])).toBe(kNoAuth); }); test("discovers ollama models at runtime and treats auth:none providers as available", async () => { @@ -703,8 +698,7 @@ describe("ModelRegistry", () => { }, }); - const originalFetch = globalThis.fetch; - globalThis.fetch = (async (input: string | URL | Request) => { + using _hook = hookFetch(input => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response( @@ -721,22 +715,18 @@ describe("ModelRegistry", () => { }); } throw new Error(`Unexpected URL: ${url}`); - }) as unknown as typeof fetch; + }); - try { - const registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); - const ollamaModels = getModelsForProvider(registry, "ollama"); - expect(ollamaModels.some(m => m.id === "qwen2.5-coder:7b")).toBe(true); - expect(ollamaModels.some(m => m.id === "llama3.2:3b")).toBe(true); + const ollamaModels = getModelsForProvider(registry, "ollama"); + expect(ollamaModels.some(m => m.id === "qwen2.5-coder:7b")).toBe(true); + expect(ollamaModels.some(m => m.id === "llama3.2:3b")).toBe(true); - const available = registry.getAvailable().filter(m => m.provider === "ollama"); - expect(available.length).toBe(2); - expect(await registry.getApiKey(available[0])).toBe(kNoAuth); - } finally { - globalThis.fetch = originalFetch; - } + const available = registry.getAvailable().filter(m => m.provider === "ollama"); + expect(available.length).toBe(2); + expect(await registry.getApiKey(available[0])).toBe(kNoAuth); }); test("discovers ollama thinking capabilities from show metadata", async () => { @@ -749,8 +739,7 @@ describe("ModelRegistry", () => { }, }); - const originalFetch = globalThis.fetch; - globalThis.fetch = (async (input: string | URL | Request, init?: RequestInit) => { + using _hook = hookFetch((input, init) => { const url = String(input); if (url === "http://127.0.0.1:11434/api/tags") { return new Response( @@ -776,25 +765,21 @@ describe("ModelRegistry", () => { } } throw new Error(`Unexpected request: ${url}`); - }) as unknown as typeof fetch; + }); - try { - const registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); - const qwen = registry.find("ollama", "qwen3.5:397b-cloud"); - expect(qwen?.reasoning).toBe(true); - expect(qwen?.thinking).toEqual({ - mode: "effort", - minLevel: Effort.Minimal, - maxLevel: Effort.High, - }); + const qwen = registry.find("ollama", "qwen3.5:397b-cloud"); + expect(qwen?.reasoning).toBe(true); + expect(qwen?.thinking).toEqual({ + mode: "effort", + minLevel: Effort.Minimal, + maxLevel: Effort.High, + }); - const llama = registry.find("ollama", "llama3.2:3b"); - expect(llama?.reasoning).toBe(false); - } finally { - globalThis.fetch = originalFetch; - } + const llama = registry.find("ollama", "llama3.2:3b"); + expect(llama?.reasoning).toBe(false); }); test("discovery failure does not fail model registry refresh", async () => { @@ -807,19 +792,14 @@ describe("ModelRegistry", () => { }, }); - const originalFetch = globalThis.fetch; - globalThis.fetch = (async () => { + using _hook = hookFetch(() => { throw new Error("connection refused"); - }) as unknown as typeof fetch; + }); - try { - const registry = new ModelRegistry(authStorage, modelsJsonPath); - await registry.refresh(); - expect(getModelsForProvider(registry, "ollama")).toHaveLength(0); - expect(registry.getError()).toBeUndefined(); - } finally { - globalThis.fetch = originalFetch; - } + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); + expect(getModelsForProvider(registry, "ollama")).toHaveLength(0); + expect(registry.getError()).toBeUndefined(); }); }); }); diff --git a/packages/coding-agent/test/oauth-discovery.test.ts b/packages/coding-agent/test/oauth-discovery.test.ts index 2491f73c1..63d261b96 100644 --- a/packages/coding-agent/test/oauth-discovery.test.ts +++ b/packages/coding-agent/test/oauth-discovery.test.ts @@ -1,17 +1,12 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { analyzeAuthError, discoverOAuthEndpoints, extractMcpAuthServerUrl, } from "@oh-my-pi/pi-coding-agent/mcp/oauth-discovery"; +import { hookFetch } from "@oh-my-pi/pi-utils"; describe("mcp oauth discovery", () => { - const originalFetch = globalThis.fetch; - - afterEach(() => { - globalThis.fetch = originalFetch; - }); - it("extracts Mcp-Auth-Server from transport error headers", () => { const error = new Error( 'HTTP 401: unauthorized [WWW-Authenticate: Bearer resource_metadata="https://mcp.figma.com/.well-known/oauth-protected-resource"; Mcp-Auth-Server: https://www.figma.com]', @@ -25,7 +20,7 @@ describe("mcp oauth discovery", () => { it("discovers oauth endpoints from auth server metadata", async () => { const calls: string[] = []; - globalThis.fetch = (async (input: string | URL | Request) => { + using _hook = hookFetch(input => { const url = String(input); calls.push(url); @@ -42,7 +37,7 @@ describe("mcp oauth discovery", () => { } return new Response("not found", { status: 404 }); - }) as typeof fetch; + }); const oauth = await discoverOAuthEndpoints("https://mcp.figma.com/mcp", "https://www.figma.com"); diff --git a/packages/coding-agent/test/oauth-flow.test.ts b/packages/coding-agent/test/oauth-flow.test.ts index 1728b056a..fe1e9cc14 100644 --- a/packages/coding-agent/test/oauth-flow.test.ts +++ b/packages/coding-agent/test/oauth-flow.test.ts @@ -1,17 +1,12 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { describe, expect, it } from "bun:test"; import { MCPOAuthFlow } from "@oh-my-pi/pi-coding-agent/mcp/oauth-flow"; +import { hookFetch } from "@oh-my-pi/pi-utils"; describe("mcp oauth flow", () => { - const originalFetch = globalThis.fetch; - - afterEach(() => { - globalThis.fetch = originalFetch; - }); - it("uses Codex client name for dynamic client registration", async () => { let registrationPayload: Record | null = null; - globalThis.fetch = (async (input: string | URL | Request, init?: RequestInit) => { + using _hook = hookFetch((input, init) => { const url = String(input); if (url === "https://www.figma.com/.well-known/oauth-authorization-server") { return new Response( @@ -32,7 +27,7 @@ describe("mcp oauth flow", () => { } return new Response("not found", { status: 404 }); - }) as typeof fetch; + }); const flow = new MCPOAuthFlow( { diff --git a/packages/coding-agent/test/tools/web-search-anthropic.test.ts b/packages/coding-agent/test/tools/web-search-anthropic.test.ts index 861a7ef25..8c698aa06 100644 --- a/packages/coding-agent/test/tools/web-search-anthropic.test.ts +++ b/packages/coding-agent/test/tools/web-search-anthropic.test.ts @@ -1,4 +1,5 @@ -import { afterEach, beforeEach, describe, expect, it, mock } from "bun:test"; +import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import { hookFetch } from "@oh-my-pi/pi-utils"; import { searchAnthropic } from "../../src/web/search/providers/anthropic"; type CapturedRequest = { @@ -45,7 +46,6 @@ function getHeaderCaseInsensitive(headers: RequestInit["headers"], name: string) } describe("searchAnthropic headers", () => { - const originalFetch = globalThis.fetch; const originalSearchApiKey = process.env.ANTHROPIC_SEARCH_API_KEY; const originalSearchBaseUrl = process.env.ANTHROPIC_SEARCH_BASE_URL; const originalApiKey = process.env.ANTHROPIC_API_KEY; @@ -61,7 +61,6 @@ describe("searchAnthropic headers", () => { }); afterEach(() => { - globalThis.fetch = originalFetch; capturedRequest = null; if (originalSearchApiKey === undefined) { @@ -89,8 +88,8 @@ describe("searchAnthropic headers", () => { } }); - function mockFetch(responseBody: unknown) { - globalThis.fetch = mock(async (url: string | URL | Request, init?: RequestInit) => { + function mockFetch(responseBody: unknown): Disposable { + return hookFetch((url, init) => { capturedRequest = { url: typeof url === "string" ? url : url.toString(), headers: init?.headers, @@ -101,12 +100,12 @@ describe("searchAnthropic headers", () => { status: 200, headers: { "Content-Type": "application/json" }, }); - }) as unknown as typeof fetch; + }); } it("includes web-search beta header and sends API key in X-Api-Key mode", async () => { process.env.ANTHROPIC_SEARCH_API_KEY = "sk-ant-api-test"; - mockFetch(makeAnthropicResponse()); + using _hook = mockFetch(makeAnthropicResponse()); await searchAnthropic({ query: "test api key mode" }); @@ -120,7 +119,7 @@ describe("searchAnthropic headers", () => { it("includes web-search beta header and sends OAuth token in Authorization mode", async () => { process.env.ANTHROPIC_SEARCH_API_KEY = "sk-ant-oat-test"; - mockFetch(makeAnthropicResponse()); + using _hook = mockFetch(makeAnthropicResponse()); await searchAnthropic({ query: "test oauth mode" }); 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 2f8b4bc52..b48a2af9b 100644 --- a/packages/coding-agent/test/tools/web-search-gemini.test.ts +++ b/packages/coding-agent/test/tools/web-search-gemini.test.ts @@ -1,4 +1,5 @@ -import { afterEach, beforeEach, describe, expect, it, mock, vi } from "bun:test"; +import { afterEach, describe, expect, it, vi } from "bun:test"; +import { hookFetch } from "@oh-my-pi/pi-utils"; import { AgentStorage } from "../../src/session/agent-storage"; import { searchGemini } from "../../src/web/search/providers/gemini"; @@ -10,10 +11,9 @@ const SSE_RESPONSE = 'data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"Gemini answer"}]}}],"modelVersion":"gemini-2.5-flash"}}\n\n'; describe("searchGemini tools serialization", () => { - const originalFetch = globalThis.fetch; let capturedRequest: CapturedRequest | null = null; - beforeEach(() => { + function mockGeminiFetch() { capturedRequest = null; vi.spyOn(AgentStorage, "open").mockResolvedValue({ listAuthCredentials: () => [ @@ -29,8 +29,7 @@ describe("searchGemini tools serialization", () => { ], updateAuthCredential: () => undefined, } as unknown as AgentStorage); - - globalThis.fetch = mock(async (_url: string | URL | Request, init?: RequestInit) => { + return hookFetch((_url, init) => { capturedRequest = { body: init?.body ? (JSON.parse(init.body as string) as Record) : null, }; @@ -38,16 +37,16 @@ describe("searchGemini tools serialization", () => { status: 200, headers: { "Content-Type": "text/event-stream" }, }); - }) as unknown as typeof fetch; - }); + }); + } afterEach(() => { vi.restoreAllMocks(); - globalThis.fetch = originalFetch; capturedRequest = null; }); it("sends default googleSearch tool when no passthrough payloads are provided", async () => { + using _hook = mockGeminiFetch(); await searchGemini({ query: "default tools" }); expect(capturedRequest).not.toBeNull(); @@ -57,6 +56,7 @@ describe("searchGemini tools serialization", () => { }); it("passes through google_search payload into googleSearch tool", async () => { + using _hook = mockGeminiFetch(); await searchGemini({ query: "google payload", google_search: { dynamicRetrievalConfig: { mode: "MODE_DYNAMIC" } }, @@ -69,6 +69,7 @@ describe("searchGemini tools serialization", () => { }); it("includes codeExecution and urlContext tools when provided", async () => { + using _hook = mockGeminiFetch(); await searchGemini({ query: "extended tools", code_execution: {}, diff --git a/packages/utils/src/hook-fetch.ts b/packages/utils/src/hook-fetch.ts new file mode 100644 index 000000000..0891bac53 --- /dev/null +++ b/packages/utils/src/hook-fetch.ts @@ -0,0 +1,30 @@ +/** + * 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 31f97afb0..cf1d05c88 100644 --- a/packages/utils/src/index.ts +++ b/packages/utils/src/index.ts @@ -6,6 +6,7 @@ export * from "./env"; export * from "./format"; export * from "./fs-error"; export * from "./glob"; +export * from "./hook-fetch"; export * from "./indent"; export * from "./json"; export * as logger from "./logger";