From e6893515978ef1bd57ee647b46e8908279b9d567 Mon Sep 17 00:00:00 2001 From: can1357 Date: Tue, 26 May 2026 03:59:13 +0200 Subject: [PATCH] fix(coding-agent): resolved OAuth token expiry flow in AuthStorage - Centralized OAuth access lifecycle in `AuthStorage`, returning identity metadata and new access-result types. - Added 60-second skew and strict expiry checks, returning undefined/throws for stale or expired OAuth credentials. - Removed provider-local token refresh flows from Gemini, Gemini CLI, Antigravity, Kimi, and related OAuth helpers. - Migrated web-search providers from `AgentStorage` to `AuthStorage` session-aware lookup with `authStorage`/`sessionId`/`signal` flow. - Replaced `findAnthropicAuth`/DB auth lookup with `buildAnthropicAuthConfig` and explicit base-url override/env fallback ordering. --- packages/ai/CHANGELOG.md | 17 ++ packages/ai/src/auth-storage.ts | 123 ++++++-- .../ai/src/providers/google-gemini-cli.ts | 53 ++-- packages/ai/src/usage/gemini.ts | 28 +- packages/ai/src/usage/google-antigravity.ts | 25 +- packages/ai/src/usage/kimi.ts | 23 +- packages/ai/src/utils/anthropic-auth.ts | 165 ++--------- packages/ai/src/utils/oauth/index.ts | 40 +-- packages/ai/test/anthropic-oauth.test.ts | 155 ++++------ .../auth-storage-broker-no-sentinel.test.ts | 104 +++++++ .../ai/test/auth-storage-refresh-skew.test.ts | 128 ++++++++ .../test/google-gemini-cli-alignment.test.ts | 6 + packages/coding-agent/CHANGELOG.md | 4 + packages/coding-agent/src/web/kagi.ts | 18 +- packages/coding-agent/src/web/search/index.ts | 72 +++-- .../coding-agent/src/web/search/provider.ts | 8 +- .../src/web/search/providers/anthropic.ts | 57 ++-- .../src/web/search/providers/base.ts | 40 ++- .../src/web/search/providers/brave.ts | 7 +- .../src/web/search/providers/codex.ts | 275 ++++++++---------- .../src/web/search/providers/exa.ts | 8 +- .../src/web/search/providers/gemini.ts | 234 +++++---------- .../src/web/search/providers/jina.ts | 7 +- .../src/web/search/providers/kagi.ts | 47 ++- .../src/web/search/providers/kimi.ts | 62 ++-- .../src/web/search/providers/parallel.ts | 168 ++++++++++- .../src/web/search/providers/perplexity.ts | 97 +++--- .../src/web/search/providers/searxng.ts | 7 +- .../src/web/search/providers/synthetic.ts | 44 ++- .../src/web/search/providers/tavily.ts | 47 +-- .../src/web/search/providers/zai.ts | 42 +-- .../test/tools/web-search-codex.test.ts | 103 ++++--- .../test/tools/web-search-gemini.test.ts | 71 +++-- .../test/tools/web-search-kagi.test.ts | 23 +- .../test/tools/web-search-parallel.test.ts | 22 +- .../test/tools/web-search-tavily.test.ts | 34 ++- .../test/web/search/abort-and-timeout.test.ts | 6 +- .../test/web/search/codex-broker.test.ts | 62 ++++ .../test/web/search/tavily.test.ts | 37 ++- 39 files changed, 1409 insertions(+), 1060 deletions(-) create mode 100644 packages/ai/test/auth-storage-broker-no-sentinel.test.ts create mode 100644 packages/ai/test/auth-storage-refresh-skew.test.ts create mode 100644 packages/coding-agent/test/web/search/codex-broker.test.ts diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index fd86a7790..fe3aa47fa 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,9 +1,26 @@ # Changelog ## [Unreleased] +### Breaking Changes + +- Removed `findAnthropicAuth` from `anthropic-auth` and replaced store-driven auth discovery with `buildAnthropicAuthConfig`, requiring callers to provide an already-resolved API key before building Anthropic auth config + +### Added + +- Added `AuthStorage.getOAuthAccess` to return a refreshed OAuth access token with identity metadata (`accountId`, `email`, `projectId`, `enterpriseUrl`) for callers that need bearer-token headers together + +### Changed + +- Changed OAuth selection in `AuthStorage` to treat credentials as stale when they are within 60 seconds of expiry and rotate them preemptively +- Changed Google Gemini CLI, Google Gemini usage, Antigravity usage, and Kimi usage flows to stop refreshing OAuth tokens directly and rely on `AuthStorage` for token rotation + +### Removed + +- Removed provider-local OAuth refresh helpers from Google Gemini CLI and Google/Kimi/Antigravity usage probes, preventing direct refresh calls from those usage paths ### Fixed +- Fixed expired OAuth handling so provider-level paths no longer attempt direct token refresh calls for expired credentials and instead rely on `AuthStorage` for rotation - Fixed Claude Opus 4.7 on Amazon Bedrock streaming no reasoning output (and appearing to hang on long reasoning runs) because Anthropic silently switched the adaptive-thinking display default to `"omitted"`. The Bedrock provider now sends `thinking.display = "summarized"` by default on Opus 4.7+ adaptive models and on budget-based Claude models, mirroring the existing direct-Anthropic behavior. `BedrockOptions.thinkingDisplay` (`"summarized" | "omitted"`) is exposed for callers that want to opt out, and `hideThinkingSummary` now wires through to the Bedrock case ([#1373](https://github.com/can1357/oh-my-pi/issues/1373)). ## [15.3.2] - 2026-05-25 diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index 58a477887..4ff6265e8 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -394,6 +394,19 @@ const USAGE_FAILURE_BACKOFF_MS = 10_000; // (~3.5s total worst case); a tight per-request budget aborts retries mid-cycle. const DEFAULT_USAGE_REQUEST_TIMEOUT_MS = 10_000; const DEFAULT_OAUTH_REFRESH_TIMEOUT_MS = 10_000; +/** + * Refresh OAuth access tokens this many ms before their stated expiry. The + * skew exists so callers downstream of {@link AuthStorage} (stream providers, + * usage probes, web_search) never observe a credential that is expired or + * about to expire mid-request — there's a single rotation point and everyone + * downstream trusts the token they receive. + * + * Set to 60s: comfortably absorbs request RTT + a clock-skew window without + * triggering a refresh on every request. Provider token endpoints typically + * mint access tokens with 30-60min lifetimes, so refreshing 60s early changes + * the rotation cadence by <4%. + */ +const OAUTH_REFRESH_SKEW_MS = 60_000; /** * Cap on the buffered credential_disabled backlog held while no handler is attached. * In practice the backlog is 0–N where N ≈ active providers (≤ ~20). The cap exists so @@ -429,6 +442,23 @@ type AuthApiKeyOptions = { */ signal?: AbortSignal; }; +type OAuthResolutionResult = { apiKey: string; credential: OAuthCredential }; + +/** + * Refreshed OAuth access plus identity metadata returned by + * {@link AuthStorage.getOAuthAccess}. Callers that authenticate via a bearer + * AND need the credential's identity (Codex `chatgpt-account-id`, Google + * `projectId`, GitHub `enterpriseUrl`) consume this shape directly; the + * refresh slot is deliberately omitted because rotating refresh tokens never + * leave {@link AuthStorage}. + */ +export interface OAuthAccess { + accessToken: string; + accountId?: string; + email?: string; + projectId?: string; + enterpriseUrl?: string; +} export interface InvalidateCredentialMatchingOptions { signal?: AbortSignal; sessionId?: string; @@ -2501,15 +2531,19 @@ export class AuthStorage { } /** - * Resolves an OAuth API key, trying credentials in priority order. + * Resolves an OAuth credential, trying credentials in priority order. * Skips blocked credentials and checks usage limits for providers with usage data. * Falls back to earliest-unblocking credential if all are blocked. + * + * Returns both the API key bytes for outbound requests AND the refreshed + * {@link OAuthCredential} so callers needing identity metadata (account id, + * project id, etc.) do not have to dereference the snapshot themselves. */ - async #resolveOAuthApiKey( + async #resolveOAuthSelection( provider: string, sessionId?: string, options?: AuthApiKeyOptions, - ): Promise { + ): Promise { const credentials = this.#getCredentialsForProvider(provider) .map((credential, index) => ({ credential, index })) .filter((entry): entry is { credential: OAuthCredential; index: number } => entry.credential.type === "oauth"); @@ -2550,9 +2584,9 @@ export class AuthStorage { } await Promise.all( candidates.map(async candidate => { - if (Date.now() < candidate.selection.credential.expires) return; + if (Date.now() + OAUTH_REFRESH_SKEW_MS < candidate.selection.credential.expires) return; const latestCredential = this.#getCredentialsForProvider(provider)[candidate.selection.index]; - if (latestCredential?.type === "oauth" && Date.now() < latestCredential.expires) { + if (latestCredential?.type === "oauth" && Date.now() + OAUTH_REFRESH_SKEW_MS < latestCredential.expires) { candidate.selection.credential = latestCredential; return; } @@ -2583,14 +2617,21 @@ export class AuthStorage { const fallback = candidates[0]; for (const candidate of candidates) { - const apiKey = await this.#tryOAuthCredential(provider, candidate.selection, providerKey, sessionId, options, { - checkUsage, - allowBlocked: false, - prefetchedUsage: candidate.usage, - usagePrechecked: candidate.usageChecked, - enforceProRequirement, - }); - if (apiKey) return apiKey; + const resolved = await this.#tryOAuthCredential( + provider, + candidate.selection, + providerKey, + sessionId, + options, + { + checkUsage, + allowBlocked: false, + prefetchedUsage: candidate.usage, + usagePrechecked: candidate.usageChecked, + enforceProRequirement, + }, + ); + if (resolved) return resolved; } if (fallback && this.#isCredentialBlocked(providerKey, fallback.selection.index)) { @@ -2616,7 +2657,7 @@ export class AuthStorage { const existing = this.#oauthCredentialRefreshInFlight.get(credentialId); if (existing) return raceCredentialRefreshWithSignal(existing, signal); } - if (Date.now() < credential.expires) return credential; + if (Date.now() + OAUTH_REFRESH_SKEW_MS < credential.expires) return credential; if (credentialId === undefined) { return this.#refreshOAuthCredentialUnshared(provider, credential, undefined, signal); } @@ -2719,7 +2760,7 @@ export class AuthStorage { usagePrechecked?: boolean; enforceProRequirement?: boolean; }, - ): Promise { + ): Promise { const { checkUsage, allowBlocked, @@ -2820,7 +2861,7 @@ export class AuthStorage { } } this.#recordSessionCredential(provider, sessionId, "oauth", selection.index); - return result.apiKey; + return { apiKey: result.apiKey, credential: updated }; } catch (error) { const errorMsg = String(error); // Only remove credentials for definitive auth failures @@ -2855,7 +2896,7 @@ export class AuthStorage { credentialId, }); await this.reload(); - return this.getApiKey(provider, sessionId, options); + return this.#resolveOAuthSelection(provider, sessionId, options); } } // Permanently disable invalid credentials with an explicit cause for inspection/debugging. @@ -2874,10 +2915,10 @@ export class AuthStorage { index: selection.index, }); await this.reload(); - return this.getApiKey(provider, sessionId, options); + return this.#resolveOAuthSelection(provider, sessionId, options); } if (this.#getCredentialsForProvider(provider).some(credential => credential.type === "oauth")) { - return this.getApiKey(provider, sessionId, options); + return this.#resolveOAuthSelection(provider, sessionId, options); } } else { // Block temporarily for transient failures (5 minutes) @@ -2964,9 +3005,9 @@ export class AuthStorage { return this.#configValueResolver(apiKeySelection.credential.key); } - const oauthKey = await this.#resolveOAuthApiKey(provider, sessionId, options); - if (oauthKey) { - return oauthKey; + const oauthResolved = await this.#resolveOAuthSelection(provider, sessionId, options); + if (oauthResolved) { + return oauthResolved.apiKey; } // Fall back to environment variable or custom resolver. If we reach here after @@ -2981,6 +3022,44 @@ export class AuthStorage { return this.#fallbackResolver?.(provider) ?? undefined; } + /** + * Resolve the OAuth credential for `provider`, refreshing through the same + * pipeline as {@link AuthStorage.getApiKey} but returning the refreshed + * {@link OAuthAccess} (raw access token + identity metadata) instead of + * the API-key bytes. + * + * Use this when the caller needs to inject identity headers alongside the + * bearer (Codex `chatgpt-account-id`, Google `project`, GitHub + * `enterpriseUrl`). For pure "give me the bytes for `Authorization`" + * scenarios, prefer {@link AuthStorage.getApiKey}. + * + * Returns `undefined` when no OAuth credential is available, the + * credential fails to refresh, or runtime/config overrides have replaced + * OAuth with an explicit API key. + */ + async getOAuthAccess( + provider: string, + sessionId?: string, + options?: AuthApiKeyOptions, + ): Promise { + // Runtime / config overrides intentionally short-circuit OAuth: when the + // user has pinned an API key, they expect the OAuth identity to be + // suppressed (same contract as `getOAuthAccountId`). + if (this.#runtimeOverrides.has(provider) || this.#configOverrides.has(provider)) { + return undefined; + } + const resolved = await this.#resolveOAuthSelection(provider, sessionId, options); + if (!resolved) return undefined; + const { credential } = resolved; + return { + accessToken: credential.access, + accountId: credential.accountId, + email: credential.email, + projectId: credential.projectId, + enterpriseUrl: credential.enterpriseUrl, + }; + } + #extractStructuredApiKeyToken(apiKey: string): string | undefined { if (!apiKey.startsWith("{")) return undefined; try { diff --git a/packages/ai/src/providers/google-gemini-cli.ts b/packages/ai/src/providers/google-gemini-cli.ts index 50b305282..82300a84d 100644 --- a/packages/ai/src/providers/google-gemini-cli.ts +++ b/packages/ai/src/providers/google-gemini-cli.ts @@ -21,8 +21,8 @@ import type { import { normalizeSystemPrompts } from "../utils"; import { AssistantMessageEventStream } from "../utils/event-stream"; import { appendRawHttpRequestDumpFor400, type RawHttpRequestDump, withHttpStatus } from "../utils/http-inspector"; -import { refreshAntigravityToken } from "../utils/oauth/google-antigravity"; -import { refreshGoogleCloudToken } from "../utils/oauth/google-gemini-cli"; +// Refresh is the sole responsibility of AuthStorage (broker-aware, single-flighted); +// the stream provider trusts the access token threaded through `options.apiKey`. import { normalizeSchemaForCCA } from "../utils/schema"; import { ANTIGRAVITY_SYSTEM_INSTRUCTION, getAntigravityUserAgent, getGeminiCliHeaders } from "./google-gemini-headers"; import type { Content, FunctionCallingConfigMode, ThinkingConfig } from "./google-shared"; @@ -191,38 +191,6 @@ export function shouldRefreshGeminiCliCredentials( return nowMs + skewMs >= expiresAt; } -async function refreshGeminiCliCredentialsIfNeeded( - credentials: ParsedGeminiCliCredentials, - isAntigravity: boolean, -): Promise { - if (!credentials.refreshToken || !shouldRefreshGeminiCliCredentials(credentials.expiresAt, isAntigravity)) { - return credentials; - } - - try { - const refreshed = isAntigravity - ? await refreshAntigravityToken(credentials.refreshToken, credentials.projectId) - : await refreshGoogleCloudToken(credentials.refreshToken, credentials.projectId); - return { - accessToken: refreshed.access, - projectId: credentials.projectId, - refreshToken: refreshed.refresh, - expiresAt: refreshed.expires, - }; - } catch (error) { - const reason = error instanceof Error ? error.message : String(error); - // Permanent auth failure (revoked/invalid token) — re-authentication required. - // Google returns 400 invalid_grant when a token is revoked or expired server-side. - if (/invalid_grant|invalid_token|token.*revoked|account.*disabled/i.test(reason)) { - throw new Error(`OAuth token has been revoked or invalidated. Use /login to re-authenticate. (${reason})`); - } - // Transient failure (network, 5xx) — fall back to existing token if not yet expired. - if (credentials.expiresAt !== undefined && Date.now() >= credentials.expiresAt) { - throw new Error(`OAuth token refresh failed before request: ${reason}`); - } - return credentials; - } -} interface CloudCodeAssistRequest { project: string; model: string; @@ -321,8 +289,21 @@ export const streamGoogleGeminiCli: StreamFunction<"google-gemini-cli"> = ( const isAntigravity = model.provider === "google-antigravity"; const parsedCredentials = parseGeminiCliCredentials(apiKeyRaw); - const activeCredentials = await refreshGeminiCliCredentialsIfNeeded(parsedCredentials, isAntigravity); - const { accessToken, projectId } = activeCredentials; + // AuthStorage already refreshed credentials before threading them + // here (see {@link OAUTH_REFRESH_SKEW_MS}). If the credential lands + // expired we bail rather than POSTing a stale token; the next call + // — driven by AuthStorage's invalidate+retry path — will carry a + // fresh credential. + if ( + shouldRefreshGeminiCliCredentials(parsedCredentials.expiresAt, isAntigravity) && + parsedCredentials.expiresAt !== undefined && + Date.now() >= parsedCredentials.expiresAt + ) { + throw new Error( + "OAuth token expired before request — please retry; AuthStorage will refresh on the next attempt.", + ); + } + const { accessToken, projectId } = parsedCredentials; const baseUrl = model.baseUrl?.trim(); const endpoints = baseUrl ? [baseUrl] : isAntigravity ? ANTIGRAVITY_ENDPOINT_FALLBACKS : [DEFAULT_ENDPOINT]; diff --git a/packages/ai/src/usage/gemini.ts b/packages/ai/src/usage/gemini.ts index 89cdaf95b..525c134e7 100644 --- a/packages/ai/src/usage/gemini.ts +++ b/packages/ai/src/usage/gemini.ts @@ -8,7 +8,8 @@ import type { UsageReport, UsageWindow, } from "../usage"; -import { refreshGoogleCloudToken } from "../utils/oauth/google-gemini-cli"; + +// (Refresh is the sole responsibility of AuthStorage; no provider-direct refresh here.) const DEFAULT_ENDPOINT = "https://cloudcode-pa.googleapis.com"; @@ -100,20 +101,21 @@ function buildAmount(remainingFraction: number | undefined): UsageAmount { }; } -async function resolveAccessToken(params: UsageFetchParams, ctx: UsageFetchContext): Promise { +/** + * Return the OAuth access token to use against `/v1internal:*`. AuthStorage is + * the sole refresh authority (broker-aware, single-flighted, rotation-safe); + * if the token landed here expired or near-expired, the next usage cycle will + * carry a freshly-refreshed credential. Returning `undefined` short-circuits + * the probe rather than POSTing a stale token to Google. + */ +function resolveAccessToken(params: UsageFetchParams): string | undefined { const { credential } = params; if (credential.type !== "oauth") return undefined; - if (credential.accessToken && (!credential.expiresAt || credential.expiresAt > Date.now() + 60_000)) { - return credential.accessToken; - } - if (!credential.refreshToken || !credential.projectId) return credential.accessToken; - try { - const refreshed = await refreshGoogleCloudToken(credential.refreshToken, credential.projectId); - return refreshed.access; - } catch (error) { - ctx.logger?.warn("Gemini CLI token refresh failed", { error: String(error) }); - return credential.accessToken; + if (!credential.accessToken) return undefined; + if (credential.expiresAt !== undefined && credential.expiresAt <= Date.now()) { + return undefined; } + return credential.accessToken; } async function loadCodeAssist( @@ -191,7 +193,7 @@ export const googleGeminiCliUsageProvider: UsageProvider = { if (credential.type !== "oauth") { return null; } - const accessToken = await resolveAccessToken(params, ctx); + const accessToken = resolveAccessToken(params); if (!accessToken) { return null; } diff --git a/packages/ai/src/usage/google-antigravity.ts b/packages/ai/src/usage/google-antigravity.ts index 7f4f677cc..26963d7b5 100644 --- a/packages/ai/src/usage/google-antigravity.ts +++ b/packages/ai/src/usage/google-antigravity.ts @@ -9,7 +9,8 @@ import type { UsageStatus, UsageWindow, } from "../usage"; -import { refreshAntigravityToken } from "../utils/oauth/google-antigravity"; + +// (Refresh is the sole responsibility of AuthStorage; no provider-direct refresh here.) interface AntigravityQuotaInfo { remainingFraction?: number; @@ -101,19 +102,19 @@ function normalizeQuotaInfos(info: AntigravityModelInfo): AntigravityQuotaInfo[] return results; } -async function resolveAccessToken(params: UsageFetchParams, ctx: UsageFetchContext): Promise { +/** + * Return the OAuth access token to use against `/v1internal:*`. AuthStorage is + * the sole refresh authority (broker-aware, single-flighted, rotation-safe); + * an expired token short-circuits the probe rather than POSTing the broker + * sentinel back to Google. + */ +function resolveAccessToken(params: UsageFetchParams): string | undefined { const { credential } = params; - if (credential.accessToken && (!credential.expiresAt || credential.expiresAt > Date.now() + 60_000)) { - return credential.accessToken; - } - if (!credential.refreshToken || !credential.projectId) return undefined; - try { - const refreshed = await refreshAntigravityToken(credential.refreshToken, credential.projectId); - return refreshed.access; - } catch (error) { - ctx.logger?.warn("Antigravity usage token refresh failed", { error: String(error) }); + if (!credential.accessToken) return undefined; + if (credential.expiresAt !== undefined && credential.expiresAt <= Date.now()) { return undefined; } + return credential.accessToken; } async function fetchAntigravityUsage(params: UsageFetchParams, ctx: UsageFetchContext): Promise { @@ -122,7 +123,7 @@ async function fetchAntigravityUsage(params: UsageFetchParams, ctx: UsageFetchCo const nowMs = Date.now(); - const accessToken = await resolveAccessToken(params, ctx); + const accessToken = resolveAccessToken(params); if (!accessToken) return null; const baseUrl = params.baseUrl?.replace(/\/+$/, "") || DEFAULT_ENDPOINT; diff --git a/packages/ai/src/usage/kimi.ts b/packages/ai/src/usage/kimi.ts index f36c7f82d..edbfe4549 100644 --- a/packages/ai/src/usage/kimi.ts +++ b/packages/ai/src/usage/kimi.ts @@ -10,7 +10,8 @@ import type { UsageWindow, } from "../usage"; import { isRecord } from "../utils"; -import { getKimiCommonHeaders, refreshKimiToken } from "../utils/oauth/kimi"; +import { getKimiCommonHeaders } from "../utils/oauth/kimi"; +// (Refresh is the sole responsibility of AuthStorage; no provider-direct refresh here.) import { toNumber } from "./shared"; const DEFAULT_BASE_URL = "https://api.kimi.com/coding/v1"; @@ -213,23 +214,17 @@ export const kimiUsageProvider: UsageProvider = { const { credential } = params; if (credential.type !== "oauth") return null; - let accessToken = credential.accessToken; + const accessToken = credential.accessToken; if (!accessToken) return null; const nowMs = Date.now(); + // AuthStorage refreshes OAuth credentials pre-emptively (60s skew). If the + // usage probe lands with an expired token, short-circuit rather than POST + // the broker sentinel back to Kimi — the next cycle will carry a freshly + // refreshed credential. if (credential.expiresAt !== undefined && credential.expiresAt <= nowMs) { - if (!credential.refreshToken) { - ctx.logger?.warn("Kimi usage token expired, no refresh token", { provider: params.provider }); - return null; - } - try { - ctx.logger?.debug("Kimi usage token expired, refreshing", { provider: params.provider }); - const refreshed = await refreshKimiToken(credential.refreshToken); - accessToken = refreshed.access; - } catch (error) { - ctx.logger?.warn("Kimi usage token refresh failed", { provider: params.provider, error: String(error) }); - return null; - } + ctx.logger?.debug("Kimi usage token expired; skipping probe", { provider: params.provider }); + return null; } const baseUrl = normalizeBaseUrl(params.baseUrl); diff --git a/packages/ai/src/utils/anthropic-auth.ts b/packages/ai/src/utils/anthropic-auth.ts index a00b20a50..280b3e4db 100644 --- a/packages/ai/src/utils/anthropic-auth.ts +++ b/packages/ai/src/utils/anthropic-auth.ts @@ -1,20 +1,18 @@ /** * Anthropic Authentication * - * 5-tier auth resolution: - * 1. ANTHROPIC_SEARCH_API_KEY / ANTHROPIC_SEARCH_BASE_URL env vars - * 2. ANTHROPIC_FOUNDRY_API_KEY override when Foundry mode is enabled - * 3. OAuth credentials in ~/.omp/agent/agent.db (with expiry check) - * 4. API key credentials in ~/.omp/agent/agent.db - * 5. Generic Anthropic fallback (ANTHROPIC_API_KEY / ANTHROPIC_BASE_URL) + * Thin helper for turning an already-resolved API key into the request-shaping + * config consumed by {@link buildAnthropicSearchHeaders} / {@link buildAnthropicUrl}. + * + * Credential storage and refresh live in `AuthStorage` — call + * `authStorage.getApiKey("anthropic", sessionId)` first, then pass the result + * through {@link buildAnthropicAuthConfig} for header/URL shaping. */ -import { $env, getAgentDbPath } from "@oh-my-pi/pi-utils"; -import { type AuthCredential, type AuthCredentialStore, SqliteAuthCredentialStore } from "../auth-storage"; +import { $env } from "@oh-my-pi/pi-utils"; import { buildAnthropicHeaders as buildProviderAnthropicHeaders, normalizeAnthropicBaseUrl, } from "../providers/anthropic"; -import { getEnvApiKey } from "../stream"; import { isFoundryEnabled } from "./foundry"; /** Auth configuration for Anthropic */ @@ -24,22 +22,14 @@ export interface AnthropicAuthConfig { isOAuth: boolean; } -/** OAuth credential for Anthropic API access */ -export interface AnthropicOAuthCredential { - type: "oauth"; - access: string; - refresh?: string; - /** Expiry timestamp in milliseconds */ - expires: number; -} - const DEFAULT_BASE_URL = "https://api.anthropic.com"; function normalizeBaseUrl(baseUrl: string | undefined): string | undefined { const trimmed = baseUrl?.trim(); return trimmed ? trimmed.replace(/\/+$/, "") : undefined; } -function resolveAnthropicBaseUrlFromEnv(): string | undefined { + +export function resolveAnthropicBaseUrlFromEnv(): string | undefined { if (isFoundryEnabled()) { const foundryBaseUrl = normalizeBaseUrl($env.FOUNDRY_BASE_URL); if (foundryBaseUrl) return foundryBaseUrl; @@ -50,141 +40,32 @@ function resolveAnthropicBaseUrlFromEnv(): string | undefined { /** * Checks if a token is an OAuth token by looking for sk-ant-oat prefix. - * @param apiKey - The API key to check - * @returns True if the token is an OAuth token */ export function isOAuthToken(apiKey: string): boolean { return apiKey.includes("sk-ant-oat"); } /** - * Converts a generic AuthCredential to AnthropicOAuthCredential if it's a valid OAuth entry. - * @param credential - The credential to convert - * @returns The converted OAuth credential, or null if not a valid OAuth type + * Build an {@link AnthropicAuthConfig} from an already-resolved API key. + * + * `apiKey` is whatever the caller chose for `Authorization`/`x-api-key` — + * usually `authStorage.getApiKey("anthropic")`. `baseUrl` overrides the + * env-derived base; pass `undefined` to fall back to FOUNDRY/ANTHROPIC env + * resolution and finally `DEFAULT_BASE_URL`. + * + * `isOAuth` is derived from the token prefix so the helper stays pure: callers + * never have to thread the OAuth flag through their own resolution logic. */ -function toAnthropicOAuthCredential(credential: AuthCredential): AnthropicOAuthCredential | null { - if (credential.type !== "oauth") return null; - if (typeof credential.access !== "string" || typeof credential.expires !== "number") return null; +export function buildAnthropicAuthConfig(apiKey: string, baseUrl?: string): AnthropicAuthConfig { return { - type: "oauth", - access: credential.access, - refresh: credential.refresh, - expires: credential.expires, + apiKey, + baseUrl: normalizeBaseUrl(baseUrl) ?? resolveAnthropicBaseUrlFromEnv() ?? DEFAULT_BASE_URL, + isOAuth: isOAuthToken(apiKey), }; } -/** - * Reads Anthropic OAuth credentials from an AuthCredentialStore. - * @param store - Credential store to read from (creates AuthCredentialStore if not provided) - * @returns Array of valid Anthropic OAuth credentials - */ -async function readAnthropicOAuthCredentials(store?: AuthCredentialStore): Promise { - const ownsStore = !store; - const effectiveStore = store ?? (await SqliteAuthCredentialStore.open(getAgentDbPath())); - try { - const records = effectiveStore.listAuthCredentials("anthropic"); - const credentials: AnthropicOAuthCredential[] = []; - for (const record of records) { - const mapped = toAnthropicOAuthCredential(record.credential); - if (mapped) { - credentials.push(mapped); - } - } - - return credentials; - } finally { - if (ownsStore) { - effectiveStore.close(); - } - } -} - -/** - * Finds Anthropic auth config using priority: - * 1. ANTHROPIC_SEARCH_API_KEY / ANTHROPIC_SEARCH_BASE_URL - * 2. ANTHROPIC_FOUNDRY_API_KEY override when Foundry mode is enabled - * 3. OAuth in agent.db (with 5-minute expiry buffer) - * 4. API key in agent.db - * 5. ANTHROPIC_API_KEY / ANTHROPIC_BASE_URL fallback - * @param store - Optional credential store (creates one from default db path if not provided) - * @returns The first valid auth configuration found, or null if none available - */ -export async function findAnthropicAuth(store?: AuthCredentialStore): Promise { - // 1. Explicit search-specific env vars - const searchApiKey = $env.ANTHROPIC_SEARCH_API_KEY; - const searchBaseUrl = $env.ANTHROPIC_SEARCH_BASE_URL; - if (searchApiKey) { - return { - apiKey: searchApiKey, - baseUrl: searchBaseUrl ?? DEFAULT_BASE_URL, - isOAuth: isOAuthToken(searchApiKey), - }; - } - - // 2. Foundry explicit env override - const foundryApiKey = isFoundryEnabled() ? $env.ANTHROPIC_FOUNDRY_API_KEY?.trim() : undefined; - if (foundryApiKey) { - return { - apiKey: foundryApiKey, - baseUrl: resolveAnthropicBaseUrlFromEnv() ?? DEFAULT_BASE_URL, - isOAuth: isOAuthToken(foundryApiKey), - }; - } - - // Tiers 3-4 use the credential store; manage lifecycle once - const ownsStore = !store; - const effectiveStore = store ?? (await SqliteAuthCredentialStore.open(getAgentDbPath())); - try { - // 3. OAuth credentials in agent.db (with 5-minute expiry buffer) - const expiryBuffer = 5 * 60 * 1000; // 5 minutes - const now = Date.now(); - const credentials = await readAnthropicOAuthCredentials(effectiveStore); - for (const credential of credentials) { - if (!credential.access) continue; - if (credential.expires > now + expiryBuffer) { - return { - apiKey: credential.access, - baseUrl: DEFAULT_BASE_URL, - isOAuth: true, - }; - } - } - - // 4. API key credentials in agent.db - const apiKeyRecord = effectiveStore - .listAuthCredentials("anthropic") - .find(record => record.credential.type === "api_key"); - if (apiKeyRecord && apiKeyRecord.credential.type === "api_key") { - return { - apiKey: apiKeyRecord.credential.key, - baseUrl: resolveAnthropicBaseUrlFromEnv() ?? DEFAULT_BASE_URL, - isOAuth: isOAuthToken(apiKeyRecord.credential.key), - }; - } - } finally { - if (ownsStore) { - effectiveStore.close(); - } - } - - // 5. Generic ANTHROPIC_API_KEY fallback - const apiKey = getEnvApiKey("anthropic"); - const baseUrl = resolveAnthropicBaseUrlFromEnv(); - if (apiKey) { - return { - apiKey, - baseUrl: baseUrl ?? DEFAULT_BASE_URL, - isOAuth: isOAuthToken(apiKey), - }; - } - - return null; -} - /** * Builds HTTP headers for Anthropic API requests (search variant). - * @param auth - The authentication configuration - * @returns Headers object ready for use in fetch requests */ export function buildAnthropicSearchHeaders(auth: AnthropicAuthConfig): Record { return buildProviderAnthropicHeaders({ @@ -198,8 +79,6 @@ export function buildAnthropicSearchHeaders(auth: AnthropicAuthConfig): Record= creds.expires) { - try { - creds = await refreshOAuthToken(provider, creds); - } catch (refreshError) { - if (provider === "perplexity") { - const jwtExpiry = getPerplexityJwtExpiryMs(creds.access); - if (jwtExpiry && Date.now() < jwtExpiry) { - const fallbackCredentials = { ...creds, expires: jwtExpiry }; - return { newCredentials: fallbackCredentials, apiKey: fallbackCredentials.access }; - } + if (provider === "perplexity") { + const jwtExpiry = getPerplexityJwtExpiryMs(creds.access); + if (jwtExpiry && Date.now() < jwtExpiry) { + const fallbackCredentials = { ...creds, expires: jwtExpiry }; + return { newCredentials: fallbackCredentials, apiKey: fallbackCredentials.access }; } - const reason = refreshError instanceof Error ? refreshError.message : String(refreshError); - throw new Error(`Failed to refresh OAuth token for ${provider}: ${reason}`); } + throw new Error( + `OAuth credential for ${provider} is expired and must be refreshed via AuthStorage before getOAuthApiKey is called`, + ); } // For providers that need request-time credential metadata, return JSON. const needsStructuredApiKey = diff --git a/packages/ai/test/anthropic-oauth.test.ts b/packages/ai/test/anthropic-oauth.test.ts index 66345cb4d..0916eb13f 100644 --- a/packages/ai/test/anthropic-oauth.test.ts +++ b/packages/ai/test/anthropic-oauth.test.ts @@ -1,9 +1,5 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; -import * as fs from "node:fs"; -import * as os from "node:os"; -import * as path from "node:path"; -import { SqliteAuthCredentialStore } from "../src/auth-storage"; -import { buildAnthropicUrl, findAnthropicAuth } from "../src/utils/anthropic-auth"; +import { buildAnthropicAuthConfig, buildAnthropicUrl } from "../src/utils/anthropic-auth"; import { AnthropicOAuthFlow, refreshAnthropicToken } from "../src/utils/oauth/anthropic"; import { withEnv } from "./helpers"; @@ -194,95 +190,72 @@ describe("anthropic oauth alignment", () => { }); }); -describe("anthropic auth resolution", () => { - it("prefers explicit Foundry env key over stored OAuth and normalizes Foundry base URL", async () => { - const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); - fs.mkdirSync(tmpDir, { recursive: true }); - const dbPath = path.join(tmpDir, "agent.db"); - const store = await SqliteAuthCredentialStore.open(dbPath); - try { - store.replaceAuthCredentialsForProvider("anthropic", [ - { type: "oauth", access: "sk-ant-oat-db", refresh: "refresh", expires: Date.now() + 20 * 60 * 1000 }, - ]); - await withEnv( - { - CLAUDE_CODE_USE_FOUNDRY: "true", - ANTHROPIC_FOUNDRY_API_KEY: "foundry-explicit-key", - FOUNDRY_BASE_URL: "https://foundry.example.com/anthropic/", - ANTHROPIC_API_KEY: undefined, - ANTHROPIC_OAUTH_TOKEN: undefined, - }, - async () => { - const auth = await findAnthropicAuth(store); - expect(auth).not.toBeNull(); - expect(auth?.apiKey).toBe("foundry-explicit-key"); - expect(auth?.isOAuth).toBe(false); - expect(auth?.baseUrl).toBe("https://foundry.example.com/anthropic"); - expect(buildAnthropicUrl(auth!)).toBe("https://foundry.example.com/anthropic/v1/messages?beta=true"); - }, - ); - } finally { - store.close(); - fs.rmSync(tmpDir, { recursive: true, force: true }); - } +describe("buildAnthropicAuthConfig", () => { + it("classifies sk-ant-oat tokens as OAuth", () => { + const config = buildAnthropicAuthConfig("sk-ant-oat-foobar"); + expect(config.isOAuth).toBe(true); + expect(config.apiKey).toBe("sk-ant-oat-foobar"); }); - it("keeps non-Foundry OAuth precedence unchanged", async () => { - const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); - fs.mkdirSync(tmpDir, { recursive: true }); - const dbPath = path.join(tmpDir, "agent.db"); - const store = await SqliteAuthCredentialStore.open(dbPath); - try { - store.replaceAuthCredentialsForProvider("anthropic", [ - { type: "oauth", access: "sk-ant-oat-db", refresh: "refresh", expires: Date.now() + 20 * 60 * 1000 }, - ]); - await withEnv( - { - CLAUDE_CODE_USE_FOUNDRY: undefined, - ANTHROPIC_FOUNDRY_API_KEY: "foundry-explicit-key", - ANTHROPIC_API_KEY: "sk-ant-api-env", - ANTHROPIC_OAUTH_TOKEN: undefined, - }, - async () => { - const auth = await findAnthropicAuth(store); - expect(auth).not.toBeNull(); - expect(auth?.apiKey).toBe("sk-ant-oat-db"); - expect(auth?.isOAuth).toBe(true); - expect(auth?.baseUrl).toBe("https://api.anthropic.com"); - }, - ); - } finally { - store.close(); - fs.rmSync(tmpDir, { recursive: true, force: true }); - } + it("treats sk-ant-api tokens as non-OAuth", () => { + const config = buildAnthropicAuthConfig("sk-ant-api-foobar"); + expect(config.isOAuth).toBe(false); }); - it("prefers stored API key over generic env fallback", async () => { - const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); - fs.mkdirSync(tmpDir, { recursive: true }); - const dbPath = path.join(tmpDir, "agent.db"); - const store = await SqliteAuthCredentialStore.open(dbPath); - try { - store.replaceAuthCredentialsForProvider("anthropic", [{ type: "api_key", key: "sk-ant-api-db" }]); - await withEnv( - { - CLAUDE_CODE_USE_FOUNDRY: undefined, - ANTHROPIC_FOUNDRY_API_KEY: undefined, - ANTHROPIC_API_KEY: "sk-ant-api-env", - ANTHROPIC_BASE_URL: "https://anthropic.example.com/", - ANTHROPIC_OAUTH_TOKEN: undefined, - }, - async () => { - const auth = await findAnthropicAuth(store); - expect(auth).not.toBeNull(); - expect(auth?.apiKey).toBe("sk-ant-api-db"); - expect(auth?.isOAuth).toBe(false); - expect(auth?.baseUrl).toBe("https://anthropic.example.com"); - }, - ); - } finally { - store.close(); - fs.rmSync(tmpDir, { recursive: true, force: true }); - } + it("normalizes the explicit baseUrl override (trailing slash, env precedence)", async () => { + await withEnv( + { + CLAUDE_CODE_USE_FOUNDRY: "true", + FOUNDRY_BASE_URL: "https://foundry.example.com/anthropic/", + ANTHROPIC_BASE_URL: undefined, + }, + async () => { + const explicit = buildAnthropicAuthConfig("sk-ant-api-key", "https://override.example.com/"); + expect(explicit.baseUrl).toBe("https://override.example.com"); + expect(buildAnthropicUrl(explicit)).toBe("https://override.example.com/v1/messages?beta=true"); + }, + ); + }); + + it("falls back to FOUNDRY_BASE_URL when Foundry mode is enabled and no explicit override is given", async () => { + await withEnv( + { + CLAUDE_CODE_USE_FOUNDRY: "true", + FOUNDRY_BASE_URL: "https://foundry.example.com/anthropic/", + ANTHROPIC_BASE_URL: undefined, + }, + async () => { + const config = buildAnthropicAuthConfig("sk-ant-api-key"); + expect(config.baseUrl).toBe("https://foundry.example.com/anthropic"); + }, + ); + }); + + it("falls back to ANTHROPIC_BASE_URL when Foundry mode is disabled", async () => { + await withEnv( + { + CLAUDE_CODE_USE_FOUNDRY: undefined, + FOUNDRY_BASE_URL: undefined, + ANTHROPIC_BASE_URL: "https://anthropic.example.com/", + }, + async () => { + const config = buildAnthropicAuthConfig("sk-ant-api-key"); + expect(config.baseUrl).toBe("https://anthropic.example.com"); + }, + ); + }); + + it("uses the default Anthropic base URL when no env or override is set", async () => { + await withEnv( + { + CLAUDE_CODE_USE_FOUNDRY: undefined, + FOUNDRY_BASE_URL: undefined, + ANTHROPIC_BASE_URL: undefined, + }, + async () => { + const config = buildAnthropicAuthConfig("sk-ant-api-key"); + expect(config.baseUrl).toBe("https://api.anthropic.com"); + }, + ); }); }); diff --git a/packages/ai/test/auth-storage-broker-no-sentinel.test.ts b/packages/ai/test/auth-storage-broker-no-sentinel.test.ts new file mode 100644 index 000000000..5f6507eed --- /dev/null +++ b/packages/ai/test/auth-storage-broker-no-sentinel.test.ts @@ -0,0 +1,104 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { + type AuthCredentialStore, + AuthStorage, + REMOTE_REFRESH_SENTINEL, + SqliteAuthCredentialStore, +} from "../src/auth-storage"; +import * as oauthUtils from "../src/utils/oauth"; + +describe("AuthStorage broker sentinel refresh", () => { + let tempDir = ""; + let store: AuthCredentialStore | undefined; + let authStorage: AuthStorage | undefined; + let brokerRefreshCalls = 0; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-broker-no-sentinel-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + brokerRefreshCalls = 0; + authStorage = new AuthStorage(store, { + refreshOAuthCredential: async (_provider, credentialId, credential) => { + brokerRefreshCalls += 1; + expect(credentialId).toBe(1); + expect(credential.refresh).toBe(REMOTE_REFRESH_SENTINEL); + return { + access: "broker-access-rotated", + refresh: REMOTE_REFRESH_SENTINEL, + expires: Date.now() + 60 * 60_000, + accountId: "broker-account", + email: "broker@example.com", + projectId: "broker-project", + }; + }, + }); + }); + + afterEach(async () => { + vi.restoreAllMocks(); + store?.close(); + store = undefined; + authStorage = undefined; + if (tempDir) { + await fs.rm(tempDir, { recursive: true, force: true }); + tempDir = ""; + } + }); + + test("getOAuthAccess refreshes expired broker credentials through the store hook only", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + await authStorage.set("anthropic", [ + { + type: "oauth", + access: "broker-access-stale", + refresh: REMOTE_REFRESH_SENTINEL, + expires: Date.now() - 60_000, + accountId: "broker-account-old", + }, + ]); + + const providerRefresh = vi.spyOn(oauthUtils, "refreshOAuthToken").mockImplementation(async () => { + throw new Error("provider-direct refresh must not be called"); + }); + + const access = await authStorage.getOAuthAccess("anthropic", "broker-session"); + + expect(access).toEqual({ + accessToken: "broker-access-rotated", + accountId: "broker-account", + email: "broker@example.com", + projectId: "broker-project", + enterpriseUrl: undefined, + }); + expect(brokerRefreshCalls).toBe(1); + expect(providerRefresh).not.toHaveBeenCalled(); + const persisted = store.listAuthCredentials("anthropic"); + expect(persisted).toHaveLength(1); + expect(persisted[0]?.credential.type).toBe("oauth"); + if (persisted[0]?.credential.type === "oauth") { + expect(persisted[0].credential.access).toBe("broker-access-rotated"); + expect(persisted[0].credential.refresh).toBe(REMOTE_REFRESH_SENTINEL); + } + }); + + test("getOAuthApiKey refuses expired broker sentinels instead of provider-direct refresh", async () => { + const providerRefresh = vi.spyOn(oauthUtils, "refreshOAuthToken").mockImplementation(async () => { + throw new Error("provider-direct refresh must not be called"); + }); + + await expect( + oauthUtils.getOAuthApiKey("anthropic", { + anthropic: { + access: "broker-access-stale", + refresh: REMOTE_REFRESH_SENTINEL, + expires: Date.now() - 60_000, + }, + }), + ).rejects.toThrow("must be refreshed via AuthStorage"); + expect(providerRefresh).not.toHaveBeenCalled(); + }); +}); diff --git a/packages/ai/test/auth-storage-refresh-skew.test.ts b/packages/ai/test/auth-storage-refresh-skew.test.ts new file mode 100644 index 000000000..d3ed0517a --- /dev/null +++ b/packages/ai/test/auth-storage-refresh-skew.test.ts @@ -0,0 +1,128 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { type AuthCredentialStore, AuthStorage, SqliteAuthCredentialStore } from "../src/auth-storage"; +import { registerOAuthProvider, unregisterOAuthProviders } from "../src/utils/oauth"; + +describe("AuthStorage OAuth refresh skew", () => { + let tempDir = ""; + let store: AuthCredentialStore | undefined; + let authStorage: AuthStorage | undefined; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-refresh-skew-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + authStorage = new AuthStorage(store); + }); + + afterEach(async () => { + unregisterOAuthProviders("auth-storage-refresh-skew-test"); + store?.close(); + store = undefined; + authStorage = undefined; + if (tempDir) { + await fs.rm(tempDir, { recursive: true, force: true }); + tempDir = ""; + } + }); + + test("refreshes before strict expiry when the credential is inside the 60s skew", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + let refreshCalls = 0; + const refreshedExpires = Date.now() + 60 * 60_000; + registerOAuthProvider({ + id: "unit-oauth-skew", + name: "Unit OAuth Skew", + sourceId: "auth-storage-refresh-skew-test", + async login() { + return { access: "unused", refresh: "unused", expires: refreshedExpires }; + }, + async refreshToken(credentials) { + refreshCalls += 1; + return { + ...credentials, + access: "access-after-skew-refresh", + refresh: "refresh-after-skew-refresh", + expires: refreshedExpires, + }; + }, + getApiKey(credentials) { + return credentials.access; + }, + }); + + await authStorage.set("unit-oauth-skew", [ + { + type: "oauth", + access: "access-before-skew-refresh", + refresh: "refresh-before-skew-refresh", + expires: Date.now() + 30_000, + }, + ]); + + const apiKey = await authStorage.getApiKey("unit-oauth-skew", "skew-session"); + + expect(apiKey).toBe("access-after-skew-refresh"); + expect(refreshCalls).toBe(1); + const stored = store.listAuthCredentials("unit-oauth-skew"); + expect(stored).toHaveLength(1); + expect(stored[0]?.credential.type).toBe("oauth"); + if (stored[0]?.credential.type === "oauth") { + expect(stored[0].credential.access).toBe("access-after-skew-refresh"); + expect(stored[0].credential.refresh).toBe("refresh-after-skew-refresh"); + } + }); + + test("coalesces concurrent skew refreshes for the same credential", async () => { + if (!authStorage) throw new Error("test setup failed"); + + const refreshedExpires = Date.now() + 60 * 60_000; + const refreshStarted = Promise.withResolvers(); + const allowRefresh = Promise.withResolvers(); + let refreshCalls = 0; + + registerOAuthProvider({ + id: "unit-oauth-skew-mutex", + name: "Unit OAuth Skew Mutex", + sourceId: "auth-storage-refresh-skew-test", + async login() { + return { access: "unused", refresh: "unused", expires: refreshedExpires }; + }, + async refreshToken(credentials) { + refreshCalls += 1; + refreshStarted.resolve(); + await allowRefresh.promise; + return { + ...credentials, + access: "access-after-shared-skew-refresh", + refresh: "refresh-after-shared-skew-refresh", + expires: refreshedExpires, + }; + }, + getApiKey(credentials) { + return credentials.access; + }, + }); + + await authStorage.set("unit-oauth-skew-mutex", [ + { + type: "oauth", + access: "access-before-shared-skew-refresh", + refresh: "refresh-before-shared-skew-refresh", + expires: Date.now() + 30_000, + }, + ]); + + const first = authStorage.getApiKey("unit-oauth-skew-mutex", "same-session"); + const second = authStorage.getApiKey("unit-oauth-skew-mutex", "same-session"); + + await refreshStarted.promise; + allowRefresh.resolve(); + + await expect(first).resolves.toBe("access-after-shared-skew-refresh"); + await expect(second).resolves.toBe("access-after-shared-skew-refresh"); + expect(refreshCalls).toBe(1); + }); +}); diff --git a/packages/ai/test/google-gemini-cli-alignment.test.ts b/packages/ai/test/google-gemini-cli-alignment.test.ts index fde0a30f8..bfadc6171 100644 --- a/packages/ai/test/google-gemini-cli-alignment.test.ts +++ b/packages/ai/test/google-gemini-cli-alignment.test.ts @@ -1,5 +1,6 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { hookFetch } from "@oh-my-pi/pi-utils"; +import * as geminiCliProvider from "../src/providers/google-gemini-cli"; import { ANTIGRAVITY_SYSTEM_INSTRUCTION, buildRequest, @@ -115,6 +116,11 @@ describe("Google Gemini CLI alignment", () => { expect(shouldRefreshGeminiCliCredentials(preBufferedExpiry, true, issuedAt + 54 * 60 * 1000)).toBe(true); expect(shouldRefreshGeminiCliCredentials(preBufferedExpiry, false, issuedAt + 54 * 60 * 1000)).toBe(true); }); + + it("does not export provider-direct refresh helper", () => { + expect(shouldRefreshGeminiCliCredentials).toBe(geminiCliProvider.shouldRefreshGeminiCliCredentials); + expect(Object.hasOwn(geminiCliProvider, "refreshGeminiCliCredentialsIfNeeded")).toBe(false); + }); it("omits antigravity-only metadata in non-antigravity request payloads", () => { const model = createModel("google-gemini-cli"); const payload = buildRequest(model, createContext(), "proj-123", {}, false) as { diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 3abd90bc9..c4921806e 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Added - Added `codex` and `gemini` to the web search provider settings so users can configure OpenAI and Gemini web search directly from provider selection @@ -8,11 +9,14 @@ ### Changed +- Changed web search provider credential lookup to use the shared `AuthStorage` pipeline (`getApiKey`/`getOAuthAccess`) for API-key and OAuth auth instead of direct `AgentStorage` access - Changed the `codex` web search provider display label from `Codex` to `OpenAI` - Updated `anthropic` and `openai`/`gemini` web search option descriptions to reflect their native `web_search`/OAuth requirements ### Fixed +- Fixed web search OAuth-backed providers (including Codex and Gemini) to use broker-managed token retrieval and account metadata, avoiding direct token-store refresh behavior that could cause search authentication failures +- Updated Tavily missing-credential feedback to prompt users to configure an API-key provider setting instead of referencing `agent.db` directly - Refreshed expired OpenAI Codex OAuth tokens during `web_search` execution and persisted the updated credentials so searches continue working after token expiry ## [15.3.2] - 2026-05-25 diff --git a/packages/coding-agent/src/web/kagi.ts b/packages/coding-agent/src/web/kagi.ts index 1c350e9eb..06ac5df5c 100644 --- a/packages/coding-agent/src/web/kagi.ts +++ b/packages/coding-agent/src/web/kagi.ts @@ -1,6 +1,5 @@ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../session/agent-storage"; -import { findCredential, withHardTimeout } from "./search/providers/utils"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import { withHardTimeout } from "./search/providers/utils"; const KAGI_SEARCH_URL = "https://kagi.com/api/v0/search"; @@ -98,6 +97,7 @@ function parseKagiErrorResponse(statusCode: number, responseText: string): KagiA export interface KagiSearchOptions { limit?: number; + sessionId?: string; signal?: AbortSignal; } @@ -114,8 +114,12 @@ export interface KagiSearchResult { relatedQuestions: string[]; } -export function findKagiApiKey(storage: AgentStorage): string | null { - return findCredential(storage, getEnvApiKey("kagi"), "kagi"); +export async function findKagiApiKey( + authStorage: AuthStorage, + sessionId?: string, + signal?: AbortSignal, +): Promise { + return (await authStorage.getApiKey("kagi", sessionId, { signal })) ?? null; } function getAuthHeaders(apiKey: string): Record { @@ -128,9 +132,9 @@ function getAuthHeaders(apiKey: string): Record { export async function searchWithKagi( query: string, options: KagiSearchOptions = {}, - storage: AgentStorage, + authStorage: AuthStorage, ): Promise { - const apiKey = findKagiApiKey(storage); + 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'."); } diff --git a/packages/coding-agent/src/web/search/index.ts b/packages/coding-agent/src/web/search/index.ts index 7428c3570..9c6c3f073 100644 --- a/packages/coding-agent/src/web/search/index.ts +++ b/packages/coding-agent/src/web/search/index.ts @@ -3,16 +3,16 @@ * * Single tool supporting Anthropic, Perplexity, Exa, Brave, Jina, Kimi, Gemini, Codex, Tavily, Kagi, Z.AI, SearXNG, and Synthetic * providers with provider-specific parameters exposed conditionally. - * */ import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; -import { getAgentDbPath, prompt } from "@oh-my-pi/pi-utils"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import { prompt } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; import type { CustomTool, CustomToolContext, RenderResultOptions } from "../../extensibility/custom-tools/types"; import type { Theme } from "../../modes/theme/theme"; import webSearchSystemPrompt from "../../prompts/system/web-search.md" with { type: "text" }; import webSearchDescription from "../../prompts/tools/web-search.md" with { type: "text" }; -import { AgentStorage } from "../../session/agent-storage"; +import { discoverAuthStorage } from "../../sdk"; import type { ToolSession } from "../../tools"; import { formatAge } from "../../tools/render-utils"; import { throwIfAborted } from "../../tools/tool-errors"; @@ -115,19 +115,25 @@ function formatForLLM(response: SearchResponse): string { return parts.join("\n"); } +interface ExecuteSearchOptions { + authStorage: AuthStorage; + sessionId?: string; + signal?: AbortSignal; +} + /** Execute web search */ async function executeSearch( _toolCallId: string, params: SearchQueryParams, - signal?: AbortSignal, + options: ExecuteSearchOptions, ): Promise<{ content: Array<{ type: "text"; text: string }>; details: SearchRenderDetails }> { - const storage = await AgentStorage.open(getAgentDbPath()); + const { authStorage, sessionId, signal } = options; const providers = params.provider && params.provider !== "auto" ? await getSearchProvider(params.provider).then(async provider => - (await provider.isAvailable(storage)) ? [provider] : resolveProviderChain(storage, "auto"), + (await provider.isAvailable(authStorage)) ? [provider] : resolveProviderChain(authStorage, "auto"), ) - : await resolveProviderChain(storage); + : await resolveProviderChain(authStorage); if (providers.length === 0) { const message = "No web search provider configured."; return { @@ -141,19 +147,18 @@ async function executeSearch( for (const provider of providers) { lastProvider = provider; try { - const response = await provider.search( - { - query: params.query.replace(/202\d/g, String(new Date().getFullYear())), // LUL - limit: params.limit, - recency: params.recency, - systemPrompt: webSearchSystemPrompt, - maxOutputTokens: params.max_tokens, - numSearchResults: params.num_search_results, - temperature: params.temperature, - signal, - }, - storage, - ); + const response = await provider.search({ + query: params.query.replace(/202\d/g, String(new Date().getFullYear())), // LUL + limit: params.limit, + recency: params.recency, + systemPrompt: webSearchSystemPrompt, + maxOutputTokens: params.max_tokens, + numSearchResults: params.num_search_results, + temperature: params.temperature, + signal, + authStorage, + sessionId, + }); const text = formatForLLM(response); @@ -195,18 +200,27 @@ async function executeSearch( /** * Execute a web search query for CLI/testing workflows. + * + * `authStorage` may be omitted; in that case we discover one via the standard + * factory (`discoverAuthStorage`), which honours `OMP_AUTH_BROKER_URL` and + * otherwise opens the local SQLite credential store. */ export async function runSearchQuery( params: SearchQueryParams, + options: { authStorage?: AuthStorage; sessionId?: string; signal?: AbortSignal } = {}, ): Promise<{ content: Array<{ type: "text"; text: string }>; details: SearchRenderDetails }> { - return executeSearch("cli-web-search", params); + const authStorage = options.authStorage ?? (await discoverAuthStorage()); + return executeSearch("cli-web-search", params, { + authStorage, + sessionId: options.sessionId, + signal: options.signal, + }); } /** * Web search tool implementation. * * Supports Anthropic, Perplexity, Exa, Brave, Jina, Kimi, Gemini, Codex, Z.AI, SearXNG, and Synthetic providers with automatic fallback. - * Session is accepted for interface consistency but not used. */ export class WebSearchTool implements AgentTool { readonly name = "web_search"; @@ -217,7 +231,10 @@ export class WebSearchTool implements AgentTool, _context?: AgentToolContext, ): Promise> { - return executeSearch(_toolCallId, params, signal); + const authStorage = this.#session.authStorage ?? (await discoverAuthStorage()); + const sessionId = this.#session.getSessionId?.() ?? undefined; + return executeSearch(_toolCallId, params, { authStorage, sessionId, signal }); } } @@ -243,10 +262,11 @@ export const webSearchCustomTool: CustomTool { const providers: SearchProvider[] = []; if (preferredProvider !== "auto") { const provider = await getSearchProvider(preferredProvider); - if (await provider.isAvailable(storage)) { + if (await provider.isAvailable(authStorage)) { providers.push(provider); } } @@ -164,7 +164,7 @@ export async function resolveProviderChain( for (const id of SEARCH_PROVIDER_ORDER) { if (id === preferredProvider) continue; const provider = await getSearchProvider(id); - if (await provider.isAvailable(storage)) { + if (await provider.isAvailable(authStorage)) { providers.push(provider); } } diff --git a/packages/coding-agent/src/web/search/providers/anthropic.ts b/packages/coding-agent/src/web/search/providers/anthropic.ts index 905262747..8e245c7bb 100644 --- a/packages/coding-agent/src/web/search/providers/anthropic.ts +++ b/packages/coding-agent/src/web/search/providers/anthropic.ts @@ -7,14 +7,14 @@ import { type AnthropicAuthConfig, type AnthropicSystemBlock, + type AuthStorage, + buildAnthropicAuthConfig, buildAnthropicSearchHeaders, buildAnthropicSystemBlocks, buildAnthropicUrl, - findAnthropicAuth, stripClaudeToolPrefix, } from "@oh-my-pi/pi-ai"; import { $env } from "@oh-my-pi/pi-utils"; -import type { AgentStorage } from "../../../session/agent-storage"; import type { AnthropicApiResponse, AnthropicCitation, @@ -35,9 +35,7 @@ export interface AnthropicSearchParams { query: string; system_prompt?: string; num_results?: number; - /** Maximum output tokens. Defaults to 4096. */ max_tokens?: number; - /** Sampling temperature (0–1). Lower = more focused/factual. */ temperature?: number; signal?: AbortSignal; } @@ -243,30 +241,47 @@ function parseResponse(response: AnthropicApiResponse): SearchResponse { * @returns Search response with synthesized answer, sources, and citations * @throws {Error} If no Anthropic credentials are configured */ -export async function searchAnthropic(params: AnthropicSearchParams, storage: AgentStorage): Promise { - const auth = await findAnthropicAuth(storage.authStore); +export async function searchAnthropic( + params: SearchParams | AnthropicSearchParams, + _legacyStorage?: unknown, +): Promise { + const searchApiKey = $env.ANTHROPIC_SEARCH_API_KEY; + const searchBaseUrl = $env.ANTHROPIC_SEARCH_BASE_URL; + let auth: AnthropicAuthConfig | undefined; + + if (searchApiKey) { + auth = buildAnthropicAuthConfig(searchApiKey, searchBaseUrl); + } else if ("authStorage" in params) { + const apiKey = await params.authStorage.getApiKey("anthropic", params.sessionId, { + signal: params.signal, + }); + if (apiKey) auth = buildAnthropicAuthConfig(apiKey); + } + if (!auth) { throw new Error( - "No Anthropic credentials found. Set ANTHROPIC_API_KEY or configure OAuth in ~/.omp/agent/agent.db", + "No Anthropic credentials found. Set ANTHROPIC_SEARCH_API_KEY or ANTHROPIC_API_KEY, or configure Anthropic OAuth.", ); } const model = getModel(); + const systemPrompt = "authStorage" in params ? params.systemPrompt : params.system_prompt; + const maxTokens = "authStorage" in params ? params.maxOutputTokens : params.max_tokens; const response = await callSearch( auth, model, params.query, - params.system_prompt, - params.max_tokens, + systemPrompt, + maxTokens, params.temperature, params.signal, ); const result = parseResponse(response); - // Apply num_results limit if specified - if (params.num_results && result.sources.length > params.num_results) { - result.sources = result.sources.slice(0, params.num_results); + const numResults = "authStorage" in params ? (params.numSearchResults ?? params.limit) : params.num_results; + if (numResults && result.sources.length > numResults) { + result.sources = result.sources.slice(0, numResults); } return result; @@ -277,21 +292,11 @@ export class AnthropicProvider extends SearchProvider { readonly id = "anthropic"; readonly label = "Anthropic"; - isAvailable(storage: AgentStorage) { - return findAnthropicAuth(storage.authStore).then(Boolean); + isAvailable(authStorage: AuthStorage): Promise | boolean { + return Boolean($env.ANTHROPIC_SEARCH_API_KEY) || authStorage.hasAuth("anthropic"); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchAnthropic( - { - query: params.query, - system_prompt: params.systemPrompt, - num_results: params.numSearchResults ?? params.limit, - max_tokens: params.maxOutputTokens, - temperature: params.temperature, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchAnthropic(params); } } diff --git a/packages/coding-agent/src/web/search/providers/base.ts b/packages/coding-agent/src/web/search/providers/base.ts index 6653614b8..c3821e71a 100644 --- a/packages/coding-agent/src/web/search/providers/base.ts +++ b/packages/coding-agent/src/web/search/providers/base.ts @@ -1,7 +1,16 @@ -import type { AgentStorage } from "../../../session/agent-storage"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import type { SearchProviderId, SearchResponse } from "../types"; -/** Shared web search parameters passed to providers. */ +/** + * Shared web search parameters passed to providers. + * + * `authStorage` is the **only** credential source providers may consult. + * Opening a sibling SQLite handle or calling provider-direct refresh helpers + * (e.g. `refreshOpenAICodexToken`, `refreshGoogleCloudToken`) is prohibited: + * it races the broker's per-credential refresh and POSTs the broker sentinel + * (`REMOTE_REFRESH_SENTINEL`) to the upstream token endpoint, which classifies + * as `invalid_grant` and disables the row. + */ export interface SearchParams { query: string; limit?: number; @@ -27,6 +36,20 @@ export interface SearchParams { googleSearch?: Record; codeExecution?: Record; urlContext?: Record; + /** + * The single source of truth for credentials. Providers MUST consult this + * handle exclusively (`getApiKey` for bearer-style auth, `getOAuthAccess` + * when identity metadata is required). Do not open `AgentStorage` or any + * `AuthCredentialStore` directly — that bypasses the broker pipeline and + * the per-credential single-flight refresh. + */ + authStorage: AuthStorage; + /** + * Optional session id used as the round-robin / sticky key when selecting + * among multiple credentials for the same provider. Pass through from the + * caller's agent session when available; otherwise omit. + */ + sessionId?: string; } /** Base class for web search providers. */ @@ -36,16 +59,13 @@ export abstract class SearchProvider { /** * Indicates whether this provider has the credentials/config it needs to - * service a request right now. Implementations may consult the shared - * {@link AgentStorage} handle for stored OAuth credentials and avoid - * opening their own connection. + * service a request right now. Implementations consult the passed + * {@link AuthStorage} — never a sibling store. */ - abstract isAvailable(storage: AgentStorage): Promise | boolean; + abstract isAvailable(authStorage: AuthStorage): Promise | boolean; /** - * Execute a search. Implementations that read credentials from - * {@link AgentStorage} MUST use the passed handle instead of opening a - * second one. + * Execute a search. Credentials MUST be resolved through `params.authStorage`. */ - abstract search(params: SearchParams, storage: AgentStorage): Promise; + abstract search(params: SearchParams): Promise; } diff --git a/packages/coding-agent/src/web/search/providers/brave.ts b/packages/coding-agent/src/web/search/providers/brave.ts index b3a7e15cb..283228311 100644 --- a/packages/coding-agent/src/web/search/providers/brave.ts +++ b/packages/coding-agent/src/web/search/providers/brave.ts @@ -4,8 +4,7 @@ * Calls Brave's web search REST API and maps results into the unified * SearchResponse shape used by the web search tool. */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; @@ -135,11 +134,11 @@ export class BraveProvider extends SearchProvider { readonly id = "brave"; readonly label = "Brave"; - isAvailable(_storage: AgentStorage): boolean { + isAvailable(_authStorage: AuthStorage): boolean { return !!findApiKey(); } - search(params: SearchParams, _storage: AgentStorage): Promise { + search(params: SearchParams): Promise { return searchBrave({ query: params.query, num_results: params.numSearchResults ?? params.limit, diff --git a/packages/coding-agent/src/web/search/providers/codex.ts b/packages/coding-agent/src/web/search/providers/codex.ts index c210cf6e7..485bbf7af 100644 --- a/packages/coding-agent/src/web/search/providers/codex.ts +++ b/packages/coding-agent/src/web/search/providers/codex.ts @@ -2,21 +2,15 @@ * OpenAI Codex Web Search Provider * * Uses Codex's built-in web_search tool via the Responses API. - * Requires OAuth credentials stored in agent.db for provider "openai-codex". - * Returns synthesized answers with web search sources. + * Auth is resolved through `AuthStorage.getOAuthAccess("openai-codex")` so the + * broker is the sole refresh authority — this module never opens a sibling + * SQLite store, never POSTs the broker sentinel to an OpenAI token endpoint. */ import * as os from "node:os"; -import { - type AuthCredential, - getBundledModels, - type OAuthCredential, - type OAuthCredentials, - REMOTE_REFRESH_SENTINEL, -} from "@oh-my-pi/pi-ai"; -import { decodeJwt, refreshOpenAICodexToken } from "@oh-my-pi/pi-ai/utils/oauth/openai-codex"; -import { $env, logger, readSseJson } from "@oh-my-pi/pi-utils"; +import { type AuthStorage, getBundledModels } from "@oh-my-pi/pi-ai"; +import { decodeJwt } from "@oh-my-pi/pi-ai/utils/oauth/openai-codex"; +import { $env, readSseJson } from "@oh-my-pi/pi-utils"; import packageJson from "../../../../package.json" with { type: "json" }; -import type { AgentStorage } from "../../../session/agent-storage"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import type { SearchParams } from "./base"; @@ -25,29 +19,48 @@ import { classifyProviderHttpError, withHardTimeout } from "./utils"; const CODEX_BASE_URL = "https://chatgpt.com/backend-api"; const CODEX_RESPONSES_PATH = "/codex/responses"; -const FALLBACK_MODEL = "gpt-5-codex-mini"; +const FALLBACK_MODEL = "gpt-5.4"; const DEFAULT_MODEL_PREFERENCES = [ - "gpt-5-codex-mini", "gpt-5.4", + "gpt-5-codex", + "gpt-5", "gpt-5.3-codex", "gpt-5.2-codex", "gpt-5.1-codex", - "gpt-5-codex", + "gpt-5-codex-mini", ]; const JWT_CLAIM_PATH = "https://api.openai.com/auth"; const DEFAULT_INSTRUCTIONS = "You are a helpful assistant with web search capabilities. Search the web to answer the user's question accurately and cite your sources."; -function getModel(): string { +function getConfiguredModel(): string | undefined { const configuredModel = $env.PI_CODEX_WEB_SEARCH_MODEL?.trim(); - if (configuredModel) return configuredModel; + return configuredModel ? configuredModel : undefined; +} +function getDefaultModelCandidates(): string[] { const bundledModels = getBundledModels("openai-codex"); const bundledIds = new Set(bundledModels.map(model => model.id)); - const preferred = DEFAULT_MODEL_PREFERENCES.find(modelId => bundledIds.has(modelId)); - if (preferred) return preferred; + const candidates = DEFAULT_MODEL_PREFERENCES.filter(modelId => bundledIds.has(modelId)); + + if (candidates.length > 0) { + return candidates; + } + const nonMini = bundledModels.find(model => !model.id.includes("mini") && !model.id.includes("spark")); - return nonMini?.id ?? bundledModels[0]?.id ?? FALLBACK_MODEL; + if (nonMini) { + return [nonMini.id]; + } + + return bundledModels[0]?.id ? [bundledModels[0].id] : [FALLBACK_MODEL]; +} + +function shouldRetryWithNextDefaultModel(error: unknown): boolean { + if (!(error instanceof SearchProviderError)) return false; + if (error.provider !== "codex" || error.status !== 400) return false; + return /model is not supported|requested model is not supported|not supported when using codex with a chatgpt account/i.test( + error.message, + ); } export interface CodexSearchParams { @@ -59,15 +72,6 @@ export interface CodexSearchParams { search_context_size?: "low" | "medium" | "high"; } -/** OAuth credential stored in agent.db */ -interface CodexOAuthCredential { - type: "oauth"; - access: string; - refresh?: string; - expires: number; - accountId?: string; -} - /** Codex API response structure */ interface CodexResponseItem { type: string; @@ -237,7 +241,7 @@ function extractTextSources(text: string): SearchSource[] { * @param accessToken - JWT access token * @returns Account ID string, or null if not found */ -function getAccountId(accessToken: string): string | null { +function getAccountIdFromJwt(accessToken: string): string | null { const payload = decodeJwt(accessToken); const auth = payload?.[JWT_CLAIM_PATH] as { chatgpt_account_id?: string } | undefined; const accountId = auth?.chatgpt_account_id; @@ -245,100 +249,25 @@ function getAccountId(accessToken: string): string | null { } /** - * Finds valid Codex OAuth credentials in the given storage handle. - * - * Walks each stored "openai-codex" OAuth row; for the first usable row, - * returns its access token + account ID. When the access token is expired (or - * near expiry) and a refresh token is available, refreshes it via the OpenAI - * token endpoint and persists the refreshed credential back to storage so the - * next call sees the updated row. - * - * Credentials with a broker-redacted refresh slot - * ({@link REMOTE_REFRESH_SENTINEL}) cannot be refreshed from this provider — - * the refresh token lives on the auth-broker — so they are skipped once - * expired. The main agent session refreshes such credentials through - * `AuthStorage.refreshOAuthCredential` and persists them to the same row, - * which means subsequent web_search calls will pick them up. - * - * @returns OAuth credential with access token and account ID, or null if none found + * Resolve a Codex bearer + accountId through {@link AuthStorage} — the single + * refresh authority. Returns `null` when no OAuth credential is configured, + * when the credential cannot be refreshed (broker error, revoked token, etc.), + * or when the access token carries no `chatgpt_account_id` claim. */ -async function findCodexAuth(storage: AgentStorage): Promise<{ accessToken: string; accountId: string } | null> { - const expiryBuffer = 5 * 60 * 1000; // 5 minutes - const now = Date.now(); - - let records: ReadonlyArray<{ id: number; credential: AuthCredential }>; - try { - records = storage.listAuthCredentials("openai-codex"); - } catch { - return null; - } - - for (const record of records) { - const credential = record.credential; - if (credential.type !== "oauth") continue; - - const oauthCred = credential as CodexOAuthCredential; - if (!oauthCred.access) continue; - - let accessToken = oauthCred.access; - let expires = oauthCred.expires; - let refreshToken = oauthCred.refresh; - - if (expires <= now + expiryBuffer) { - // Stored access token is expired or about to expire. Refresh - // inline so codex web_search stays available between agent - // sessions and after long-running turns. The main agent's - // AuthStorage refreshes the same row on its own schedule; both - // writers are safe because updateAuthCredential is a row-level - // UPDATE keyed by the stored record id. - if (!refreshToken || refreshToken === REMOTE_REFRESH_SENTINEL) continue; - - let refreshed: OAuthCredentials; - try { - refreshed = await refreshOpenAICodexToken(refreshToken); - } catch (error) { - logger.warn("codex web_search: token refresh failed; skipping credential", { - credentialId: record.id, - error: error instanceof Error ? error.message : String(error), - }); - continue; - } - - accessToken = refreshed.access; - expires = refreshed.expires; - refreshToken = refreshed.refresh || refreshToken; - - const updated: OAuthCredential = { - ...oauthCred, - access: accessToken, - refresh: refreshToken, - expires, - accountId: refreshed.accountId ?? oauthCred.accountId, - }; - try { - storage.updateAuthCredential(record.id, updated); - } catch (error) { - logger.warn("codex web_search: failed to persist refreshed credential", { - credentialId: record.id, - error: error instanceof Error ? error.message : String(error), - }); - } - } - - const accountId = oauthCred.accountId ?? getAccountId(accessToken); - if (!accountId) continue; - - return { accessToken, accountId }; - } - - return null; +async function findCodexAuth( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, +): Promise<{ accessToken: string; accountId: string } | null> { + const access = await authStorage.getOAuthAccess("openai-codex", sessionId, { signal }); + if (!access) return null; + const accountId = access.accountId ?? getAccountIdFromJwt(access.accessToken); + if (!accountId) return null; + return { accessToken: access.accessToken, accountId }; } /** * Builds HTTP headers for Codex API requests. - * @param accessToken - OAuth access token - * @param accountId - ChatGPT account ID - * @returns Headers object for fetch requests */ function buildCodexHeaders(accessToken: string, accountId: string): Record { return { @@ -354,17 +283,19 @@ function buildCodexHeaders(accessToken: string, accountId: string): Record = { model: requestedModel, @@ -520,29 +451,67 @@ async function callCodexSearch( /** * Executes a web search using OpenAI Codex's built-in web search tool. - * Requires OAuth credentials stored in agent.db for provider "openai-codex". - * @param params - Search parameters including query and optional settings - * @returns Search response with synthesized answer, sources, and usage - * @throws {Error} If no Codex OAuth credentials are configured + * + * Default-model behavior: + * - If `PI_CODEX_WEB_SEARCH_MODEL` is set, use it exactly once and surface any + * upstream error verbatim. + * - Otherwise prefer ChatGPT-account-safe bundled defaults (GPT-5.4, GPT-5 + * Codex, GPT-5, …) and retry the next candidate only when Codex returns the + * known 400 "model is not supported" family. This avoids selecting + * `gpt-5-codex-mini` first on ChatGPT accounts, which OpenAI rejects. */ -export async function searchCodex(params: CodexSearchParams, storage: AgentStorage): Promise { - const auth = await findCodexAuth(storage); +export async function searchCodex(params: SearchParams): Promise { + const auth = await findCodexAuth(params.authStorage, params.sessionId, params.signal); if (!auth) { throw new Error( "No Codex OAuth credentials found. Login with 'omp /login openai-codex' to enable Codex web search.", ); } - const result = await callCodexSearch(auth, params.query, { - systemPrompt: params.system_prompt, - searchContextSize: params.search_context_size ?? "high", - }); + const configuredModel = getConfiguredModel(); + const modelCandidates = configuredModel ? [configuredModel] : getDefaultModelCandidates(); + + let result: + | { + answer: string; + sources: SearchSource[]; + model: string; + requestId: string; + usage?: { inputTokens: number; outputTokens: number; totalTokens: number }; + } + | undefined; + let lastError: unknown; + + for (let index = 0; index < modelCandidates.length; index += 1) { + const modelId = modelCandidates[index]; + if (!modelId) continue; + + try { + result = await callCodexSearch(auth, params.query, { + signal: params.signal, + systemPrompt: params.systemPrompt, + searchContextSize: "high", + modelId, + }); + break; + } catch (error) { + lastError = error; + const isLastCandidate = index === modelCandidates.length - 1; + if (configuredModel || isLastCandidate || !shouldRetryWithNextDefaultModel(error)) { + throw error; + } + } + } + + if (!result) { + throw lastError ?? new Error("Codex search failed without returning a result"); + } let sources = result.sources; - // Apply num_results limit if specified - if (params.num_results && sources.length > params.num_results) { - sources = sources.slice(0, params.num_results); + const numResults = params.numSearchResults ?? params.limit; + if (numResults && sources.length > numResults) { + sources = sources.slice(0, numResults); } return { @@ -563,11 +532,13 @@ export async function searchCodex(params: CodexSearchParams, storage: AgentStora /** * Checks if Codex web search is available. - * @returns True if valid OAuth credentials exist for openai-codex */ -export async function hasCodexSearch(storage: AgentStorage): Promise { - const auth = await findCodexAuth(storage); - return auth !== null; +export async function hasCodexSearch(authStorage: AuthStorage): Promise { + // `isAvailable` runs before every request — keep the probe cheap. + // `hasOAuth(...)` is a synchronous in-memory check that returns true as soon + // as a Codex OAuth credential is loaded, without driving the refresh + // pipeline. The actual refresh happens lazily in `searchCodex`. + return authStorage.hasOAuth("openai-codex"); } /** Search provider for OpenAI Codex web search. */ @@ -575,19 +546,11 @@ export class CodexProvider extends SearchProvider { readonly id = "codex"; readonly label = "OpenAI"; - isAvailable(storage: AgentStorage): Promise { - return Promise.resolve(hasCodexSearch(storage)); + isAvailable(authStorage: AuthStorage): Promise | boolean { + return hasCodexSearch(authStorage); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchCodex( - { - signal: params.signal, - query: params.query, - system_prompt: params.systemPrompt, - num_results: params.numSearchResults ?? params.limit, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchCodex(params); } } diff --git a/packages/coding-agent/src/web/search/providers/exa.ts b/packages/coding-agent/src/web/search/providers/exa.ts index 23a1c3016..2397aa2d4 100644 --- a/packages/coding-agent/src/web/search/providers/exa.ts +++ b/packages/coding-agent/src/web/search/providers/exa.ts @@ -6,10 +6,10 @@ * Requests per-result summaries via `contents.summary` and synthesizes * them into a combined `answer` string on the SearchResponse. */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import { settings } from "../../../config/settings"; import { callExaTool, findApiKey, isSearchResponse } from "../../../exa/mcp-client"; -import type { AgentStorage } from "../../../session/agent-storage"; + import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { dateToAgeSeconds } from "../utils"; @@ -250,7 +250,7 @@ export class ExaProvider extends SearchProvider { readonly id = "exa"; readonly label = "Exa"; - isAvailable(_storage: AgentStorage): boolean { + isAvailable(_authStorage: AuthStorage): boolean { try { if (settings.get("exa.enabled") === false || settings.get("exa.enableSearch") === false) { return false; @@ -261,7 +261,7 @@ export class ExaProvider extends SearchProvider { return true; } - search(params: SearchParams, _storage: AgentStorage): Promise { + search(params: SearchParams): Promise { return searchExa({ query: params.query, num_results: params.numSearchResults ?? params.limit, diff --git a/packages/coding-agent/src/web/search/providers/gemini.ts b/packages/coding-agent/src/web/search/providers/gemini.ts index 17fa65994..215e2a5cf 100644 --- a/packages/coding-agent/src/web/search/providers/gemini.ts +++ b/packages/coding-agent/src/web/search/providers/gemini.ts @@ -2,15 +2,20 @@ * Google Gemini Web Search Provider * * Uses Gemini's Google Search grounding via Cloud Code Assist API. - * Requires OAuth credentials stored in agent.db for provider "google-gemini-cli" or "google-antigravity". - * Returns synthesized answers with citations and source metadata from grounding chunks. + * Auth is resolved through `AuthStorage.getOAuthAccess(...)` for both + * `google-gemini-cli` (stable prod) and `google-antigravity` (daily sandbox) + * — the broker is the sole refresh authority, so this module never opens a + * sibling SQLite store and never POSTs the broker sentinel to a Google token + * endpoint. */ -import { ANTIGRAVITY_SYSTEM_INSTRUCTION, getAntigravityUserAgent, getGeminiCliHeaders } from "@oh-my-pi/pi-ai"; -import { refreshAntigravityToken } from "@oh-my-pi/pi-ai/utils/oauth/google-antigravity"; -import { refreshGoogleCloudToken } from "@oh-my-pi/pi-ai/utils/oauth/google-gemini-cli"; +import { + ANTIGRAVITY_SYSTEM_INSTRUCTION, + type AuthStorage, + getAntigravityUserAgent, + getGeminiCliHeaders, +} from "@oh-my-pi/pi-ai"; import { fetchWithRetry } from "@oh-my-pi/pi-utils"; -import type { AgentStorage } from "../../../session/agent-storage"; import type { SearchCitation, SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import type { SearchParams } from "./base"; @@ -26,6 +31,9 @@ const MAX_RETRIES = 3; const BASE_DELAY_MS = 1000; const RATE_LIMIT_BUDGET_MS = 5 * 60 * 1000; +const GEMINI_PROVIDERS = ["google-gemini-cli", "google-antigravity"] as const; +type GeminiProviderId = (typeof GEMINI_PROVIDERS)[number]; + interface GeminiToolParams { google_search?: Record; code_execution?: Record; @@ -41,6 +49,8 @@ export interface GeminiSearchParams extends GeminiToolParams { /** Sampling temperature (0–1). Lower = more focused/factual. */ temperature?: number; signal?: AbortSignal; + authStorage: AuthStorage; + sessionId?: string; } export function buildGeminiRequestTools(params: GeminiToolParams): Array>> { @@ -54,125 +64,40 @@ export function buildGeminiRequestTools(params: GeminiToolParams): Array { - if (!auth.refreshToken) return false; - try { - const refreshed = auth.isAntigravity - ? await refreshAntigravityToken(auth.refreshToken, auth.projectId) - : await refreshGoogleCloudToken(auth.refreshToken, auth.projectId); - auth.accessToken = refreshed.access; - auth.refreshToken = refreshed.refresh ?? auth.refreshToken; - auth.storage.updateAuthCredential(auth.credentialId, { - ...auth.credential, - access: auth.accessToken, - refresh: auth.refreshToken, - expires: refreshed.expires, - }); - auth.credential.access = auth.accessToken; - auth.credential.refresh = auth.refreshToken; - auth.credential.expires = refreshed.expires; - return true; - } catch { - return false; - } } /** - * Finds valid Gemini OAuth credentials from agent.db. - * Checks google-gemini-cli first (stable prod), then google-antigravity (daily sandbox). - * @returns OAuth credential with access token and project ID, or null if none found + * Walks the configured Gemini OAuth providers in deterministic order and + * returns the first one that yields a usable access token + projectId via + * {@link AuthStorage.getOAuthAccess}. AuthStorage handles refresh + broker + * routing internally; this helper never touches refresh tokens directly. */ -export async function findGeminiAuth(storage: AgentStorage): Promise { - const expiryBuffer = 5 * 60 * 1000; // 5 minutes - const now = Date.now(); - - // Try providers in deterministic order: gemini-cli first, then antigravity - const providers = ["google-gemini-cli", "google-antigravity"] as const; - - for (const provider of providers) { - const records = storage.listAuthCredentials(provider); - - for (const record of records) { - const credential = record.credential; - if (credential.type !== "oauth") continue; - - const oauthCred = credential as GeminiOAuthCredential; - if (!oauthCred.access) continue; - - // Get projectId from credential - const projectId = oauthCred.projectId; - if (!projectId) continue; - - // Check if token is expired (or about to expire) - if (oauthCred.expires <= now + expiryBuffer) { - // Try to refresh if we have a refresh token - if (oauthCred.refresh) { - try { - const refreshed = - provider === "google-antigravity" - ? await refreshAntigravityToken(oauthCred.refresh, projectId) - : await refreshGoogleCloudToken(oauthCred.refresh, projectId); - // Update the credential in storage - const updated = { - ...oauthCred, - access: refreshed.access, - refresh: refreshed.refresh ?? oauthCred.refresh, - expires: refreshed.expires, - }; - storage.updateAuthCredential(record.id, updated); - return { - accessToken: refreshed.access, - refreshToken: refreshed.refresh ?? oauthCred.refresh, - projectId, - isAntigravity: provider === "google-antigravity", - storage, - credentialId: record.id, - credential: updated, - }; - } catch { - // Refresh failed, skip this credential - continue; - } - } - // No refresh token or refresh failed - continue; - } - - return { - accessToken: oauthCred.access, - refreshToken: oauthCred.refresh, - projectId, - isAntigravity: provider === "google-antigravity", - storage, - credentialId: record.id, - credential: oauthCred, - }; - } +export async function findGeminiAuth( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, +): Promise { + for (const provider of GEMINI_PROVIDERS) { + const access = await authStorage.getOAuthAccess(provider, sessionId, { signal }); + if (!access?.accessToken || !access.projectId) continue; + return { + accessToken: access.accessToken, + projectId: access.projectId, + isAntigravity: provider === "google-antigravity", + }; } - return null; } +function hasGeminiOAuth(authStorage: AuthStorage): boolean { + return GEMINI_PROVIDERS.some((provider: GeminiProviderId) => authStorage.hasOAuth(provider)); +} + /** Cloud Code Assist API response types */ interface GeminiGroundingChunk { web?: { @@ -218,20 +143,20 @@ interface CloudCodeResponseChunk { /** * Calls the Cloud Code Assist API with Google Search grounding enabled. - * @param auth - Authentication info (access token and project ID) - * @param query - Search query from the user - * @param systemPrompt - Optional system prompt - * @returns Parsed response with answer, sources, and usage - * @throws {SearchProviderError} If the API request fails + * + * If a request returns a refreshable auth failure (401/403/auth-flavoured 400), + * we ask AuthStorage to invalidate + refresh the credential and retry once. + * Provider-direct refresh helpers are intentionally not used: AuthStorage owns + * the single-flight refresh and broker round-trip. */ async function callGeminiSearch( auth: GeminiAuth, query: string, - systemPrompt?: string, - maxOutputTokens?: number, - temperature?: number, - toolParams: GeminiToolParams = {}, - signal?: AbortSignal, + systemPrompt: string | undefined, + maxOutputTokens: number | undefined, + temperature: number | undefined, + toolParams: GeminiToolParams, + signal: AbortSignal | undefined, ): Promise<{ answer: string; sources: SearchSource[]; @@ -310,29 +235,13 @@ async function callGeminiSearch( const urlFor = (attempt: number) => `${endpoints[Math.min(attempt, endpoints.length - 1)]}/v1internal:streamGenerateContent?alt=sse`; - let response = await fetchWithRetry(urlFor, { + const response = await fetchWithRetry(urlFor, { ...buildInit(), maxAttempts: MAX_RETRIES + 1, defaultDelayMs: attempt => BASE_DELAY_MS * 2 ** attempt, maxDelayMs: RATE_LIMIT_BUDGET_MS, }); - if (!response.ok) { - const errorText = await response.clone().text(); - const canRefreshAuth = - response.status === 401 || - response.status === 403 || - (response.status === 400 && /api key not valid|invalid credentials|invalid authentication/i.test(errorText)); - if (canRefreshAuth && (await refreshGeminiAuth(auth))) { - response = await fetchWithRetry(urlFor, { - ...buildInit(), - maxAttempts: MAX_RETRIES + 1, - defaultDelayMs: attempt => BASE_DELAY_MS * 2 ** attempt, - maxDelayMs: RATE_LIMIT_BUDGET_MS, - }); - } - } - if (!response.ok) { const errorText = await response.text(); const classified = classifyProviderHttpError("gemini", response.status, errorText); @@ -476,13 +385,9 @@ async function callGeminiSearch( /** * Executes a web search using Google Gemini with Google Search grounding. - * Requires OAuth credentials stored in agent.db for provider "google-gemini-cli" or "google-antigravity". - * @param params - Search parameters including query and optional settings - * @returns Search response with synthesized answer, sources, and citations - * @throws {Error} If no Gemini OAuth credentials are configured */ -export async function searchGemini(params: GeminiSearchParams, storage: AgentStorage): Promise { - const auth = await findGeminiAuth(storage); +export async function searchGemini(params: GeminiSearchParams): Promise { + const auth = await findGeminiAuth(params.authStorage, params.sessionId, params.signal); if (!auth) { throw new Error( "No Gemini OAuth credentials found. Login with 'omp /login google-gemini-cli' or 'omp /login google-antigravity' to enable Gemini web search.", @@ -505,7 +410,6 @@ export async function searchGemini(params: GeminiSearchParams, storage: AgentSto let sources = result.sources; - // Apply num_results limit if specified if (params.num_results && sources.length > params.num_results) { sources = sources.slice(0, params.num_results); } @@ -526,24 +430,26 @@ export class GeminiProvider extends SearchProvider { readonly id = "gemini"; readonly label = "Gemini"; - isAvailable(storage: AgentStorage) { - return findGeminiAuth(storage).then(Boolean); + isAvailable(authStorage: AuthStorage): boolean { + // Cheap, in-memory check — avoids driving the refresh pipeline during + // the provider-chain probe. `searchGemini` calls `getOAuthAccess` which + // will refresh lazily on the actual request. + return hasGeminiOAuth(authStorage); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchGemini( - { - query: params.query, - system_prompt: params.systemPrompt, - num_results: params.numSearchResults ?? params.limit, - max_output_tokens: params.maxOutputTokens, - temperature: params.temperature, - google_search: params.googleSearch, - code_execution: params.codeExecution, - url_context: params.urlContext, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchGemini({ + query: params.query, + system_prompt: params.systemPrompt, + num_results: params.numSearchResults ?? params.limit, + max_output_tokens: params.maxOutputTokens, + temperature: params.temperature, + google_search: params.googleSearch, + code_execution: params.codeExecution, + url_context: params.urlContext, + signal: params.signal, + authStorage: params.authStorage, + sessionId: params.sessionId, + }); } } diff --git a/packages/coding-agent/src/web/search/providers/jina.ts b/packages/coding-agent/src/web/search/providers/jina.ts index 3de174e4c..49559ff1e 100644 --- a/packages/coding-agent/src/web/search/providers/jina.ts +++ b/packages/coding-agent/src/web/search/providers/jina.ts @@ -5,8 +5,7 @@ * cleaned content. */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import type { SearchParams } from "./base"; @@ -88,11 +87,11 @@ export class JinaProvider extends SearchProvider { readonly id = "jina"; readonly label = "Jina"; - isAvailable(_storage: AgentStorage): boolean { + isAvailable(_authStorage: AuthStorage): boolean { return !!findApiKey(); } - search(params: SearchParams, _storage: AgentStorage): Promise { + search(params: SearchParams): Promise { return searchJina({ query: params.query, num_results: params.numSearchResults ?? params.limit, diff --git a/packages/coding-agent/src/web/search/providers/kagi.ts b/packages/coding-agent/src/web/search/providers/kagi.ts index 7e209764f..b47285436 100644 --- a/packages/coding-agent/src/web/search/providers/kagi.ts +++ b/packages/coding-agent/src/web/search/providers/kagi.ts @@ -3,10 +3,10 @@ * * Thin wrapper that adapts shared Kagi API utilities to SearchResponse shape. */ -import type { AgentStorage } from "../../../session/agent-storage"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import type { SearchResponse } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; -import { findKagiApiKey, KagiApiError, searchWithKagi } from "../../kagi"; +import { KagiApiError, searchWithKagi } from "../../kagi"; import { clampNumResults } from "../utils"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; @@ -16,14 +16,13 @@ const DEFAULT_NUM_RESULTS = 10; const MAX_NUM_RESULTS = 40; /** Execute Kagi web search. */ -export async function searchKagi( - params: { - query: string; - num_results?: number; - signal?: AbortSignal; - }, - storage: AgentStorage, -): Promise { +export async function searchKagi(params: { + query: string; + num_results?: number; + signal?: AbortSignal; + authStorage: AuthStorage; + sessionId?: string; +}): Promise { const numResults = clampNumResults(params.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); try { @@ -31,9 +30,10 @@ export async function searchKagi( params.query, { limit: numResults, + sessionId: params.sessionId, signal: params.signal, }, - storage, + params.authStorage, ); return { @@ -59,22 +59,17 @@ export class KagiProvider extends SearchProvider { readonly id = "kagi"; readonly label = "Kagi"; - async isAvailable(storage: AgentStorage) { - try { - return !!findKagiApiKey(storage); - } catch { - return false; - } + isAvailable(authStorage: AuthStorage): boolean { + return authStorage.hasAuth("kagi"); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchKagi( - { - query: params.query, - num_results: params.numSearchResults ?? params.limit, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchKagi({ + query: params.query, + num_results: params.numSearchResults ?? params.limit, + signal: params.signal, + authStorage: params.authStorage, + sessionId: params.sessionId, + }); } } diff --git a/packages/coding-agent/src/web/search/providers/kimi.ts b/packages/coding-agent/src/web/search/providers/kimi.ts index 115c3bde5..2bc8d53b1 100644 --- a/packages/coding-agent/src/web/search/providers/kimi.ts +++ b/packages/coding-agent/src/web/search/providers/kimi.ts @@ -4,15 +4,15 @@ * Uses Moonshot Kimi Code search API to retrieve web results. * Endpoint: POST https://api.kimi.com/coding/v1/search */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { $env } from "@oh-my-pi/pi-utils"; -import type { AgentStorage } from "../../../session/agent-storage"; + import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; -import { classifyProviderHttpError, findCredential, withHardTimeout } from "./utils"; +import { classifyProviderHttpError, withHardTimeout } from "./utils"; const KIMI_SEARCH_URL = "https://api.kimi.com/coding/v1/search"; @@ -25,6 +25,8 @@ export interface KimiSearchParams { num_results?: number; include_content?: boolean; signal?: AbortSignal; + authStorage: AuthStorage; + sessionId?: string; } interface KimiSearchResult { @@ -52,14 +54,20 @@ function resolveBaseUrl(): string { return asTrimmed($env.MOONSHOT_SEARCH_BASE_URL) ?? asTrimmed($env.KIMI_SEARCH_BASE_URL) ?? KIMI_SEARCH_URL; } -/** Find Kimi search credentials from environment or agent.db credentials. */ -function findApiKey(storage: AgentStorage): string | null { - const envKey = - asTrimmed($env.MOONSHOT_SEARCH_API_KEY) ?? - asTrimmed($env.KIMI_SEARCH_API_KEY) ?? - getEnvApiKey("moonshot") ?? - null; - return findCredential(storage, envKey, "moonshot", "kimi-code"); +/** Find Kimi search credentials from environment or AuthStorage. */ +async function findApiKey( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, +): Promise { + const envKey = asTrimmed($env.MOONSHOT_SEARCH_API_KEY) ?? asTrimmed($env.KIMI_SEARCH_API_KEY); + if (envKey) return envKey; + + return ( + (await authStorage.getApiKey("moonshot", sessionId, { signal })) ?? + (await authStorage.getApiKey("kimi-code", sessionId, { signal })) ?? + null + ); } async function callKimiSearch( @@ -99,8 +107,8 @@ async function callKimiSearch( } /** Execute Kimi web search. */ -export async function searchKimi(params: KimiSearchParams, storage: AgentStorage): Promise { - const apiKey = findApiKey(storage); +export async function searchKimi(params: KimiSearchParams): Promise { + const apiKey = await findApiKey(params.authStorage, params.sessionId, params.signal); if (!apiKey) { throw new Error( "Kimi search credentials not found. Set MOONSHOT_SEARCH_API_KEY, KIMI_SEARCH_API_KEY, MOONSHOT_API_KEY, or login with 'omp /login moonshot'.", @@ -142,18 +150,22 @@ export class KimiProvider extends SearchProvider { readonly id = "kimi"; readonly label = "Kimi"; - isAvailable(storage: AgentStorage): boolean { - return !!findApiKey(storage); - } - - search(params: SearchParams, storage: AgentStorage): Promise { - return searchKimi( - { - query: params.query, - num_results: params.numSearchResults ?? params.limit, - signal: params.signal, - }, - storage, + isAvailable(authStorage: AuthStorage): boolean { + return ( + !!asTrimmed($env.MOONSHOT_SEARCH_API_KEY) || + !!asTrimmed($env.KIMI_SEARCH_API_KEY) || + authStorage.hasAuth("moonshot") || + authStorage.hasAuth("kimi-code") ); } + + search(params: SearchParams): Promise { + return searchKimi({ + query: params.query, + num_results: params.numSearchResults ?? params.limit, + signal: params.signal, + authStorage: params.authStorage, + sessionId: params.sessionId, + }); + } } diff --git a/packages/coding-agent/src/web/search/providers/parallel.ts b/packages/coding-agent/src/web/search/providers/parallel.ts index ba4e273a8..19f1b8164 100644 --- a/packages/coding-agent/src/web/search/providers/parallel.ts +++ b/packages/coding-agent/src/web/search/providers/parallel.ts @@ -1,14 +1,153 @@ -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; -import { findParallelApiKey, ParallelApiError, searchWithParallel } from "../../parallel"; +import { ParallelApiError, type ParallelSearchResult, type ParallelSearchSource } from "../../parallel"; import { clampNumResults } from "../utils"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; -import { classifyProviderHttpError, toSearchSources } from "./utils"; +import { classifyProviderHttpError, toSearchSources, withHardTimeout } from "./utils"; const DEFAULT_NUM_RESULTS = 10; const MAX_NUM_RESULTS = 40; +const PARALLEL_SEARCH_URL = "https://api.parallel.ai/v1beta/search"; +const PARALLEL_BETA_HEADER = "search-extract-2025-10-10"; + +function isObject(value: unknown): value is object { + return typeof value === "object" && value !== null; +} + +function getOwnValue(value: object, key: string): unknown { + return Object.getOwnPropertyDescriptor(value, key)?.value; +} + +function getString(value: object, key: string): string | undefined { + const field = getOwnValue(value, key); + return typeof field === "string" ? field : undefined; +} + +function getObjectArray(value: object, key: string): object[] { + const field = getOwnValue(value, key); + return Array.isArray(field) ? field.filter(isObject) : []; +} + +function getStringArray(value: object, key: string): string[] { + const field = getOwnValue(value, key); + return Array.isArray(field) ? field.filter((item): item is string => typeof item === "string") : []; +} + +function extractParallelErrorMessage(payload: unknown): string | null { + if (!isObject(payload)) return null; + + const directMessage = getString(payload, "message") ?? getString(payload, "detail") ?? getString(payload, "error"); + if (directMessage && directMessage.trim().length > 0) { + return directMessage.trim(); + } + + const errorObject = getOwnValue(payload, "error"); + if (isObject(errorObject)) { + const nestedMessage = getString(errorObject, "message") ?? getString(errorObject, "detail"); + if (nestedMessage && nestedMessage.trim().length > 0) { + return nestedMessage.trim(); + } + } + + return null; +} + +function createParallelApiError(statusCode: number, detail?: string): ParallelApiError { + return new ParallelApiError( + detail ? `Parallel API error (${statusCode}): ${detail}` : `Parallel API error (${statusCode})`, + statusCode, + ); +} + +function parseParallelErrorResponse(statusCode: number, responseText: string): ParallelApiError { + const trimmedResponseText = responseText.trim(); + if (trimmedResponseText.length === 0) { + return createParallelApiError(statusCode); + } + + try { + const payload: unknown = JSON.parse(trimmedResponseText); + return createParallelApiError(statusCode, extractParallelErrorMessage(payload) ?? trimmedResponseText); + } catch { + return createParallelApiError(statusCode, trimmedResponseText); + } +} + +function parseSearchPayload(payload: unknown): ParallelSearchResult { + if (!isObject(payload)) { + throw new ParallelApiError("Parallel search returned an invalid response payload."); + } + + const requestId = getString(payload, "search_id") ?? ""; + const rawResults = getObjectArray(payload, "results"); + const sources: ParallelSearchSource[] = []; + + for (const item of rawResults) { + const url = getString(item, "url"); + if (!url) continue; + + const excerpts = getStringArray(item, "excerpts"); + const snippet = excerpts.length > 0 ? excerpts.join("\n\n") : undefined; + sources.push({ + title: getString(item, "title") ?? url, + url, + snippet, + publishedDate: getString(item, "publish_date"), + excerpts, + }); + } + + return { + requestId, + sources, + warnings: [], + usage: [], + }; +} + +async function searchWithAuthStorage( + objective: string, + queries: string[], + params: { + signal?: AbortSignal; + }, + authStorage: AuthStorage, + sessionId?: string, +): Promise { + const apiKey = await authStorage.getApiKey("parallel", sessionId, { signal: params.signal }); + if (!apiKey) { + throw new ParallelApiError( + "Parallel credentials not found. Set PARALLEL_API_KEY or login with 'omp /login parallel'.", + ); + } + + const response = await fetch(PARALLEL_SEARCH_URL, { + method: "POST", + headers: { + Accept: "application/json", + "Content-Type": "application/json", + "x-api-key": apiKey, + "parallel-beta": PARALLEL_BETA_HEADER, + }, + body: JSON.stringify({ + objective, + search_queries: queries, + mode: "fast", + excerpts: { + max_chars_per_result: 10_000, + }, + }), + signal: withHardTimeout(params.signal), + }); + if (!response.ok) { + throw parseParallelErrorResponse(response.status, await response.text()); + } + + const payload: unknown = await response.json(); + return parseSearchPayload(payload); +} export async function searchParallel( params: { @@ -16,20 +155,20 @@ export async function searchParallel( num_results?: number; signal?: AbortSignal; }, - storage: AgentStorage, + authStorage: AuthStorage, + sessionId?: string, ): Promise { const numResults = clampNumResults(params.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); try { - const result = await searchWithParallel( + const result = await searchWithAuthStorage( params.query, [params.query], { - mode: "fast", - maxCharsPerResult: 10_000, signal: params.signal, }, - storage, + authStorage, + sessionId, ); return { @@ -53,22 +192,19 @@ export class ParallelProvider extends SearchProvider { readonly id = "parallel"; readonly label = "Parallel"; - async isAvailable(storage: AgentStorage) { - try { - return !!findParallelApiKey(storage); - } catch { - return false; - } + isAvailable(authStorage: AuthStorage) { + return !!getEnvApiKey("parallel") || authStorage.hasAuth("parallel"); } - search(params: SearchParams, storage: AgentStorage): Promise { + search(params: SearchParams): Promise { return searchParallel( { query: params.query, num_results: params.numSearchResults ?? params.limit, signal: params.signal, }, - storage, + params.authStorage, + params.sessionId, ); } } diff --git a/packages/coding-agent/src/web/search/providers/perplexity.ts b/packages/coding-agent/src/web/search/providers/perplexity.ts index 548ba7dee..b2c046563 100644 --- a/packages/coding-agent/src/web/search/providers/perplexity.ts +++ b/packages/coding-agent/src/web/search/providers/perplexity.ts @@ -3,13 +3,12 @@ * * Supports three auth modes: * - Cookies (`PERPLEXITY_COOKIES`) via `www.perplexity.ai/rest/sse/perplexity_ask` - * - OAuth JWT (stored in `agent.db`) via `www.perplexity.ai/rest/sse/perplexity_ask` + * - OAuth/session bearer via `AuthStorage` and `www.perplexity.ai/rest/sse/perplexity_ask` * - API key (`PERPLEXITY_API_KEY`) via `api.perplexity.ai/chat/completions` */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import { $env, readSseJson } from "@oh-my-pi/pi-utils"; -import type { AgentStorage } from "../../../session/agent-storage"; import type { PerplexityMessageOutput, PerplexityRequest, @@ -34,12 +33,6 @@ const OAUTH_EXPIRY_BUFFER_MS = 5 * 60 * 1000; const OAUTH_API_VERSION = "2.18"; const OAUTH_USER_AGENT = "Perplexity/641 CFNetwork/1568 Darwin/25.2.0"; -interface PerplexityOAuthCredential { - type: "oauth"; - access: string; - expires: number; -} - type PerplexityAuth = | { type: "api_key"; @@ -168,6 +161,8 @@ export interface PerplexitySearchParams { temperature?: number; /** Number of search results to retrieve. Defaults to 10. */ num_search_results?: number; + authStorage: AuthStorage; + sessionId?: string; } /** Find PERPLEXITY_API_KEY from environment or .env files (also checks PPLX_API_KEY) */ @@ -194,40 +189,47 @@ function jwtExpiryMs(token: string): number | undefined { } } -async function findOAuthToken(storage: AgentStorage): Promise { - const now = Date.now(); +async function findOAuthToken( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, + envApiKey: string | null, +): Promise { try { - const records = storage.listAuthCredentials("perplexity"); - for (const record of records) { - if (record.credential.type !== "oauth") continue; - const credential = record.credential as PerplexityOAuthCredential; - if (!credential.access) continue; - // Trust the JWT's own `exp` claim if it has one; otherwise treat as - // non-expiring. The stored `expires` field is unreliable: older logins - // wrote `loginTime + 1h` even though Perplexity JWTs typically lack `exp`. - const jwtExpiry = jwtExpiryMs(credential.access); - if (jwtExpiry !== undefined && jwtExpiry <= now + OAUTH_EXPIRY_BUFFER_MS) continue; - return credential.access; - } + const token = await authStorage.getApiKey("perplexity", sessionId, { signal }); + // `getApiKey` falls back to PERPLEXITY_API_KEY; do not route that env key + // through the OAuth/web endpoint. + if (!token || (envApiKey && token === envApiKey)) return null; + // Trust the JWT's own `exp` claim if it has one; otherwise treat as + // non-expiring. Perplexity session JWTs commonly omit `exp`. + const jwtExpiry = jwtExpiryMs(token); + if (jwtExpiry !== undefined && jwtExpiry <= Date.now() + OAUTH_EXPIRY_BUFFER_MS) return null; + return token; } catch { return null; } - return null; } -async function findPerplexityAuth(storage: AgentStorage): Promise { +async function findPerplexityAuth( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, +): Promise { // 1. PERPLEXITY_COOKIES env var const cookies = $env.PERPLEXITY_COOKIES?.trim(); if (cookies) { return { type: "cookies", cookies }; } - // 2. OAuth token from agent.db - const oauthToken = await findOAuthToken(storage); + + const apiKey = findApiKey(); + + // 2. OAuth/session bearer from AuthStorage. + const oauthToken = await findOAuthToken(authStorage, sessionId, signal, apiKey); if (oauthToken) { return { type: "oauth", token: oauthToken }; } + // 3. PERPLEXITY_API_KEY env var - const apiKey = findApiKey(); if (apiKey) { return { type: "api_key", token: apiKey }; } @@ -491,8 +493,8 @@ function applySourceLimit(result: SearchResponse, limit?: number): SearchRespons } /** Execute Perplexity web search */ -export async function searchPerplexity(params: PerplexitySearchParams, storage: AgentStorage): Promise { - const auth = await findPerplexityAuth(storage); +export async function searchPerplexity(params: PerplexitySearchParams): Promise { + const auth = await findPerplexityAuth(params.authStorage, params.sessionId, params.signal); if (!auth) { throw new Error("Perplexity auth not found. Set PERPLEXITY_COOKIES, PERPLEXITY_API_KEY, or login via OAuth."); } @@ -550,27 +552,22 @@ export class PerplexityProvider extends SearchProvider { readonly id = "perplexity"; readonly label = "Perplexity"; - async isAvailable(storage: AgentStorage) { - try { - return !!(await findPerplexityAuth(storage)); - } catch { - return false; - } + isAvailable(authStorage: AuthStorage): boolean { + return !!$env.PERPLEXITY_COOKIES?.trim() || authStorage.hasAuth("perplexity") || !!findApiKey(); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchPerplexity( - { - signal: params.signal, - query: params.query, - temperature: params.temperature, - max_tokens: params.maxOutputTokens, - num_search_results: params.numSearchResults, - system_prompt: params.systemPrompt, - search_recency_filter: params.recency, - num_results: params.limit, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchPerplexity({ + signal: params.signal, + query: params.query, + temperature: params.temperature, + max_tokens: params.maxOutputTokens, + num_search_results: params.numSearchResults, + system_prompt: params.systemPrompt, + search_recency_filter: params.recency, + num_results: params.limit, + authStorage: params.authStorage, + sessionId: params.sessionId, + }); } } diff --git a/packages/coding-agent/src/web/search/providers/searxng.ts b/packages/coding-agent/src/web/search/providers/searxng.ts index 51733258b..8eb05e500 100644 --- a/packages/coding-agent/src/web/search/providers/searxng.ts +++ b/packages/coding-agent/src/web/search/providers/searxng.ts @@ -25,8 +25,9 @@ * Reference: https://docs.searxng.org/dev/search_api.html */ +import type { AuthStorage } from "@oh-my-pi/pi-ai"; + import { settings } from "../../../config/settings"; -import type { AgentStorage } from "../../../session/agent-storage"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; @@ -289,7 +290,7 @@ export class SearXNGProvider extends SearchProvider { readonly id = "searxng"; readonly label = "SearXNG"; - isAvailable(_storage: AgentStorage): boolean { + isAvailable(_authStorage: AuthStorage): boolean { try { return !!findEndpoint(); } catch { @@ -297,7 +298,7 @@ export class SearXNGProvider extends SearchProvider { } } - search(params: SearchParams, _storage: AgentStorage): Promise { + search(params: SearchParams): Promise { return searchSearXNG({ query: params.query, num_results: params.numSearchResults ?? params.limit, diff --git a/packages/coding-agent/src/web/search/providers/synthetic.ts b/packages/coding-agent/src/web/search/providers/synthetic.ts index e74bb6c4b..c531a9a65 100644 --- a/packages/coding-agent/src/web/search/providers/synthetic.ts +++ b/packages/coding-agent/src/web/search/providers/synthetic.ts @@ -5,13 +5,12 @@ * Endpoint: POST https://api.synthetic.new/v2/search */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; -import { classifyProviderHttpError, findCredential, withHardTimeout } from "./utils"; +import { classifyProviderHttpError, withHardTimeout } from "./utils"; const SYNTHETIC_SEARCH_URL = "https://api.synthetic.new/v2/search"; @@ -26,9 +25,13 @@ interface SyntheticSearchResponse { results: SyntheticSearchResult[]; } -/** Find Synthetic API key from environment or agent.db credentials. */ -export function findApiKey(storage: AgentStorage): string | null { - return findCredential(storage, getEnvApiKey("synthetic"), "synthetic"); +/** Resolve Synthetic API key through the shared auth storage pipeline. */ +export function findApiKey( + authStorage: AuthStorage, + sessionId?: string, + signal?: AbortSignal, +): Promise { + return authStorage.getApiKey("synthetic", sessionId, { signal }); } /** Call Synthetic search API. */ @@ -62,15 +65,8 @@ async function callSyntheticSearch( } /** Execute Synthetic web search. */ -export async function searchSynthetic( - params: { - query: string; - num_results?: number; - signal?: AbortSignal; - }, - storage: AgentStorage, -): Promise { - const apiKey = findApiKey(storage); +export async function searchSynthetic(params: SearchParams): Promise { + const apiKey = await findApiKey(params.authStorage, params.sessionId, params.signal); if (!apiKey) { throw new Error("Synthetic credentials not found. Set SYNTHETIC_API_KEY or login with 'omp /login synthetic'."); } @@ -88,7 +84,8 @@ export async function searchSynthetic( }); } - const limitedSources = params.num_results ? sources.slice(0, params.num_results) : sources; + const numResults = params.numSearchResults ?? params.limit; + const limitedSources = numResults ? sources.slice(0, numResults) : sources; return { provider: "synthetic", @@ -101,18 +98,11 @@ export class SyntheticProvider extends SearchProvider { readonly id = "synthetic"; readonly label = "Synthetic"; - isAvailable(storage: AgentStorage): boolean { - return !!findApiKey(storage); + isAvailable(authStorage: AuthStorage): boolean { + return authStorage.hasAuth("synthetic") || !!getEnvApiKey("synthetic"); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchSynthetic( - { - query: params.query, - num_results: params.numSearchResults ?? params.limit, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchSynthetic(params); } } diff --git a/packages/coding-agent/src/web/search/providers/tavily.ts b/packages/coding-agent/src/web/search/providers/tavily.ts index 00f9fe555..a32a27856 100644 --- a/packages/coding-agent/src/web/search/providers/tavily.ts +++ b/packages/coding-agent/src/web/search/providers/tavily.ts @@ -4,14 +4,13 @@ * Uses Tavily's agent-focused search API to return structured results with an * optional synthesized answer. */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { clampNumResults, dateToAgeSeconds } from "../utils"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; -import { classifyProviderHttpError, findCredential, withHardTimeout } from "./utils"; +import { classifyProviderHttpError, withHardTimeout } from "./utils"; const TAVILY_SEARCH_URL = "https://api.tavily.com/search"; const DEFAULT_NUM_RESULTS = 5; @@ -59,9 +58,13 @@ function getErrorMessage(value: unknown): string | null { return null; } -/** Find Tavily API key from environment or agent.db credentials. */ -export function findApiKey(storage: AgentStorage): string | null { - return findCredential(storage, getEnvApiKey("tavily"), "tavily"); +/** Find Tavily API key through AuthStorage's unified refresh pipeline. */ +export async function findApiKey( + authStorage: AuthStorage, + sessionId: string | undefined, + signal: AbortSignal | undefined, +): Promise { + return (await authStorage.getApiKey("tavily", sessionId, { signal })) ?? null; } /** Exported for testing. Builds the Tavily request body from unified params. */ @@ -117,16 +120,22 @@ async function callTavilySearch(apiKey: string, params: TavilySearchParams): Pro } /** Execute Tavily web search. */ -export async function searchTavily(params: TavilySearchParams, storage: AgentStorage): Promise { - const apiKey = findApiKey(storage); +export async function searchTavily(params: SearchParams): Promise { + const tavilyParams: TavilySearchParams = { + query: params.query, + num_results: params.numSearchResults ?? params.limit, + recency: params.recency, + signal: params.signal, + }; + const apiKey = await findApiKey(params.authStorage, params.sessionId, params.signal); if (!apiKey) { throw new Error( - 'Tavily credentials not found. Set TAVILY_API_KEY or store an API key for provider "tavily" in agent.db.', + 'Tavily credentials not found. Set TAVILY_API_KEY or configure an API key for provider "tavily".', ); } - const numResults = clampNumResults(params.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); - const response = await callTavilySearch(apiKey, params); + const numResults = clampNumResults(tavilyParams.num_results, DEFAULT_NUM_RESULTS, MAX_NUM_RESULTS); + const response = await callTavilySearch(apiKey, tavilyParams); const sources: SearchSource[] = []; for (const result of response.results ?? []) { @@ -154,19 +163,11 @@ export class TavilyProvider extends SearchProvider { readonly id = "tavily"; readonly label = "Tavily"; - isAvailable(storage: AgentStorage): boolean { - return !!findApiKey(storage); + isAvailable(authStorage: AuthStorage): boolean { + return authStorage.hasAuth("tavily") || !!getEnvApiKey("tavily"); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchTavily( - { - query: params.query, - num_results: params.numSearchResults ?? params.limit, - recency: params.recency, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchTavily(params); } } diff --git a/packages/coding-agent/src/web/search/providers/zai.ts b/packages/coding-agent/src/web/search/providers/zai.ts index f8a466804..4795038b8 100644 --- a/packages/coding-agent/src/web/search/providers/zai.ts +++ b/packages/coding-agent/src/web/search/providers/zai.ts @@ -4,15 +4,14 @@ * Calls Z.AI's remote MCP server (`webSearchPrime`) and adapts results into * the unified SearchResponse shape used by the web search tool. */ -import { getEnvApiKey } from "@oh-my-pi/pi-ai"; -import type { AgentStorage } from "../../../session/agent-storage"; +import { type AuthStorage, getEnvApiKey } from "@oh-my-pi/pi-ai"; import { asRecord, asString } from "../../../web/scrapers/utils"; import type { SearchResponse, SearchSource } from "../../../web/search/types"; import { SearchProviderError } from "../../../web/search/types"; import { dateToAgeSeconds } from "../utils"; import type { SearchParams } from "./base"; import { SearchProvider } from "./base"; -import { classifyProviderHttpError, findCredential, withHardTimeout } from "./utils"; +import { classifyProviderHttpError, withHardTimeout } from "./utils"; const ZAI_MCP_URL = "https://api.z.ai/api/mcp/web_search_prime/mcp"; const ZAI_TOOL_NAME = "web_search_prime"; @@ -22,6 +21,8 @@ export interface ZaiSearchParams { query: string; num_results?: number; signal?: AbortSignal; + authStorage: AuthStorage; + sessionId?: string; } interface ZaiSearchResult { @@ -52,9 +53,13 @@ interface JsonRpcPayload { error?: JsonRpcError; } -/** Find Z.AI API credentials from environment or saved auth storage. */ -export function findApiKey(storage: AgentStorage): string | null { - return findCredential(storage, getEnvApiKey("zai"), "zai"); +/** Resolve Z.AI API credentials through the unified auth storage pipeline. */ +export async function findApiKey( + authStorage: AuthStorage, + sessionId?: string, + signal?: AbortSignal, +): Promise { + return (await authStorage.getApiKey("zai", sessionId, { signal })) ?? null; } async function callZaiTool(apiKey: string, args: Record, signal?: AbortSignal): Promise { @@ -272,8 +277,8 @@ function toSources(results: ZaiSearchResult[]): SearchSource[] { } /** Execute Z.AI web search via remote MCP endpoint. */ -export async function searchZai(params: ZaiSearchParams, storage: AgentStorage): Promise { - const apiKey = findApiKey(storage); +export async function searchZai(params: ZaiSearchParams): Promise { + const apiKey = await findApiKey(params.authStorage, params.sessionId, params.signal); if (!apiKey) { throw new Error("Z.AI credentials not found. Set ZAI_API_KEY or login with 'omp /login zai'."); } @@ -299,18 +304,17 @@ export class ZaiProvider extends SearchProvider { readonly id = "zai"; readonly label = "Z.AI"; - isAvailable(storage: AgentStorage): boolean { - return !!findApiKey(storage); + isAvailable(authStorage: AuthStorage): Promise | boolean { + return authStorage.hasAuth("zai") || !!getEnvApiKey("zai"); } - search(params: SearchParams, storage: AgentStorage): Promise { - return searchZai( - { - query: params.query, - num_results: params.numSearchResults ?? params.limit, - signal: params.signal, - }, - storage, - ); + search(params: SearchParams): Promise { + return searchZai({ + query: params.query, + num_results: params.numSearchResults ?? params.limit, + signal: params.signal, + authStorage: params.authStorage, + sessionId: params.sessionId, + }); } } diff --git a/packages/coding-agent/test/tools/web-search-codex.test.ts b/packages/coding-agent/test/tools/web-search-codex.test.ts index c131ccafd..849d206a7 100644 --- a/packages/coding-agent/test/tools/web-search-codex.test.ts +++ b/packages/coding-agent/test/tools/web-search-codex.test.ts @@ -1,6 +1,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { hookFetch } from "@oh-my-pi/pi-utils"; -import { AgentStorage } from "../../src/session/agent-storage"; +import type { SearchParams } from "../../src/web/search/providers/base"; import { searchCodex } from "../../src/web/search/providers/codex"; type CapturedRequest = { @@ -174,28 +175,29 @@ function makePlainUrlPunctuationSseResponse(model: string): string { } describe("searchCodex model selection", () => { - const fakeStorage = { - listAuthCredentials: () => [ - { - id: 1, - credential: { - type: "oauth", - access: "test-access-token", - expires: Date.now() + 600_000, - accountId: "acct-test", - }, - }, - ], - updateAuthCredential: () => undefined, - get authStore() { - return null as never; + const fakeAuthStorage = { + async getOAuthAccess() { + return { + accessToken: "test-access-token", + accountId: "acct-test", + }; }, - } as unknown as AgentStorage; + hasOAuth() { + return true; + }, + } as unknown as AuthStorage; let capturedRequest: CapturedRequest | null = null; + function makeSearchParams(query: string): SearchParams { + return { + query, + systemPrompt: "Codex test system prompt", + authStorage: fakeAuthStorage, + }; + } + function mockCodexFetch(responseModel: string, responseBody?: string): Disposable { capturedRequest = null; - vi.spyOn(AgentStorage, "open").mockResolvedValue(fakeStorage); return hookFetch((url, init) => { capturedRequest = { url: typeof url === "string" ? url : url.toString(), @@ -221,33 +223,69 @@ describe("searchCodex model selection", () => { it("uses the built-in default model when PI_CODEX_WEB_SEARCH_MODEL is unset", async () => { delete process.env.PI_CODEX_WEB_SEARCH_MODEL; - using _hook = mockCodexFetch("gpt-5-codex-mini"); + using _hook = mockCodexFetch("gpt-5.4"); - const result = await searchCodex({ query: "default codex model" }, fakeStorage); + const result = await searchCodex(makeSearchParams("default codex model")); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.url).toBe("https://chatgpt.com/backend-api/codex/responses"); - expect(capturedRequest?.body?.model).toBe("gpt-5-codex-mini"); - expect(result.model).toBe("gpt-5-codex-mini"); + expect(capturedRequest?.body?.model).toBe("gpt-5.4"); + expect(result.model).toBe("gpt-5.4"); expect(result.sources).toEqual([{ title: "Example Article", url: "https://example.com/article" }]); }); it("falls back to the default model when PI_CODEX_WEB_SEARCH_MODEL is blank", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = " "; - using _hook = mockCodexFetch("gpt-5-codex-mini"); + using _hook = mockCodexFetch("gpt-5.4"); - const result = await searchCodex({ query: "blank codex model" }, fakeStorage); + const result = await searchCodex(makeSearchParams("blank codex model")); expect(capturedRequest).not.toBeNull(); - expect(capturedRequest?.body?.model).toBe("gpt-5-codex-mini"); - expect(result.model).toBe("gpt-5-codex-mini"); + expect(capturedRequest?.body?.model).toBe("gpt-5.4"); + expect(result.model).toBe("gpt-5.4"); + }); + it("retries the next bundled default when Codex rejects a model for ChatGPT accounts", async () => { + delete process.env.PI_CODEX_WEB_SEARCH_MODEL; + let calls = 0; + capturedRequest = null; + using _hook = hookFetch((url, init) => { + calls += 1; + capturedRequest = { + url: typeof url === "string" ? url : url.toString(), + headers: init?.headers, + body: init?.body ? (JSON.parse(init.body as string) as Record) : null, + }; + + const requestedModel = capturedRequest.body?.model; + if (calls === 1) { + expect(requestedModel).toBe("gpt-5.4"); + return new Response( + JSON.stringify({ + detail: "The 'gpt-5.4' model is not supported when using Codex with a ChatGPT account.", + }), + { status: 400, headers: { "Content-Type": "application/json" } }, + ); + } + + expect(requestedModel).toBe("gpt-5-codex"); + return new Response(makeSseResponse("gpt-5-codex"), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }); + }); + + const result = await searchCodex(makeSearchParams("retry unsupported default")); + + expect(calls).toBe(2); + expect(result.model).toBe("gpt-5-codex"); + expect(result.sources).toEqual([{ title: "Example Article", url: "https://example.com/article" }]); }); it("uses PI_CODEX_WEB_SEARCH_MODEL when provided", async () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4-mini"; using _hook = mockCodexFetch("gpt-5.4-mini"); - const result = await searchCodex({ query: "overridden codex model" }, fakeStorage); + const result = await searchCodex(makeSearchParams("overridden codex model")); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.model).toBe("gpt-5.4-mini"); @@ -258,7 +296,7 @@ describe("searchCodex model selection", () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; using _hook = mockCodexFetch("gpt-5.4", makeMarkdownLinkSseResponse("gpt-5.4")); - const result = await searchCodex({ query: "markdown citations" }, fakeStorage); + const result = await searchCodex(makeSearchParams("markdown citations")); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.tool_choice).toEqual({ type: "web_search" }); @@ -269,7 +307,7 @@ describe("searchCodex model selection", () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; using _hook = mockCodexFetch("gpt-5.4", makePlainUrlSseResponse("gpt-5.4")); - const result = await searchCodex({ query: "plain url citations" }, fakeStorage); + const result = await searchCodex(makeSearchParams("plain url citations")); expect(result.sources).toEqual([ { title: "https://example.com/article", url: "https://example.com/article" }, @@ -281,7 +319,7 @@ describe("searchCodex model selection", () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; using _hook = mockCodexFetch("gpt-5.4", makeMarkdownParenthesesSseResponse("gpt-5.4")); - const result = await searchCodex({ query: "markdown parentheses citations" }, fakeStorage); + const result = await searchCodex(makeSearchParams("markdown parentheses citations")); expect(result.sources).toEqual([ { title: "Function", url: "https://en.wikipedia.org/wiki/Function_(mathematics)" }, @@ -292,7 +330,7 @@ describe("searchCodex model selection", () => { process.env.PI_CODEX_WEB_SEARCH_MODEL = "gpt-5.4"; using _hook = mockCodexFetch("gpt-5.4", makePlainUrlPunctuationSseResponse("gpt-5.4")); - const result = await searchCodex({ query: "plain url punctuation" }, fakeStorage); + const result = await searchCodex(makeSearchParams("plain url punctuation")); expect(result.sources).toEqual([ { title: "https://example.com/article", url: "https://example.com/article" }, @@ -305,7 +343,6 @@ describe("searchCodex model selection", () => { }); it("prefers streamed text when the final item only contains an image placeholder", async () => { - vi.spyOn(AgentStorage, "open").mockResolvedValue(fakeStorage); using _hook = hookFetch(() => { return new Response(makeImagePlaceholderSseResponse("gpt-5.4-mini"), { status: 200, @@ -313,7 +350,7 @@ describe("searchCodex model selection", () => { }); }); - const result = await searchCodex({ query: "responses api store semantics" }, fakeStorage); + const result = await searchCodex(makeSearchParams("responses api store semantics")); expect(result.answer).toBe("OpenAI Responses API defaults `store` to false unless you opt in."); expect(result.sources).toEqual([ diff --git a/packages/coding-agent/test/tools/web-search-gemini.test.ts b/packages/coding-agent/test/tools/web-search-gemini.test.ts index 750ba5447..85671d311 100644 --- a/packages/coding-agent/test/tools/web-search-gemini.test.ts +++ b/packages/coding-agent/test/tools/web-search-gemini.test.ts @@ -1,37 +1,32 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { hookFetch } from "@oh-my-pi/pi-utils"; -import { AgentStorage } from "../../src/session/agent-storage"; import { searchGemini } from "../../src/web/search/providers/gemini"; +const SSE_RESPONSE = + 'data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"Gemini answer"}]}}],"modelVersion":"gemini-2.5-flash"}}\n\n'; + type CapturedRequest = { body: Record | null; }; -const SSE_RESPONSE = - 'data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"Gemini answer"}]}}],"modelVersion":"gemini-2.5-flash"}}\n\n'; - describe("searchGemini tools serialization", () => { let capturedRequest: CapturedRequest | null = null; - const fakeStorage = { - listAuthCredentials: () => [ - { - id: 1, - credential: { - type: "oauth", - access: "test-access-token", - expires: Date.now() + 600_000, - projectId: "test-project", - }, - }, - ], - updateAuthCredential: () => undefined, - authStore: null as never, - } as unknown as AgentStorage; + const fakeAuthStorage = { + async getOAuthAccess() { + return { + accessToken: "test-access-token", + projectId: "test-project", + }; + }, + hasOAuth() { + return true; + }, + } as unknown as AuthStorage; function mockGeminiFetch() { capturedRequest = null; - vi.spyOn(AgentStorage, "open").mockResolvedValue(fakeStorage); return hookFetch((_url, init) => { capturedRequest = { body: init?.body ? (JSON.parse(init.body as string) as Record) : null, @@ -48,9 +43,17 @@ describe("searchGemini tools serialization", () => { capturedRequest = null; }); + function makeParams(query: string) { + return { + query, + authStorage: fakeAuthStorage, + systemPrompt: "Gemini test prompt", + } as const; + } + it("sends default googleSearch tool when no passthrough payloads are provided", async () => { using _hook = mockGeminiFetch(); - await searchGemini({ query: "default tools" }, fakeStorage); + await searchGemini(makeParams("default tools")); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.request).toMatchObject({ @@ -58,15 +61,12 @@ describe("searchGemini tools serialization", () => { }); }); - it("passes through google_search payload into googleSearch tool", async () => { + it("passes through googleSearch payload into googleSearch tool", async () => { using _hook = mockGeminiFetch(); - await searchGemini( - { - query: "google payload", - google_search: { dynamicRetrievalConfig: { mode: "MODE_DYNAMIC" } }, - }, - fakeStorage, - ); + await searchGemini({ + ...makeParams("google payload"), + google_search: { dynamicRetrievalConfig: { mode: "MODE_DYNAMIC" } }, + }); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.request).toMatchObject({ @@ -76,14 +76,11 @@ describe("searchGemini tools serialization", () => { it("includes codeExecution and urlContext tools when provided", async () => { using _hook = mockGeminiFetch(); - await searchGemini( - { - query: "extended tools", - code_execution: {}, - url_context: { allowedDomains: ["example.com"] }, - }, - fakeStorage, - ); + await searchGemini({ + ...makeParams("extended tools"), + code_execution: {}, + url_context: { allowedDomains: ["example.com"] }, + }); expect(capturedRequest).not.toBeNull(); expect(capturedRequest?.body?.request).toMatchObject({ 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 72493eca2..f783b9f83 100644 --- a/packages/coding-agent/test/tools/web-search-kagi.test.ts +++ b/packages/coding-agent/test/tools/web-search-kagi.test.ts @@ -1,17 +1,18 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { hookFetch } from "@oh-my-pi/pi-utils"; import { searchWithKagi } from "../../src/web/kagi"; import { searchKagi } from "../../src/web/search/providers/kagi"; import { SearchProviderError } from "../../src/web/search/types"; -import type { AgentStorage } from "../../src/session/agent-storage"; -const fakeStorage = { - listAuthCredentials: () => [], - updateAuthCredential: () => undefined, - get authStore() { - return null as never; +const fakeAuthStorage = { + async getApiKey() { + return process.env.KAGI_API_KEY ?? undefined; }, -} as unknown as AgentStorage; + hasAuth() { + return Boolean(process.env.KAGI_API_KEY); + }, +} as unknown as AuthStorage; describe("Kagi web search error handling", () => { beforeEach(() => { @@ -36,7 +37,7 @@ describe("Kagi web search error handling", () => { ); try { - await searchKagi({ query: "kagi beta" }, fakeStorage); + await searchKagi({ query: "kagi beta", authStorage: fakeAuthStorage }); expect.unreachable("expected searchKagi to throw"); } catch (error) { expect(error).toBeInstanceOf(SearchProviderError); @@ -48,7 +49,9 @@ describe("Kagi web search error handling", () => { it("falls back to plain text for non-JSON error bodies", async () => { using _hook = hookFetch(() => new Response("upstream unavailable", { status: 503 })); - await expect(searchWithKagi("plain text error", {}, fakeStorage)).rejects.toThrow("Kagi API error (503): upstream unavailable"); + await expect(searchWithKagi("plain text error", {}, fakeAuthStorage)).rejects.toThrow( + "Kagi API error (503): upstream unavailable", + ); }); it("preserves successful search parsing", async () => { @@ -72,7 +75,7 @@ describe("Kagi web search error handling", () => { ), ); - await expect(searchWithKagi("success case", {}, fakeStorage)).resolves.toEqual({ + await expect(searchWithKagi("success case", {}, fakeAuthStorage)).resolves.toEqual({ requestId: "req-kagi-success", sources: [ { diff --git a/packages/coding-agent/test/tools/web-search-parallel.test.ts b/packages/coding-agent/test/tools/web-search-parallel.test.ts index 3b0e9361a..2fa5f6d55 100644 --- a/packages/coding-agent/test/tools/web-search-parallel.test.ts +++ b/packages/coding-agent/test/tools/web-search-parallel.test.ts @@ -1,12 +1,10 @@ -import type { AgentStorage } from "../../src/session/agent-storage"; - import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { AgentStorage } from "../../src/session/agent-storage"; import { searchWithParallel } from "../../src/web/parallel"; import { searchParallel } from "../../src/web/search/providers/parallel"; - - describe("Parallel web search", () => { const fakeStorage = { listAuthCredentials: () => [ @@ -21,8 +19,18 @@ describe("Parallel web search", () => { }, ], updateAuthCredential: () => undefined, - get authStore() { return null as never; }, + get authStore() { + return null as never; + }, } as unknown as AgentStorage; + const fakeAuthStorage = { + async getApiKey() { + return process.env.PARALLEL_API_KEY ?? undefined; + }, + hasAuth() { + return Boolean(process.env.PARALLEL_API_KEY); + }, + } as unknown as AuthStorage; let capturedRequestBody: unknown; @@ -102,7 +110,7 @@ describe("Parallel web search", () => { usage: null, }); - const result = await searchParallel({ query: "alpha search" }, fakeStorage); + const result = await searchParallel({ query: "alpha search" }, fakeAuthStorage); expect(result.provider).toBe("parallel"); expect(result.requestId).toBe("search-parallel-2"); expect(result.sources).toEqual([ @@ -118,7 +126,7 @@ describe("Parallel web search", () => { it("surfaces plain-text Parallel API errors", async () => { using _hook = hookFetch(() => new Response("upstream unavailable", { status: 503 })); - await expect(searchParallel({ query: "broken" }, fakeStorage)).rejects.toMatchObject({ + await expect(searchParallel({ query: "broken" }, fakeAuthStorage)).rejects.toMatchObject({ provider: "parallel", status: 503, message: "Parallel API error (503): upstream unavailable", diff --git a/packages/coding-agent/test/tools/web-search-tavily.test.ts b/packages/coding-agent/test/tools/web-search-tavily.test.ts index 887502b9c..a260179c7 100644 --- a/packages/coding-agent/test/tools/web-search-tavily.test.ts +++ b/packages/coding-agent/test/tools/web-search-tavily.test.ts @@ -1,6 +1,6 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { hookFetch } from "@oh-my-pi/pi-utils"; -import { AgentStorage } from "../../src/session/agent-storage"; import { searchTavily } from "../../src/web/search/providers/tavily"; import type { SearchProviderError } from "../../src/web/search/types"; @@ -14,11 +14,22 @@ describe("Tavily web search provider", () => { delete process.env.TAVILY_API_KEY; }); - const fakeStorage = { - listAuthCredentials: () => [], - updateAuthCredential: () => undefined, - get authStore() { return null as never; }, - } as unknown as AgentStorage; + const fakeAuthStorage = { + async getApiKey() { + return process.env.TAVILY_API_KEY ?? undefined; + }, + hasAuth() { + return Boolean(process.env.TAVILY_API_KEY); + }, + } as unknown as AuthStorage; + + function makeParams(query: string) { + return { + query, + authStorage: fakeAuthStorage, + systemPrompt: "Tavily test prompt", + } as const; + } it("maps Tavily responses into SearchResponse and forwards recency filters", async () => { let requestBody: Record | null = null; @@ -46,7 +57,7 @@ describe("Tavily web search provider", () => { ); }); - const response = await searchTavily({ query: "latest ai news", num_results: 2, recency: "week" }, fakeStorage); + const response = await searchTavily({ ...makeParams("latest ai news"), numSearchResults: 2, recency: "week" }); // Recency must not couple to topic — topic should be absent (Tavily defaults to general) expect(requestBody).toMatchObject({ query: "latest ai news", @@ -87,7 +98,7 @@ describe("Tavily web search provider", () => { }), ); - await expect(searchTavily({ query: "bad auth" }, fakeStorage)).rejects.toEqual( + await expect(searchTavily(makeParams("bad auth"))).rejects.toEqual( expect.objectContaining({ provider: "tavily", status: 401, @@ -98,11 +109,8 @@ describe("Tavily web search provider", () => { it("throws a clear error when Tavily credentials are missing", async () => { delete process.env.TAVILY_API_KEY; - vi.spyOn(AgentStorage, "open").mockResolvedValue({ - listAuthCredentials: () => [], - } as unknown as AgentStorage); - await expect(searchTavily({ query: "missing creds" }, fakeStorage)).rejects.toThrow( - 'Tavily credentials not found. Set TAVILY_API_KEY or store an API key for provider "tavily" in agent.db.', + await expect(searchTavily(makeParams("missing creds"))).rejects.toThrow( + 'Tavily credentials not found. Set TAVILY_API_KEY or configure an API key for provider "tavily".', ); }); }); diff --git a/packages/coding-agent/test/web/search/abort-and-timeout.test.ts b/packages/coding-agent/test/web/search/abort-and-timeout.test.ts index 4c83da36d..c269628d6 100644 --- a/packages/coding-agent/test/web/search/abort-and-timeout.test.ts +++ b/packages/coding-agent/test/web/search/abort-and-timeout.test.ts @@ -13,6 +13,7 @@ */ import { afterEach, describe, expect, it, vi } from "bun:test"; import { hookFetch } from "@oh-my-pi/pi-utils"; +import type { AgentStorage } from "../../../src/session/agent-storage"; import type { ToolSession } from "../../../src/tools"; import { ToolAbortError } from "../../../src/tools/tool-errors"; import { WebSearchTool } from "../../../src/web/search"; @@ -21,14 +22,15 @@ import { searchAnthropic } from "../../../src/web/search/providers/anthropic"; import type { SearchParams } from "../../../src/web/search/providers/base"; import { searchBrave } from "../../../src/web/search/providers/brave"; import { withHardTimeout } from "../../../src/web/search/providers/utils"; -import type { AgentStorage } from "../../../src/session/agent-storage"; import type { SearchProviderId, SearchResponse } from "../../../src/web/search/types"; const FAKE_SESSION = {} as ToolSession; const fakeStorage = { listAuthCredentials: () => [], updateAuthCredential: () => undefined, - get authStore() { return null as never; }, + get authStore() { + return null as never; + }, } as unknown as AgentStorage; describe("withHardTimeout", () => { diff --git a/packages/coding-agent/test/web/search/codex-broker.test.ts b/packages/coding-agent/test/web/search/codex-broker.test.ts new file mode 100644 index 000000000..c59624740 --- /dev/null +++ b/packages/coding-agent/test/web/search/codex-broker.test.ts @@ -0,0 +1,62 @@ +import { describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; +import { hookFetch } from "@oh-my-pi/pi-utils"; +import { AgentStorage } from "../../../src/session/agent-storage"; +import type { SearchParams } from "../../../src/web/search/providers/base"; +import { searchCodex } from "../../../src/web/search/providers/codex"; + +function makeSseResponse(): string { + return [ + `data: ${JSON.stringify({ + type: "response.output_item.done", + item: { + type: "message", + content: [ + { + type: "output_text", + text: "Broker-backed Codex answer", + annotations: [{ type: "url_citation", url: "https://example.com/broker", title: "Broker" }], + }, + ], + }, + })}`, + "", + `data: ${JSON.stringify({ + type: "response.completed", + response: { id: "resp_codex_broker", model: "gpt-5-codex-mini" }, + })}`, + "", + ].join("\n"); +} + +describe("Codex web search broker auth", () => { + it("uses AuthStorage.getOAuthAccess for token + account metadata without opening AgentStorage", async () => { + const getOAuthAccess = vi.fn(async () => ({ + accessToken: "broker-refreshed-access-token", + accountId: "broker-account-id", + })); + const authStorage = { getOAuthAccess } as unknown as AuthStorage; + const openSpy = vi.spyOn(AgentStorage, "open"); + let requestHeaders: Headers | undefined; + + using _hook = hookFetch(async (_url, init) => { + requestHeaders = new Headers(init?.headers); + return new Response(makeSseResponse(), { status: 200, headers: { "Content-Type": "text/event-stream" } }); + }); + + const params: SearchParams = { + query: "broker codex search", + systemPrompt: "Use web search.", + authStorage, + sessionId: "codex-broker-session", + }; + + const result = await searchCodex(params); + + expect(result.provider).toBe("codex"); + expect(getOAuthAccess).toHaveBeenCalledWith("openai-codex", "codex-broker-session", { signal: undefined }); + expect(requestHeaders?.get("authorization")).toBe("Bearer broker-refreshed-access-token"); + expect(requestHeaders?.get("chatgpt-account-id")).toBe("broker-account-id"); + expect(openSpy).not.toHaveBeenCalled(); + }); +}); diff --git a/packages/coding-agent/test/web/search/tavily.test.ts b/packages/coding-agent/test/web/search/tavily.test.ts index 8585ad169..689a0a495 100644 --- a/packages/coding-agent/test/web/search/tavily.test.ts +++ b/packages/coding-agent/test/web/search/tavily.test.ts @@ -1,11 +1,11 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { AuthStorage } from "@oh-my-pi/pi-ai"; import { buildRequestBody, searchTavily, type TavilySearchParams, } from "@oh-my-pi/pi-coding-agent/web/search/providers/tavily"; import { hookFetch } from "@oh-my-pi/pi-utils"; -import type { AgentStorage } from "../../../src/session/agent-storage"; describe("Tavily buildRequestBody", () => { afterEach(() => { @@ -53,11 +53,23 @@ describe("Tavily searchTavily request shape (integration)", () => { delete process.env.TAVILY_API_KEY; }); - const fakeStorage = { - listAuthCredentials: () => [], - updateAuthCredential: () => undefined, - get authStore() { return null as never; }, - } as unknown as AgentStorage; + const fakeAuthStorage = { + async getApiKey() { + return process.env.TAVILY_API_KEY ?? undefined; + }, + hasAuth() { + return Boolean(process.env.TAVILY_API_KEY); + }, + } as unknown as AuthStorage; + + function makeParams(query: string, extras: Partial = {}) { + return { + query, + authStorage: fakeAuthStorage, + systemPrompt: "Tavily integration test prompt", + ...extras, + }; + } it("does not send topic=news to the upstream API when recency is set", async () => { process.env.TAVILY_API_KEY = "test-key"; @@ -86,20 +98,13 @@ describe("Tavily searchTavily request shape (integration)", () => { return new Response("not mocked", { status: 500 }); }); - const params: TavilySearchParams = { - query: "Bun runtime latest release notes", - recency: "week", - }; - const response = await searchTavily(params, fakeStorage); + const response = await searchTavily(makeParams("Bun runtime latest release notes", { recency: "week" })); expect(capturedBody).toBeDefined(); - // The core regression: recency must not coerce topic to news. Topic should - // be absent entirely (Tavily defaults to "general"). expect(capturedBody).not.toHaveProperty("topic"); expect(capturedBody?.time_range).toBe("week"); expect(capturedBody?.query).toBe("Bun runtime latest release notes"); - // And the response should still be parsed correctly end-to-end. expect(response.provider).toBe("tavily"); expect(response.answer).toBe("test answer"); expect(response.sources).toHaveLength(1); @@ -122,10 +127,10 @@ describe("Tavily searchTavily request shape (integration)", () => { return new Response("not mocked", { status: 500 }); }); - await searchTavily({ query: "bun sqlite" }, fakeStorage); + await searchTavily(makeParams("bun sqlite")); expect(capturedBody).toBeDefined(); expect(capturedBody).not.toHaveProperty("topic"); expect(capturedBody).not.toHaveProperty("time_range"); }); -}); \ No newline at end of file +});