diff --git a/packages/coding-agent/src/config/model-discovery.ts b/packages/coding-agent/src/config/model-discovery.ts index c569144b4..0d8308dc8 100644 --- a/packages/coding-agent/src/config/model-discovery.ts +++ b/packages/coding-agent/src/config/model-discovery.ts @@ -5,7 +5,7 @@ * `discoverModelsByProviderType` with a `DiscoveryContext`; built-in provider * discovery lives in pi-catalog's provider-models. */ -import type { FetchImpl } from "@oh-my-pi/pi-ai"; +import { type ApiKey, type FetchImpl, withAuth } from "@oh-my-pi/pi-ai"; import type { Api, Model } from "@oh-my-pi/pi-ai/types"; import { buildModel } from "@oh-my-pi/pi-catalog/build"; import { @@ -97,10 +97,12 @@ export interface DiscoveryContext { /** Injected fetch implementation (tests stub this). */ fetch: FetchImpl; /** - * Resolve a provider's API key for `Authorization: Bearer …`. Returns - * undefined when no key is stored or it is a local/no-auth sentinel. + * Resolve a provider's bearer credential for `Authorization: Bearer …`. + * Returns undefined when no key is stored or it is a local/no-auth + * sentinel; otherwise an {@link ApiKey} whose resolver participates in the + * central force-refresh/rotate auth-retry policy on 401/usage-limit. */ - getBearerApiKey(provider: string): Promise; + getBearerApiKeyResolver(provider: string): Promise; } type OllamaDiscoveredModelMetadata = { @@ -314,22 +316,26 @@ export async function discoverLlamaCppModels( const baseUrl = normalizeLlamaCppBaseUrl(providerConfig.baseUrl); const modelsUrl = `${baseUrl}/models`; - const headers: Record = { ...(providerConfig.headers ?? {}) }; - const apiKey = await ctx.getBearerApiKey(providerConfig.provider); - if (apiKey) { - headers.Authorization = `Bearer ${apiKey}`; - } - - const [response, serverMetadata] = await Promise.all([ - ctx.fetch(modelsUrl, { - headers, - signal: AbortSignal.timeout(250), - }), - discoverLlamaCppServerMetadata(ctx, baseUrl, headers), - ]); - if (!response.ok) { - throw new Error(`HTTP ${response.status} from ${modelsUrl}`); - } + const baseHeaders: Record = { ...(providerConfig.headers ?? {}) }; + let headers = baseHeaders; + const attempt = async (h: Record) => { + const [response, metadata] = await Promise.all([ + ctx.fetch(modelsUrl, { + headers: h, + signal: AbortSignal.timeout(250), + }), + discoverLlamaCppServerMetadata(ctx, baseUrl, h), + ]); + if (!response.ok) { + throw new Error(`HTTP ${response.status} from ${modelsUrl}`); + } + headers = h; + return [response, metadata] as const; + }; + const apiKey = await ctx.getBearerApiKeyResolver(providerConfig.provider); + const [response, serverMetadata] = apiKey + ? await withAuth(apiKey, key => attempt({ ...baseHeaders, Authorization: `Bearer ${key}` })) + : await attempt(baseHeaders); const payload = (await response.json()) as { data?: Array<{ id: string }> }; const models = payload.data ?? []; const discovered: Model[] = []; @@ -370,19 +376,23 @@ export async function discoverOpenAIModelsList( const baseUrl = normalizeOpenAIModelsListBaseUrl(providerConfig.baseUrl); const modelsUrl = `${baseUrl}/models`; - const headers: Record = { ...(providerConfig.headers ?? {}) }; - const apiKey = await ctx.getBearerApiKey(providerConfig.provider); - if (apiKey) { - headers.Authorization = `Bearer ${apiKey}`; - } - - const response = await ctx.fetch(modelsUrl, { - headers, - signal: AbortSignal.timeout(10_000), - }); - if (!response.ok) { - throw new Error(`HTTP ${response.status} from ${modelsUrl}`); - } + const baseHeaders: Record = { ...(providerConfig.headers ?? {}) }; + let headers = baseHeaders; + const attempt = async (h: Record) => { + const res = await ctx.fetch(modelsUrl, { + headers: h, + signal: AbortSignal.timeout(10_000), + }); + if (!res.ok) { + throw new Error(`HTTP ${res.status} from ${modelsUrl}`); + } + headers = h; + return res; + }; + const apiKey = await ctx.getBearerApiKeyResolver(providerConfig.provider); + const response = apiKey + ? await withAuth(apiKey, key => attempt({ ...baseHeaders, Authorization: `Bearer ${key}` })) + : await attempt(baseHeaders); const payload = (await response.json()) as { data?: Array<{ id: string }> }; const models = payload.data ?? []; const discovered: Model[] = []; @@ -435,19 +445,23 @@ export async function discoverProxyModels( const baseUrl = normalizeOpenAIModelsListBaseUrl(providerConfig.baseUrl); const modelsUrl = `${baseUrl}/models`; - const headers: Record = { ...(providerConfig.headers ?? {}) }; - const apiKey = await ctx.getBearerApiKey(providerConfig.provider); - if (apiKey) { - headers.Authorization = `Bearer ${apiKey}`; - } - - const response = await ctx.fetch(modelsUrl, { - headers, - signal: AbortSignal.timeout(10_000), - }); - if (!response.ok) { - throw new Error(`HTTP ${response.status} from ${modelsUrl}`); - } + const baseHeaders: Record = { ...(providerConfig.headers ?? {}) }; + let headers = baseHeaders; + const attempt = async (h: Record) => { + const res = await ctx.fetch(modelsUrl, { + headers: h, + signal: AbortSignal.timeout(10_000), + }); + if (!res.ok) { + throw new Error(`HTTP ${res.status} from ${modelsUrl}`); + } + headers = h; + return res; + }; + const apiKey = await ctx.getBearerApiKeyResolver(providerConfig.provider); + const response = apiKey + ? await withAuth(apiKey, key => attempt({ ...baseHeaders, Authorization: `Bearer ${key}` })) + : await attempt(baseHeaders); const payload = (await response.json()) as { data?: Array<{ id?: string; name?: string; supported_endpoint_types?: string[] }>; }; diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 0882807ef..2c713865c 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -1238,9 +1238,10 @@ export class ModelRegistry { #discoveryContext(): DiscoveryContext { return { fetch: this.#fetch, - getBearerApiKey: async provider => { + getBearerApiKeyResolver: async provider => { const apiKey = await this.getApiKeyForProvider(provider); - return apiKey && apiKey !== DEFAULT_LOCAL_TOKEN && apiKey !== kNoAuth ? apiKey : undefined; + if (!apiKey || apiKey === DEFAULT_LOCAL_TOKEN || apiKey === kNoAuth) return undefined; + return this.resolver(provider); }, }; } diff --git a/packages/coding-agent/src/tools/tts.ts b/packages/coding-agent/src/tools/tts.ts index 39fc11bd9..e1f091fe6 100644 --- a/packages/coding-agent/src/tools/tts.ts +++ b/packages/coding-agent/src/tools/tts.ts @@ -1,6 +1,7 @@ // Ported from NousResearch/hermes-agent (MIT) — tools/tts_tool.py L167-171, L896-959. import type { AgentToolResult } from "@oh-my-pi/pi-agent-core"; +import { type ApiKey, ProviderHttpError, withAuth } from "@oh-my-pi/pi-ai"; import * as z from "zod/v4"; import type { CustomTool, CustomToolContext } from "../extensibility/custom-tools/types"; import { ohMyPiXAIUserAgent, resolveXAIHttpCredentials } from "../lib/xai-http"; @@ -96,27 +97,46 @@ export const ttsTool: CustomTool = { const timeoutSignal = AbortSignal.timeout(60_000); const combinedSignal = signal ? AbortSignal.any([signal, timeoutSignal]) : timeoutSignal; - const response = await fetch(`${creds.baseURL}/tts`, { - method: "POST", - headers: { - Authorization: `Bearer ${creds.apiKey}`, - "Content-Type": "application/json", - "User-Agent": ohMyPiXAIUserAgent(), - }, - body: JSON.stringify(payload), - signal: combinedSignal, + const sessionId = ctx.sessionManager.getSessionId(); + const apiKey: ApiKey = ctx.modelRegistry.resolver(creds.provider, { + sessionId, + baseUrl: creds.baseURL, }); - if (!response.ok) { - const detail = await response.text(); - return { - isError: true, - content: [ - { - type: "text", - text: `xAI TTS failed (${response.status}): ${detail.slice(0, 300)}`, - }, - ], - }; + + let response: Response; + try { + response = await withAuth( + apiKey, + async key => { + const resp = await fetch(`${creds.baseURL}/tts`, { + method: "POST", + headers: { + Authorization: `Bearer ${key}`, + "Content-Type": "application/json", + "User-Agent": ohMyPiXAIUserAgent(), + }, + body: JSON.stringify(payload), + signal: combinedSignal, + }); + if (!resp.ok) { + const detail = await resp.text(); + throw new ProviderHttpError(`xAI TTS failed (${resp.status}): ${detail.slice(0, 300)}`, resp.status, { + headers: resp.headers, + }); + } + return resp; + }, + { signal: combinedSignal }, + ); + } catch (error) { + const status = (error as { status?: unknown }).status; + if (error instanceof Error && typeof status === "number") { + return { + isError: true, + content: [{ type: "text", text: error.message }], + }; + } + throw error; } const bytes = new Uint8Array(await response.arrayBuffer()); await Bun.write(outputPath, bytes); diff --git a/packages/coding-agent/src/web/kagi.ts b/packages/coding-agent/src/web/kagi.ts index 38041913b..bc78651d6 100644 --- a/packages/coding-agent/src/web/kagi.ts +++ b/packages/coding-agent/src/web/kagi.ts @@ -6,7 +6,7 @@ * through the shared {@link AuthStorage} broker (Bearer token), and responses * are categorized result buckets rather than the legacy flat object array. */ -import type { AuthStorage, FetchImpl } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, type FetchImpl, withAuth } from "@oh-my-pi/pi-ai"; import { withHardTimeout } from "./search/providers/utils"; const KAGI_SEARCH_URL = "https://kagi.com/api/v1/search"; @@ -173,14 +173,6 @@ export interface KagiSearchResult { answer?: string; } -export async function findKagiApiKey( - authStorage: AuthStorage, - sessionId?: string, - signal?: AbortSignal, -): Promise { - return (await authStorage.getApiKey("kagi", sessionId, { signal })) ?? null; -} - /** * Compute a YYYY-MM-DD date string `recency` units before now, in UTC. * UTC keeps the recency window deterministic regardless of host timezone and @@ -247,27 +239,34 @@ export async function searchWithKagi( options: KagiSearchOptions = {}, authStorage: AuthStorage, ): Promise { - const apiKey = await findKagiApiKey(authStorage, options.sessionId, options.signal); - if (!apiKey) { - throw new KagiApiError("Kagi credentials not found. Set KAGI_API_KEY or login with 'omp /login kagi'."); - } - const fetchImpl = options.fetch ?? fetch; + const body = JSON.stringify(buildRequestBody(query, options)); - const response = await fetchImpl(KAGI_SEARCH_URL, { - method: "POST", - headers: { - Authorization: `Bearer ${apiKey}`, - "Content-Type": "application/json", - Accept: "application/json", + const response = await withAuth( + authStorage.resolver("kagi", { sessionId: options.sessionId }), + async apiKey => { + const res = await fetchImpl(KAGI_SEARCH_URL, { + method: "POST", + headers: { + Authorization: `Bearer ${apiKey}`, + "Content-Type": "application/json", + Accept: "application/json", + }, + body, + signal: withHardTimeout(options.signal), + }); + + if (!res.ok) { + throw parseKagiErrorResponse(res.status, await res.text()); + } + + return res; }, - body: JSON.stringify(buildRequestBody(query, options)), - signal: withHardTimeout(options.signal), - }); - - if (!response.ok) { - throw parseKagiErrorResponse(response.status, await response.text()); - } + { + signal: options.signal, + missingKeyMessage: "Kagi credentials not found. Set KAGI_API_KEY or login with 'omp /login kagi'.", + }, + ); const payload = (await response.json()) as KagiSearchResponse; if (payload.error && payload.error.length > 0) { diff --git a/packages/coding-agent/test/tools/web-search-kagi.test.ts b/packages/coding-agent/test/tools/web-search-kagi.test.ts index 1a3497744..d6d777a05 100644 --- a/packages/coding-agent/test/tools/web-search-kagi.test.ts +++ b/packages/coding-agent/test/tools/web-search-kagi.test.ts @@ -8,6 +8,9 @@ const fakeAuthStorage = { async getApiKey() { return process.env.KAGI_API_KEY ?? undefined; }, + resolver() { + return async () => process.env.KAGI_API_KEY ?? undefined; + }, hasAuth() { return Boolean(process.env.KAGI_API_KEY); },