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.
This commit is contained in:
can1357
2026-06-09 04:09:49 +02:00
parent af33e4055e
commit eb1a46baf5
156 changed files with 2491 additions and 2197 deletions
-51
View File
@@ -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(...)
```
+3 -1
View File
@@ -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
@@ -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) {
+5 -3
View File
@@ -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<Record<string, unknown>>,
instructions: string,
signal?: AbortSignal,
opts?: { fetch?: FetchImpl },
): Promise<OpenAiRemoteCompactionResponse> {
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<RemoteCompactionResponse> {
const response = await fetch(endpoint, {
const response = await (opts?.fetch ?? fetch)(endpoint, {
method: "POST",
headers: { "content-type": "application/json" },
body: JSON.stringify(request),
+4 -1
View File
@@ -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}`,
@@ -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);
@@ -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">> = {}): 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<Record<string, unknown>> | undefined;
using _hook = hookFetch(async (_input, init) => {
const fetchMock: FetchImpl = async (_input, init) => {
const body = JSON.parse(String(init?.body)) as { input: Array<Record<string, unknown>> };
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<Response>();
@@ -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());
+4 -2
View File
@@ -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.
+1
View File
@@ -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.
+18 -1
View File
@@ -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<typeof fetch>[0], init?: Parameters<typeof fetch>[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;
+6 -4
View File
@@ -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<OllamaShowResponse | undefined> {
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;
}
@@ -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<unknown> {
async function fetchModelsDevPayload(fetchImpl: FetchImpl = fetch): Promise<unknown> {
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<OllamaResolvedMetadata>,
fetchImpl: FetchImpl = fetch,
): Promise<Model<"openai-responses">[] | 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<OllamaShowMetadata | undefined> {
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<OllamaResolvedMetadata> {
function createOllamaMetadataResolver(
nativeBaseUrl: string,
fetchImpl?: FetchImpl,
): (modelId: string) => Promise<OllamaResolvedMetadata> {
const cache = new Map<string, Promise<OllamaResolvedMetadata>>();
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<typeof getBundledModels>[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<TApi extends Api>(models: readonly Model<Ap
return references;
}
async function loadModelsDevReferences<TApi extends Api>(): Promise<Map<string, Model<TApi>>> {
async function loadModelsDevReferences<TApi extends Api>(fetchImpl?: FetchImpl): Promise<Map<string, Model<TApi>>> {
try {
const payload = await fetchModelsDevPayload();
const payload = await fetchModelsDevPayload(fetchImpl);
return createModelsDevReferenceMap<TApi>(
mapModelsDevToModels(payload as Record<string, unknown>, 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<typeof getBundledModels>[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<string, unknown> | undefined)
export interface ZenMuxModelManagerConfig {
apiKey?: string;
baseUrl?: string;
fetch?: FetchImpl;
}
export function zenmuxModelManagerOptions(config?: ZenMuxModelManagerConfig): ModelManagerOptions<Api> {
@@ -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
);
},
@@ -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,
});
}
}
+11 -3
View File
@@ -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<void> {
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<void> {
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}`,
+3 -2
View File
@@ -18,7 +18,8 @@ interface KiloDeviceAuthPollResponse {
}
export async function loginKilo(callbacks: OAuthController): Promise<OAuthCredentials> {
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<OAuthCreden
throw new Error("Login cancelled");
}
const pollResponse = await fetch(`${KILO_DEVICE_AUTH_BASE_URL}/codes/${encodeURIComponent(userCode)}`);
const pollResponse = await fetchImpl(`${KILO_DEVICE_AUTH_BASE_URL}/codes/${encodeURIComponent(userCode)}`);
if (pollResponse.status === 202) {
await Bun.sleep(POLL_INTERVAL_MS);
continue;
@@ -1,10 +1,7 @@
import { afterEach, describe, expect, it, vi } from "bun:test";
import { isXAIAccessTokenExpiring, refreshXAIOAuthToken, validateXAIEndpoint, XAIOAuthFlow } from "../xai-oauth";
const originalFetch = global.fetch;
afterEach(() => {
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(
+36 -16
View File
@@ -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<string, string | number>,
fetchImpl: FetchImpl,
extraHeaders?: Record<string, string>,
): Promise<string> {
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<OAuthCreden
/**
* Refresh Anthropic OAuth token
*/
export async function refreshAnthropicToken(refreshToken: string): Promise<OAuthCredentials> {
export async function refreshAnthropicToken(
refreshToken: string,
fetchOverride?: FetchImpl,
): Promise<OAuthCredentials> {
const fetchImpl = fetchOverride ?? fetch;
let responseBody: string;
try {
responseBody = await postJson(
@@ -250,6 +269,7 @@ export async function refreshAnthropicToken(refreshToken: string): Promise<OAuth
client_id: CLIENT_ID,
refresh_token: refreshToken,
},
fetchImpl,
{
// CC sends these on refresh but not on the initial code exchange
"anthropic-beta": "oauth-2025-04-20",
@@ -261,7 +281,7 @@ export async function refreshAnthropicToken(refreshToken: string): Promise<OAuth
}
const data = parseOAuthTokenResponse(responseBody, "token refresh");
const { accountId, email } = await resolveAccountIdentity(data);
const { accountId, email } = await resolveAccountIdentity(data, fetchImpl);
return {
refresh: data.refresh_token || refreshToken,
@@ -3,6 +3,7 @@
*/
import { scheduler } from "node:timers/promises";
import { getBundledModels } from "../../models";
import type { FetchImpl } from "../../types";
import type { OAuthCredentials } from "./types";
const CLIENT_ID = "Ov23li8tweQw6odWQebz";
@@ -15,6 +16,16 @@ export const OPENCODE_HEADERS = {
const INITIAL_POLL_INTERVAL_MULTIPLIER = 1.2;
const SLOW_DOWN_POLL_INTERVAL_MULTIPLIER = 1.4;
type GitHubCopilotLoginOptions = {
onAuth: (url: string, instructions?: string) => void;
onPrompt: (prompt: { message: string; placeholder?: string; allowEmpty?: boolean }) => Promise<string>;
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<unknown> {
const response = await fetch(url, init);
async function fetchJson(url: string, init: RequestInit, fetchImpl: FetchImpl): Promise<unknown> {
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<unknown> {
return response.json();
}
async function startDeviceFlow(domain: string): Promise<DeviceCodeResponse> {
async function startDeviceFlow(domain: string, fetchImpl: FetchImpl): Promise<DeviceCodeResponse> {
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<boolean> {
async function enableGitHubCopilotModel(
token: string,
modelId: string,
fetchImpl: FetchImpl,
enterpriseDomain?: string,
): Promise<boolean> {
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<void> {
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<string>;
onProgress?: (message: string) => void;
signal?: AbortSignal;
pollIntervalFloorMs?: number;
pollIntervalScaleMs?: number;
}): Promise<OAuthCredentials> {
export async function loginGitHubCopilot(options: GitHubCopilotLoginOptions): Promise<OAuthCredentials> {
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;
}
@@ -38,6 +38,7 @@ async function loginMiniMaxCodeWithBaseUrl(
baseUrl: string,
providerName: string,
): Promise<string> {
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;
}
+2
View File
@@ -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<string>;
onPrompt?(prompt: OAuthPrompt): Promise<string>;
signal?: AbortSignal;
fetch?: FetchImpl;
}
export interface OAuthLoginCallbacks extends OAuthController {
+16 -8
View File
@@ -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<XAIOAuthDiscovery> {
async function xaiOAuthDiscovery(
timeoutMs: number = DISCOVERY_TIMEOUT_MS,
fetchOverride?: FetchImpl,
): Promise<XAIOAuthDiscovery> {
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<OAuthCredentials> {
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<OAuthCredent
* re-validates the cached `token_endpoint` on the refresh hot path so a
* cached-but-poisoned endpoint cannot silently leak a refresh_token.
*/
export async function refreshXAIOAuthToken(refreshToken: string): Promise<OAuthCredentials> {
export async function refreshXAIOAuthToken(refreshToken: string, fetchOverride?: FetchImpl): Promise<OAuthCredentials> {
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<OAuthC
refresh_token: refreshToken,
});
const response = await fetch(tokenEndpoint, {
const response = await fetchImpl(tokenEndpoint, {
method: "POST",
headers: {
"Content-Type": "application/x-www-form-urlencoded",
+9 -4
View File
@@ -8,6 +8,7 @@
* login opens plan management so users copy the regional `tp-...` key.
*/
import type { FetchImpl } from "../../types";
import type { OAuthController } from "./types";
const PROVIDER_ID = "xiaomi";
@@ -50,9 +51,11 @@ const VALIDATION_TIMEOUT_MS = 15_000;
async function validateXiaomiApiKey(
apiKey: string,
tokenPlanRegion: XiaomiTokenPlanRegion | undefined,
signal?: AbortSignal,
tokenPlanRegion?: XiaomiTokenPlanRegion,
fetchOverride?: FetchImpl,
): Promise<void> {
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<string> {
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<string> {
}
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<string> {
* Prompts for a token-plan API key and validates it against the selected region.
*/
export async function loginXiaomiTokenPlan(options: OAuthController, region: XiaomiTokenPlanRegion): Promise<string> {
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;
}
+2 -2
View File
@@ -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.
+2 -2
View File
@@ -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<void>;
}
+2 -2
View File
@@ -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;
}
/**
@@ -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<TApi extends Api> {
/** 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.
+2 -1
View File
@@ -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<Response | Error>): { calls: FetchCall[]; fetch: typeof fetch } {
function createFetchMock(responses: Array<Response | Error>): { calls: FetchCall[]; fetch: FetchImpl } {
const calls: FetchCall[] = [];
const fetchImpl = (async (input: string | URL | Request, init?: RequestInit) => {
calls.push({ url: String(input), init: init ?? {} });
+8 -19
View File
@@ -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");
@@ -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<T>(promise: Promise<T>, errorMessage: stri
}
afterEach(() => {
global.fetch = originalFetch;
vi.useRealTimers();
vi.restoreAllMocks();
});
@@ -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");
+10 -9
View File
@@ -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<string, str
});
}
function makeContext(fetchImpl: typeof fetch, retryWait?: UsageFetchContext["retryWait"]): UsageFetchContext {
function makeContext(fetchImpl: FetchImpl, retryWait?: UsageFetchContext["retryWait"]): UsageFetchContext {
return { fetch: fetchImpl, retryWait };
}
@@ -43,7 +44,7 @@ describe("claudeUsageProvider retry contract", () => {
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
+7 -6
View File
@@ -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<typeof fetch>[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<typeof fetch>[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;
+14 -16
View File
@@ -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 */
@@ -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<string | null> = [];
const requestedAuthorizations: Array<string | null> = [];
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<string | null> = [];
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();
@@ -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(""),
});
@@ -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"),
@@ -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<string | null> = [];
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<string | null> = [];
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<string | null> = [];
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<string | null> = [];
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();
@@ -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<UsageFetchParams["credential"]>) {
} 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({}) };
}
@@ -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();
@@ -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();
+5 -12
View File
@@ -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({
+8
View File
@@ -0,0 +1,8 @@
import type { FetchImpl } from "../../src/types";
type FetchHandler = (input: string | URL | Request, init?: RequestInit) => Response | Promise<Response>;
/** 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);
}
+9 -12
View File
@@ -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<Response> {
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 <think> 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, "<think>"),
minimaxChunk(model, "hidden reasoning"),
minimaxChunk(model, "</think>"),
@@ -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 },
+10 -11
View File
@@ -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: "<authenticated>" }).result();
const result = await streamBedrock(model, context, { apiKey: "<authenticated>", fetch: fetchMock }).result();
expect(requestHeaders?.get("authorization")).toBe("Bearer bedrock-api-key");
expect(requestHeaders?.get("authorization")).not.toBe("Bearer <authenticated>");
@@ -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: "<authenticated>" }).result();
const result = await streamBedrock(model, context, { apiKey: "<authenticated>", fetch: fetchMock }).result();
const authorization = requestHeaders?.get("authorization");
expect(authorization).toStartWith("AWS4-HMAC-SHA256 ");
+40 -41
View File
@@ -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<Response> => {
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<Response> => {
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<Response> => {
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<Response> => {
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");
+14 -14
View File
@@ -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<Response> {
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([
+5 -11
View File
@@ -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<Response> => 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<unknown> {
const { promise, resolve } = Promise.withResolvers<unknown>();
global.fetch = mockFetch();
streamOpenAICompletions(model, context, {
apiKey: "test-key",
signal: abortedSignal(),
onPayload: payload => resolve(payload),
...opts,
fetch: mockFetch(),
});
return promise;
}
+13 -19
View File
@@ -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?.();
+23 -20
View File
@@ -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<Response> {
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;
+36 -44
View File
@@ -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", () => {
+10 -19
View File
@@ -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<Response> => {
const fetchMock = (async (input: string | URL | Request, init?: RequestInit): Promise<Response> => {
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<string, unknown>) : {};
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<Response> => {
const fetchMock = (async (input: string | URL | Request): Promise<Response> => {
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<Response> => {
const fetchMock = (async (input: string | URL | Request, init?: RequestInit): Promise<Response> => {
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<string, unknown>) : {};
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
+10 -9
View File
@@ -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);
+1 -7
View File
@@ -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",
+7 -9
View File
@@ -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);
+4 -11
View File
@@ -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<Response> => {
const fetchMock = (async (input: string | Request | URL): Promise<Response> => {
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");
+15 -16
View File
@@ -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<Response> {
function createMockFetch(events: unknown[]): FetchImpl {
return (async (_input: string | URL | Request, _init?: RequestInit): Promise<Response> => {
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)
+3 -8
View File
@@ -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<Response> => {
const fetchMock = (async (input: string | URL | Request, _init?: RequestInit): Promise<Response> => {
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;
+1 -7
View File
@@ -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",
+1 -7
View File
@@ -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() }],
+37 -31
View File
@@ -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"));
+14 -14
View File
@@ -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<Response> {
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)
+5 -10
View File
@@ -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<string, unknown> | undefined;
global.fetch = Object.assign(
const fetchMock: FetchImpl = Object.assign(
async (_input: string | URL | Request, init?: RequestInit): Promise<Response> => {
payload = JSON.parse(typeof init?.body === "string" ? init.body : "{}") as Record<string, unknown>;
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");
+10 -16
View File
@@ -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");
});
});
+7 -5
View File
@@ -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");
+6 -12
View File
@@ -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)");
});
+3 -12
View File
@@ -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();
+12 -22
View File
@@ -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();
});
});
+41 -46
View File
@@ -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<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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<Record<string, unknown>> | 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<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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<Record<string, unknown>> | 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<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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<Record<string, unknown>> | undefined;
@@ -507,54 +502,54 @@ describe("ollama-cloud provider support", () => {
describe("mapToolChoice", () => {
test("omits tool_choice when undefined or auto", async () => {
let requestBody: Record<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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<string, unknown> | undefined;
global.fetch = vi.fn(async (_input, init) => {
const fetchMock: FetchImpl = vi.fn(async (_input, init) => {
requestBody = JSON.parse(String(init?.body ?? "{}")) as Record<string, unknown>;
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");
});
+11 -17
View File
@@ -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);
}
+93 -61
View File
@@ -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<string, ProviderSessionState>();
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<string, unknown> | 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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, unknown>);
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>(),
@@ -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<string, ProviderSessionState>(),
@@ -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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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<string, ProviderSessionState>();
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,
+2 -1
View File
@@ -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;
@@ -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<Response> {
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<unknown> {
const { promise, resolve } = Promise.withResolvers<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
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<unknown>();
global.fetch = createMockFetch(["[DONE]"]);
const fetchMock = createMockFetch(["[DONE]"]);
streamOpenAICompletions(model, baseContext(), {
apiKey: "test-key",
fetch: fetchMock,
signal: createAbortedSignal(),
openrouterVariant: "nitro",
onPayload: payload => resolve(payload),
@@ -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<Record<string, unknown>> {
let payload: Record<string, unknown> | undefined;
global.fetch = Object.assign(
const fetchMock: FetchImpl = Object.assign(
async (_input: string | URL | Request, init?: RequestInit): Promise<Response> => {
payload = JSON.parse(typeof init?.body === "string" ? init.body : "{}") as Record<string, unknown>;
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();
@@ -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();
@@ -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<string, unknown>): Record<string, unknown> {
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();
});
@@ -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<Response> {
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<Response> {
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<void> {
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<void> {
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<void> {
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<void> {
// 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,
);
@@ -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<string, unknown> | 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<string, unknown>) : 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();
});
@@ -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<string, unknown> {
function mockSseFetch(): { fetchMock: FetchImpl; captured: Record<string, unknown> } {
const captured: Record<string, unknown> = {};
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<string, unknown>) : {};
Object.assign(captured, body);
const event = {
@@ -29,8 +27,7 @@ function mockSseFetch(): Record<string, unknown> {
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<Record<string, unknown>> {
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();
});
@@ -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<Record<string, unknown>> {
let captured: Record<string, unknown> = {};
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<string, unknown>) : {};
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();
});
@@ -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<Response> =>
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<Response> => {
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<string, ProviderSessionState>();
const strictFlags: boolean[][] = [];
let attempt = 0;
global.fetch = Object.assign(
const fetchMock: FetchImpl = Object.assign(
async (_input: string | URL | Request, init?: RequestInit): Promise<Response> => {
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<string, ProviderSessionState>();
const strictFlags: boolean[][] = [];
global.fetch = Object.assign(
const fetchMock: FetchImpl = Object.assign(
async (_input: string | URL | Request, init?: RequestInit): Promise<Response> => {
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");
+7 -6
View File
@@ -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<typeof fetch>[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\)/);
@@ -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<never> => {
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",
+5 -10
View File
@@ -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<string, string> = {}): 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<Response> =>
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);
},
+17 -13
View File
@@ -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");
});
+3 -4
View File
@@ -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();
+45 -43
View File
@@ -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<SseChunk | "[DONE]">): Response {
});
}
function mockFetch(events: ReadonlyArray<SseChunk | "[DONE]">): typeof fetch {
function mockFetch(events: ReadonlyArray<SseChunk | "[DONE]">): FetchImpl {
const fn = async (_input: string | URL | Request, _init?: RequestInit): Promise<Response> => 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<unknown>): Response {
});
}
function mockNdjsonFetch(lines: ReadonlyArray<unknown>): typeof fetch {
function mockNdjsonFetch(lines: ReadonlyArray<unknown>): FetchImpl {
const fn = async (_input: string | URL | Request, _init?: RequestInit): Promise<Response> => 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 <thin" }),
chunk(model.id, { content: "k>hidden reasoning</think" }),
chunk(model.id, { content: ">" }),
@@ -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<string, unknown> | 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<string, unknown>;
},
@@ -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");
+4 -10
View File
@@ -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");
+7 -7
View File
@@ -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<typeof fetch>[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<typeof fetch>[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;
+17 -20
View File
@@ -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<Record<string, unknown>>): void {
global.fetch = (async () =>
function mockWaferModelsResponse(entries: Array<Record<string, unknown>>): 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();
+8 -14
View File
@@ -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<string, string>[] = [];
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<string, string>);
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<string, string>[] = [];
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<string, string>);
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);
@@ -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<string, string>; body: string }[] = [];
using _hook = hookFetch((input, init) => {
const fetchMock: FetchImpl = async (input, init) => {
seen.push({
url: String(input),
headers: (init?.headers ?? {}) as Record<string, string>,
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);
+7 -12
View File
@@ -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)");
});
+4 -5
View File
@@ -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" }),
);
+12 -15
View File
@@ -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<Model<"openai-completions">, "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<Response> => {
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<Response> => {
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?.();
+3 -1
View File
@@ -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 `@<upstream>` 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
@@ -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<string, { options: ModelManagerOptions<Api>; 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<OllamaDiscoveredModelMetadata | null> {
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<LlamaCppDiscoveredServerMetadata | null> {
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),
});
@@ -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;
}
@@ -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<OAuthEndpoints | null> {
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;
}
}
+12 -5
View File
@@ -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<string | null> {
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<void> {
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<OAuthCredentials> {
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(),
+10 -2
View File
@@ -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<FetchProvider, () => Promise<string | null>> = {
// 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(),
});

Some files were not shown because too many files have changed in this diff Show More