Files
oh-my-pi/packages/coding-agent/src/session/auth-storage.ts
T
2026-02-19 13:51:47 +01:00

1401 lines
45 KiB
TypeScript

/**
* Credential storage for API keys and OAuth tokens.
* Handles loading, saving, and refreshing credentials from agent.db.
*/
import {
antigravityUsageProvider,
claudeUsageProvider,
getEnvApiKey,
getOAuthApiKey,
getOAuthProvider,
githubCopilotUsageProvider,
googleGeminiCliUsageProvider,
kimiUsageProvider,
loginAnthropic,
loginAntigravity,
loginCerebras,
loginCloudflareAiGateway,
loginCursor,
loginGeminiCli,
loginGitHubCopilot,
loginHuggingface,
loginKimi,
loginLiteLLM,
loginMiniMaxCode,
loginMiniMaxCodeCn,
loginMoonshot,
loginNanoGPT,
loginNvidia,
loginOllama,
loginOpenAICodex,
loginOpenCode,
loginPerplexity,
loginQianfan,
loginQwenPortal,
loginSynthetic,
loginTogether,
loginVenice,
loginVllm,
loginXiaomi,
loginZai,
type OAuthController,
type OAuthCredentials,
type OAuthProvider,
type OAuthProviderId,
openaiCodexUsageProvider,
type Provider,
type UsageCache,
type UsageCacheEntry,
type UsageCredential,
type UsageLimit,
type UsageLogger,
type UsageProvider,
type UsageReport,
zaiUsageProvider,
} from "@oh-my-pi/pi-ai";
import { logger } from "@oh-my-pi/pi-utils";
import { resolveConfigValue } from "../config/resolve-config-value";
import { AgentStorage } from "./agent-storage";
export type ApiKeyCredential = {
type: "api_key";
key: string;
};
export type OAuthCredential = {
type: "oauth";
} & OAuthCredentials;
export type AuthCredential = ApiKeyCredential | OAuthCredential;
export type AuthCredentialEntry = AuthCredential | AuthCredential[];
export type AuthStorageData = Record<string, AuthCredentialEntry>;
/**
* Serialized representation of AuthStorage for passing to subagent workers.
* Contains only the essential credential data, not runtime state.
*/
export interface SerializedAuthStorage {
credentials: Record<
string,
Array<{
id: number;
type: "api_key" | "oauth";
data: Record<string, unknown>;
}>
>;
runtimeOverrides?: Record<string, string>;
dbPath?: string;
}
/**
* In-memory representation pairing DB row ID with credential.
* The ID is required for update/delete operations against agent.db.
*/
type StoredCredential = { id: number; credential: AuthCredential };
export type AuthStorageOptions = {
usageProviderResolver?: (provider: Provider) => UsageProvider | undefined;
usageCache?: UsageCache;
usageFetch?: typeof fetch;
usageNow?: () => number;
usageLogger?: UsageLogger;
};
const DEFAULT_USAGE_PROVIDERS: UsageProvider[] = [
openaiCodexUsageProvider,
kimiUsageProvider,
antigravityUsageProvider,
googleGeminiCliUsageProvider,
claudeUsageProvider,
zaiUsageProvider,
githubCopilotUsageProvider,
];
const DEFAULT_USAGE_PROVIDER_MAP = new Map<Provider, UsageProvider>(
DEFAULT_USAGE_PROVIDERS.map(provider => [provider.id, provider]),
);
const USAGE_CACHE_PREFIX = "usage_cache:";
function resolveDefaultUsageProvider(provider: Provider): UsageProvider | undefined {
return DEFAULT_USAGE_PROVIDER_MAP.get(provider);
}
function parseUsageCacheEntry(raw: string): UsageCacheEntry | undefined {
try {
const parsed = JSON.parse(raw) as { value?: UsageReport | null; expiresAt?: unknown };
const expiresAt = typeof parsed.expiresAt === "number" ? parsed.expiresAt : undefined;
if (!expiresAt || !Number.isFinite(expiresAt)) return undefined;
return { value: parsed.value ?? null, expiresAt };
} catch {
return undefined;
}
}
class AuthStorageUsageCache implements UsageCache {
constructor(private storage: AgentStorage) {}
get(key: string): UsageCacheEntry | undefined {
const raw = this.storage.getCache(`${USAGE_CACHE_PREFIX}${key}`);
if (!raw) return undefined;
const entry = parseUsageCacheEntry(raw);
if (!entry) return undefined;
if (entry.expiresAt <= Date.now()) return undefined;
return entry;
}
set(key: string, entry: UsageCacheEntry): void {
const payload = JSON.stringify({ value: entry.value ?? null, expiresAt: entry.expiresAt });
this.storage.setCache(`${USAGE_CACHE_PREFIX}${key}`, payload, Math.floor(entry.expiresAt / 1000));
}
cleanup(): void {
this.storage.cleanExpiredCache();
}
}
/**
* Credential storage backed by agent.db.
* Reads from SQLite (agent.db).
*/
export class AuthStorage {
static readonly #defaultBackoffMs = 60_000; // Default backoff when no reset time available
/** Provider -> credentials cache, populated from agent.db on reload(). */
#data: Map<string, StoredCredential[]> = new Map();
#runtimeOverrides: Map<string, string> = new Map();
/** Tracks next credential index per provider:type key for round-robin distribution (non-session use). */
#providerRoundRobinIndex: Map<string, number> = new Map();
/** Tracks the last used credential per provider for a session (used for rate-limit switching). */
#sessionLastCredential: Map<string, Map<string, { type: AuthCredential["type"]; index: number }>> = new Map();
/** Maps provider:type -> credentialIndex -> blockedUntilMs for temporary backoff. */
#credentialBackoff: Map<string, Map<number, number>> = new Map();
#usageProviderResolver?: (provider: Provider) => UsageProvider | undefined;
#usageCache?: UsageCache;
#usageFetch: typeof fetch;
#usageNow: () => number;
#usageLogger?: UsageLogger;
#fallbackResolver?: (provider: string) => string | undefined;
private constructor(
private storage: AgentStorage,
options: AuthStorageOptions = {},
) {
this.#usageProviderResolver = options.usageProviderResolver ?? resolveDefaultUsageProvider;
this.#usageCache = options.usageCache ?? new AuthStorageUsageCache(this.storage);
this.#usageFetch = options.usageFetch ?? fetch;
this.#usageNow = options.usageNow ?? Date.now;
this.#usageLogger =
options.usageLogger ??
({
debug: (message, meta) => logger.debug(message, meta),
warn: (message, meta) => logger.warn(message, meta),
} satisfies UsageLogger);
}
/**
* Create an AuthStorage instance.
* @param dbPath - Path to agent.db
*/
static async create(dbPath: string, options: AuthStorageOptions = {}): Promise<AuthStorage> {
const storage = await AgentStorage.open(dbPath);
return new AuthStorage(storage, options);
}
/**
* Set a runtime API key override (not persisted to disk).
* Used for CLI --api-key flag.
*/
setRuntimeApiKey(provider: string, apiKey: string): void {
this.#runtimeOverrides.set(provider, apiKey);
}
/**
* Remove a runtime API key override.
*/
removeRuntimeApiKey(provider: string): void {
this.#runtimeOverrides.delete(provider);
}
/**
* Set a fallback resolver for API keys not found in agent.db or env vars.
* Used for custom provider keys from models.json.
*/
setFallbackResolver(resolver: (provider: string) => string | undefined): void {
this.#fallbackResolver = resolver;
}
/**
* Reload credentials from agent.db.
* Reloads credentials from the database.
*/
async reload(): Promise<void> {
const records = this.storage.listAuthCredentials();
const grouped = new Map<string, StoredCredential[]>();
for (const record of records) {
const list = grouped.get(record.provider) ?? [];
list.push({ id: record.id, credential: record.credential });
grouped.set(record.provider, list);
}
const dedupedGrouped = new Map<string, StoredCredential[]>();
for (const [provider, entries] of grouped.entries()) {
const deduped = this.#pruneDuplicateStoredCredentials(provider, entries);
if (deduped.length > 0) {
dedupedGrouped.set(provider, deduped);
}
}
this.#data = dedupedGrouped;
}
/**
* Gets cached credentials for a provider.
* @param provider - Provider name (e.g., "anthropic", "openai")
* @returns Array of stored credentials, empty if none exist
*/
#getStoredCredentials(provider: string): StoredCredential[] {
return this.#data.get(provider) ?? [];
}
/**
* Updates in-memory credential cache for a provider.
* Removes the provider entry entirely if credentials array is empty.
* @param provider - Provider name (e.g., "anthropic", "openai")
* @param credentials - Array of stored credentials to cache
*/
#setStoredCredentials(provider: string, credentials: StoredCredential[]): void {
if (credentials.length === 0) {
this.#data.delete(provider);
} else {
this.#data.set(provider, credentials);
}
}
#getOAuthIdentifiers(credential: OAuthCredential): string[] {
const identifiers: string[] = [];
const accountId = credential.accountId?.trim();
if (accountId) identifiers.push(`account:${accountId}`);
const email = credential.email?.trim().toLowerCase();
if (email) identifiers.push(`email:${email}`);
if (identifiers.length > 0) return identifiers;
const tokenIdentifiers = this.#getOAuthIdentifiersFromToken(credential.access) ?? [];
for (const identifier of tokenIdentifiers) {
identifiers.push(identifier);
}
if (identifiers.length > 0) return identifiers;
const refreshIdentifiers = this.#getOAuthIdentifiersFromToken(credential.refresh) ?? [];
for (const identifier of refreshIdentifiers) {
identifiers.push(identifier);
}
return identifiers;
}
#getOAuthIdentifiersFromToken(token: string | undefined): string[] | undefined {
if (!token) return undefined;
const parts = token.split(".");
if (parts.length !== 3) return undefined;
const payloadRaw = parts[1];
const decoder = new TextDecoder("utf-8");
try {
const payload = JSON.parse(
decoder.decode(Uint8Array.fromBase64(payloadRaw, { alphabet: "base64url" })),
) as Record<string, unknown>;
if (!payload || typeof payload !== "object") return undefined;
const identifiers: string[] = [];
const email = typeof payload.email === "string" ? payload.email.trim().toLowerCase() : undefined;
if (email) identifiers.push(`email:${email}`);
const accountId =
typeof payload.account_id === "string"
? payload.account_id
: typeof payload.accountId === "string"
? payload.accountId
: typeof payload.user_id === "string"
? payload.user_id
: typeof payload.sub === "string"
? payload.sub
: undefined;
const trimmedAccountId = accountId?.trim();
if (trimmedAccountId) identifiers.push(`account:${trimmedAccountId}`);
return identifiers.length > 0 ? identifiers : undefined;
} catch {
return undefined;
}
}
#dedupeOAuthCredentials(credentials: AuthCredential[]): AuthCredential[] {
const seen = new Set<string>();
const deduped: AuthCredential[] = [];
for (let index = credentials.length - 1; index >= 0; index -= 1) {
const credential = credentials[index];
if (credential.type !== "oauth") {
deduped.push(credential);
continue;
}
const identifiers = this.#getOAuthIdentifiers(credential);
if (identifiers.length === 0) {
deduped.push(credential);
continue;
}
if (identifiers.some(identifier => seen.has(identifier))) {
continue;
}
for (const identifier of identifiers) {
seen.add(identifier);
}
deduped.push(credential);
}
return deduped.reverse();
}
#pruneDuplicateStoredCredentials(provider: string, entries: StoredCredential[]): StoredCredential[] {
const seen = new Set<string>();
const kept: StoredCredential[] = [];
const removed: StoredCredential[] = [];
for (let index = entries.length - 1; index >= 0; index -= 1) {
const entry = entries[index];
const credential = entry.credential;
if (credential.type !== "oauth") {
kept.push(entry);
continue;
}
const identifiers = this.#getOAuthIdentifiers(credential);
if (identifiers.length === 0) {
kept.push(entry);
continue;
}
if (identifiers.some(identifier => seen.has(identifier))) {
removed.push(entry);
continue;
}
for (const identifier of identifiers) {
seen.add(identifier);
}
kept.push(entry);
}
if (removed.length > 0) {
for (const entry of removed) {
this.storage.deleteAuthCredential(entry.id);
}
this.#resetProviderAssignments(provider);
}
return kept.reverse();
}
/** Returns all credentials for a provider as an array */
#getCredentialsForProvider(provider: string): AuthCredential[] {
return this.#getStoredCredentials(provider).map(entry => entry.credential);
}
/** Composite key for round-robin tracking: "anthropic:oauth" or "openai:api_key" */
#getProviderTypeKey(provider: string, type: AuthCredential["type"]): string {
return `${provider}:${type}`;
}
/**
* Returns next index in round-robin sequence for load distribution.
* Increments stored counter and wraps at total.
*/
#getNextRoundRobinIndex(providerKey: string, total: number): number {
if (total <= 1) return 0;
const current = this.#providerRoundRobinIndex.get(providerKey) ?? -1;
const next = (current + 1) % total;
this.#providerRoundRobinIndex.set(providerKey, next);
return next;
}
/**
* FNV-1a hash for deterministic session-to-credential mapping.
* Ensures the same session always starts with the same credential.
*/
#getHashedIndex(sessionId: string, total: number): number {
if (total <= 1) return 0;
let hash = 2166136261; // FNV offset basis
for (let i = 0; i < sessionId.length; i++) {
hash ^= sessionId.charCodeAt(i);
hash = Math.imul(hash, 16777619); // FNV prime
}
return (hash >>> 0) % total;
}
/**
* Returns credential indices in priority order for selection.
* With sessionId: starts from hashed index (consistent per session).
* Without sessionId: starts from round-robin index (load balancing).
* Order wraps around so all credentials are tried if earlier ones are blocked.
*/
#getCredentialOrder(providerKey: string, sessionId: string | undefined, total: number): number[] {
if (total <= 1) return [0];
const start = sessionId
? this.#getHashedIndex(sessionId, total)
: this.#getNextRoundRobinIndex(providerKey, total);
const order: number[] = [];
for (let i = 0; i < total; i++) {
order.push((start + i) % total);
}
return order;
}
/** Checks if a credential is temporarily blocked due to usage limits. */
#isCredentialBlocked(providerKey: string, credentialIndex: number): boolean {
const backoffMap = this.#credentialBackoff.get(providerKey);
if (!backoffMap) return false;
const blockedUntil = backoffMap.get(credentialIndex);
if (!blockedUntil) return false;
if (blockedUntil <= Date.now()) {
backoffMap.delete(credentialIndex);
if (backoffMap.size === 0) {
this.#credentialBackoff.delete(providerKey);
}
return false;
}
return true;
}
/** Marks a credential as blocked until the specified time. */
#markCredentialBlocked(providerKey: string, credentialIndex: number, blockedUntilMs: number): void {
const backoffMap = this.#credentialBackoff.get(providerKey) ?? new Map<number, number>();
const existing = backoffMap.get(credentialIndex) ?? 0;
backoffMap.set(credentialIndex, Math.max(existing, blockedUntilMs));
this.#credentialBackoff.set(providerKey, backoffMap);
}
/** Records which credential was used for a session (for rate-limit switching). */
#recordSessionCredential(
provider: string,
sessionId: string | undefined,
type: AuthCredential["type"],
index: number,
): void {
if (!sessionId) return;
const sessionMap = this.#sessionLastCredential.get(provider) ?? new Map();
sessionMap.set(sessionId, { type, index });
this.#sessionLastCredential.set(provider, sessionMap);
}
/** Retrieves the last credential used by a session. */
#getSessionCredential(
provider: string,
sessionId: string | undefined,
): { type: AuthCredential["type"]; index: number } | undefined {
if (!sessionId) return undefined;
return this.#sessionLastCredential.get(provider)?.get(sessionId);
}
/**
* Selects a credential of the specified type for a provider.
* Returns both the credential and its index in the original array (for updates/removal).
* Uses deterministic hashing for session stickiness and skips blocked credentials when possible.
*/
#selectCredentialByType<T extends AuthCredential["type"]>(
provider: string,
type: T,
sessionId?: string,
): { credential: Extract<AuthCredential, { type: T }>; index: number } | undefined {
const credentials = this.#getCredentialsForProvider(provider)
.map((credential, index) => ({ credential, index }))
.filter(
(entry): entry is { credential: Extract<AuthCredential, { type: T }>; index: number } =>
entry.credential.type === type,
);
if (credentials.length === 0) return undefined;
if (credentials.length === 1) return credentials[0];
const providerKey = this.#getProviderTypeKey(provider, type);
const order = this.#getCredentialOrder(providerKey, sessionId, credentials.length);
const fallback = credentials[order[0]];
for (const idx of order) {
const candidate = credentials[idx];
if (!this.#isCredentialBlocked(providerKey, candidate.index)) {
return candidate;
}
}
return fallback;
}
/**
* Clears round-robin and session assignment state for a provider.
* Called when credentials are added/removed to prevent stale index references.
*/
#resetProviderAssignments(provider: string): void {
for (const key of this.#providerRoundRobinIndex.keys()) {
if (key.startsWith(`${provider}:`)) {
this.#providerRoundRobinIndex.delete(key);
}
}
this.#sessionLastCredential.delete(provider);
for (const key of this.#credentialBackoff.keys()) {
if (key.startsWith(`${provider}:`)) {
this.#credentialBackoff.delete(key);
}
}
}
/** Updates credential at index in-place (used for OAuth token refresh) */
#replaceCredentialAt(provider: string, index: number, credential: AuthCredential): void {
const entries = this.#getStoredCredentials(provider);
if (index < 0 || index >= entries.length) return;
const target = entries[index];
this.storage.updateAuthCredential(target.id, credential);
const updated = [...entries];
updated[index] = { id: target.id, credential };
this.#setStoredCredentials(provider, updated);
}
/**
* Removes credential at index (used when OAuth refresh fails).
* Cleans up provider entry if last credential removed.
*/
#removeCredentialAt(provider: string, index: number): void {
const entries = this.#getStoredCredentials(provider);
if (index < 0 || index >= entries.length) return;
this.storage.deleteAuthCredential(entries[index].id);
const updated = entries.filter((_value, idx) => idx !== index);
this.#setStoredCredentials(provider, updated);
this.#resetProviderAssignments(provider);
}
/**
* Get credential for a provider (first entry if multiple).
*/
get(provider: string): AuthCredential | undefined {
return this.#getCredentialsForProvider(provider)[0];
}
/**
* Set credential for a provider.
*/
async set(provider: string, credential: AuthCredentialEntry): Promise<void> {
const normalized = Array.isArray(credential) ? credential : [credential];
const deduped = this.#dedupeOAuthCredentials(normalized);
const stored = this.storage.replaceAuthCredentialsForProvider(provider, deduped);
this.#setStoredCredentials(
provider,
stored.map(record => ({ id: record.id, credential: record.credential })),
);
this.#resetProviderAssignments(provider);
}
/**
* Remove credential for a provider.
*/
async remove(provider: string): Promise<void> {
this.storage.deleteAuthCredentialsForProvider(provider);
this.#data.delete(provider);
this.#resetProviderAssignments(provider);
}
/**
* List all providers with credentials.
*/
list(): string[] {
return [...this.#data.keys()];
}
/**
* Check if credentials exist for a provider in agent.db.
*/
has(provider: string): boolean {
return this.#getCredentialsForProvider(provider).length > 0;
}
/**
* Check if any form of auth is configured for a provider.
* Unlike getApiKey(), this doesn't refresh OAuth tokens.
*/
hasAuth(provider: string): boolean {
if (this.#runtimeOverrides.has(provider)) return true;
if (this.#getCredentialsForProvider(provider).length > 0) return true;
if (getEnvApiKey(provider)) return true;
if (this.#fallbackResolver?.(provider)) return true;
return false;
}
/**
* Check if OAuth credentials are configured for a provider.
*/
hasOAuth(provider: string): boolean {
return this.#getCredentialsForProvider(provider).some(credential => credential.type === "oauth");
}
/**
* Get OAuth credentials for a provider.
*/
getOAuthCredential(provider: string): OAuthCredential | undefined {
return this.#getCredentialsForProvider(provider).find(
(credential): credential is OAuthCredential => credential.type === "oauth",
);
}
/**
* Get all credentials.
*/
getAll(): AuthStorageData {
const result: AuthStorageData = {};
for (const [provider, entries] of this.#data.entries()) {
const credentials = entries.map(entry => entry.credential);
if (credentials.length === 1) {
result[provider] = credentials[0];
} else if (credentials.length > 1) {
result[provider] = credentials;
}
}
return result;
}
/**
* Login to an OAuth provider.
*/
async login(
provider: OAuthProviderId,
ctrl: OAuthController & {
/** onAuth is required by auth-storage but optional in OAuthController */
onAuth: (info: { url: string; instructions?: string }) => void;
/** onPrompt is required for some providers (github-copilot, openai-codex) */
onPrompt: (prompt: { message: string; placeholder?: string }) => Promise<string>;
},
): Promise<void> {
let credentials: OAuthCredentials;
const saveApiKeyCredential = async (apiKey: string): Promise<void> => {
const newCredential: ApiKeyCredential = { type: "api_key", key: apiKey };
const existing = this.#getCredentialsForProvider(provider);
if (existing.length === 0) {
await this.set(provider, newCredential);
return;
}
await this.set(provider, [...existing, newCredential]);
};
switch (provider) {
case "anthropic":
credentials = await loginAnthropic({
...ctrl,
onManualCodeInput: async () =>
ctrl.onPrompt({ message: "Paste the authorization code (or full redirect URL):" }),
});
break;
case "github-copilot":
credentials = await loginGitHubCopilot({
onAuth: (url, instructions) => ctrl.onAuth({ url, instructions }),
onPrompt: ctrl.onPrompt,
onProgress: ctrl.onProgress,
signal: ctrl.signal,
});
break;
case "google-gemini-cli":
credentials = await loginGeminiCli(ctrl);
break;
case "google-antigravity":
credentials = await loginAntigravity(ctrl);
break;
case "openai-codex":
credentials = await loginOpenAICodex(ctrl);
break;
case "kimi-code":
credentials = await loginKimi(ctrl);
break;
case "cursor":
credentials = await loginCursor(
url => ctrl.onAuth({ url }),
ctrl.onProgress ? () => ctrl.onProgress?.("Waiting for browser authentication...") : undefined,
);
break;
case "perplexity":
credentials = await loginPerplexity(ctrl);
break;
case "huggingface": {
const apiKey = await loginHuggingface(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "opencode": {
const apiKey = await loginOpenCode(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "ollama": {
const apiKey = await loginOllama(ctrl);
if (!apiKey) {
return;
}
await saveApiKeyCredential(apiKey);
return;
}
case "cerebras": {
const apiKey = await loginCerebras(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "zai": {
const apiKey = await loginZai(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "qianfan": {
const apiKey = await loginQianfan(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "minimax-code": {
const apiKey = await loginMiniMaxCode(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "minimax-code-cn": {
const apiKey = await loginMiniMaxCodeCn(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "synthetic": {
const apiKey = await loginSynthetic(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "venice": {
const apiKey = await loginVenice(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "litellm": {
const apiKey = await loginLiteLLM(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "moonshot": {
const apiKey = await loginMoonshot(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "nanogpt": {
const apiKey = await loginNanoGPT(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "together": {
const apiKey = await loginTogether(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "cloudflare-ai-gateway": {
const apiKey = await loginCloudflareAiGateway(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "vllm": {
const apiKey = await loginVllm(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "qwen-portal": {
const apiKey = await loginQwenPortal(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "nvidia": {
const apiKey = await loginNvidia(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
case "xiaomi": {
const apiKey = await loginXiaomi(ctrl);
await saveApiKeyCredential(apiKey);
return;
}
default: {
const customProvider = getOAuthProvider(provider);
if (!customProvider) {
throw new Error(`Unknown OAuth provider: ${provider}`);
}
const customLoginResult = await customProvider.login({
onAuth: info => ctrl.onAuth(info),
onProgress: ctrl.onProgress,
onPrompt: ctrl.onPrompt,
onManualCodeInput: async () =>
ctrl.onPrompt({ message: "Paste the authorization code (or full redirect URL):" }),
signal: ctrl.signal,
});
if (typeof customLoginResult === "string") {
await saveApiKeyCredential(customLoginResult);
return;
}
credentials = customLoginResult;
break;
}
}
const newCredential: OAuthCredential = { type: "oauth", ...credentials };
const existing = this.#getCredentialsForProvider(provider);
if (existing.length === 0) {
await this.set(provider, newCredential);
return;
}
await this.set(provider, [...existing, newCredential]);
}
/**
* Logout from a provider.
*/
async logout(provider: string): Promise<void> {
await this.remove(provider);
}
// ─────────────────────────────────────────────────────────────────────────────
// Usage API Integration
// Queries provider usage endpoints to detect rate limits before they occur.
// ─────────────────────────────────────────────────────────────────────────────
#buildUsageCredential(credential: OAuthCredential): UsageCredential {
return {
type: "oauth",
accessToken: credential.access,
refreshToken: credential.refresh,
expiresAt: credential.expires,
accountId: credential.accountId,
projectId: credential.projectId,
email: credential.email,
enterpriseUrl: credential.enterpriseUrl,
};
}
#getUsageReportMetadataValue(report: UsageReport, key: string): string | undefined {
const metadata = report.metadata;
if (!metadata || typeof metadata !== "object") return undefined;
const value = metadata[key];
return typeof value === "string" ? value.trim() : undefined;
}
#getUsageReportScopeAccountId(report: UsageReport): string | undefined {
const ids = new Set<string>();
for (const limit of report.limits) {
const accountId = limit.scope.accountId?.trim();
if (accountId) ids.add(accountId);
}
if (ids.size === 1) return [...ids][0];
return undefined;
}
#getUsageReportIdentifiers(report: UsageReport): string[] {
const identifiers: string[] = [];
const email = this.#getUsageReportMetadataValue(report, "email");
if (email) identifiers.push(`email:${email.toLowerCase()}`);
const accountId = this.#getUsageReportMetadataValue(report, "accountId");
if (accountId) identifiers.push(`account:${accountId}`);
const account = this.#getUsageReportMetadataValue(report, "account");
if (account) identifiers.push(`account:${account}`);
const user = this.#getUsageReportMetadataValue(report, "user");
if (user) identifiers.push(`account:${user}`);
const username = this.#getUsageReportMetadataValue(report, "username");
if (username) identifiers.push(`account:${username}`);
const scopeAccountId = this.#getUsageReportScopeAccountId(report);
if (scopeAccountId) identifiers.push(`account:${scopeAccountId}`);
return identifiers.map(identifier => `${report.provider}:${identifier.toLowerCase()}`);
}
#mergeUsageReportGroup(reports: UsageReport[]): UsageReport {
if (reports.length === 1) return reports[0];
const sorted = [...reports].sort((a, b) => {
const limitDiff = b.limits.length - a.limits.length;
if (limitDiff !== 0) return limitDiff;
return (b.fetchedAt ?? 0) - (a.fetchedAt ?? 0);
});
const base = sorted[0];
const mergedLimits = [...base.limits];
const limitIds = new Set(mergedLimits.map(limit => limit.id));
const mergedMetadata: Record<string, unknown> = { ...(base.metadata ?? {}) };
let fetchedAt = base.fetchedAt;
for (const report of sorted.slice(1)) {
fetchedAt = Math.max(fetchedAt, report.fetchedAt);
for (const limit of report.limits) {
if (!limitIds.has(limit.id)) {
limitIds.add(limit.id);
mergedLimits.push(limit);
}
}
if (report.metadata) {
for (const [key, value] of Object.entries(report.metadata)) {
if (mergedMetadata[key] === undefined) {
mergedMetadata[key] = value;
}
}
}
}
return {
...base,
fetchedAt,
limits: mergedLimits,
metadata: Object.keys(mergedMetadata).length > 0 ? mergedMetadata : undefined,
};
}
#dedupeUsageReports(reports: UsageReport[]): UsageReport[] {
const groups: UsageReport[][] = [];
const idToGroup = new Map<string, number>();
for (const report of reports) {
const identifiers = this.#getUsageReportIdentifiers(report);
let groupIndex: number | undefined;
for (const identifier of identifiers) {
const existing = idToGroup.get(identifier);
if (existing !== undefined) {
groupIndex = existing;
break;
}
}
if (groupIndex === undefined) {
groupIndex = groups.length;
groups.push([]);
}
groups[groupIndex].push(report);
for (const identifier of identifiers) {
idToGroup.set(identifier, groupIndex);
}
}
const deduped = groups.map(group => this.#mergeUsageReportGroup(group));
if (deduped.length !== reports.length) {
this.#usageLogger?.debug("Usage reports deduped", {
before: reports.length,
after: deduped.length,
});
}
return deduped;
}
#isUsageLimitExhausted(limit: UsageLimit): boolean {
if (limit.status === "exhausted") return true;
const amount = limit.amount;
if (amount.usedFraction !== undefined && amount.usedFraction >= 1) return true;
if (amount.remainingFraction !== undefined && amount.remainingFraction <= 0) return true;
if (amount.used !== undefined && amount.limit !== undefined && amount.used >= amount.limit) return true;
if (amount.remaining !== undefined && amount.remaining <= 0) return true;
if (amount.unit === "percent" && amount.used !== undefined && amount.used >= 100) return true;
return false;
}
/** Returns true if usage indicates rate limit has been reached. */
#isUsageLimitReached(report: UsageReport): boolean {
return report.limits.some(limit => this.#isUsageLimitExhausted(limit));
}
/** Extracts the earliest reset timestamp from exhausted windows (in ms). */
#getUsageResetAtMs(report: UsageReport, nowMs: number): number | undefined {
const candidates: number[] = [];
for (const limit of report.limits) {
if (!this.#isUsageLimitExhausted(limit)) continue;
const window = limit.window;
if (window?.resetsAt && window.resetsAt > nowMs) {
candidates.push(window.resetsAt);
}
if (window?.resetInMs && window.resetInMs > 0) {
const resetAt = nowMs + window.resetInMs;
if (resetAt > nowMs) candidates.push(resetAt);
}
}
if (candidates.length === 0) return undefined;
return Math.min(...candidates);
}
async #getUsageReport(
provider: Provider,
credential: OAuthCredential,
options?: { baseUrl?: string },
): Promise<UsageReport | null> {
const resolver = this.#usageProviderResolver;
const cache = this.#usageCache;
if (!resolver || !cache) return null;
const providerImpl = resolver(provider);
if (!providerImpl) return null;
const params = {
provider,
credential: this.#buildUsageCredential(credential),
baseUrl: options?.baseUrl,
};
if (providerImpl.supports && !providerImpl.supports(params)) return null;
try {
return await providerImpl.fetchUsage(params, {
cache,
fetch: this.#usageFetch,
now: this.#usageNow,
logger: this.#usageLogger,
});
} catch (error) {
logger.debug("AuthStorage usage fetch failed", {
provider,
error: String(error),
});
return null;
}
}
async fetchUsageReports(options?: {
baseUrlResolver?: (provider: Provider) => string | undefined;
}): Promise<UsageReport[] | null> {
const resolver = this.#usageProviderResolver;
const cache = this.#usageCache;
if (!resolver || !cache) return null;
const tasks: Array<Promise<UsageReport | null>> = [];
const providers = new Set<string>([
...this.#data.keys(),
...DEFAULT_USAGE_PROVIDERS.map(provider => provider.id),
]);
this.#usageLogger?.debug("Usage fetch requested", {
providers: Array.from(providers).sort(),
});
for (const provider of providers) {
const providerImpl = resolver(provider as Provider);
if (!providerImpl) continue;
const baseUrl = options?.baseUrlResolver?.(provider as Provider);
let entries = this.#getStoredCredentials(provider);
if (entries.length > 0) {
const dedupedEntries = this.#pruneDuplicateStoredCredentials(provider, entries);
if (dedupedEntries.length !== entries.length) {
this.#setStoredCredentials(provider, dedupedEntries);
}
entries = dedupedEntries;
}
if (entries.length === 0) {
const runtimeKey = this.#runtimeOverrides.get(provider);
const envKey = getEnvApiKey(provider);
const apiKey = runtimeKey ?? envKey;
if (!apiKey) {
continue;
}
const params = {
provider: provider as Provider,
credential: { type: "api_key", apiKey } satisfies UsageCredential,
baseUrl,
};
if (providerImpl.supports && !providerImpl.supports(params)) {
continue;
}
this.#usageLogger?.debug("Usage fetch queued", {
provider,
credentialType: "api_key",
baseUrl,
});
tasks.push(
providerImpl
.fetchUsage(params, {
cache,
fetch: this.#usageFetch,
now: this.#usageNow,
logger: this.#usageLogger,
})
.catch(error => {
logger.debug("AuthStorage usage fetch failed", {
provider,
error: String(error),
});
return null;
}),
);
continue;
}
for (const entry of entries) {
const credential = entry.credential;
const usageCredential: UsageCredential =
credential.type === "api_key"
? { type: "api_key", apiKey: credential.key }
: this.#buildUsageCredential(credential);
const params = {
provider: provider as Provider,
credential: usageCredential,
baseUrl,
};
if (providerImpl.supports && !providerImpl.supports(params)) {
continue;
}
this.#usageLogger?.debug("Usage fetch queued", {
provider,
credentialType: usageCredential.type,
baseUrl,
accountId: usageCredential.accountId,
email: usageCredential.email,
});
tasks.push(
providerImpl
.fetchUsage(params, {
cache,
fetch: this.#usageFetch,
now: this.#usageNow,
logger: this.#usageLogger,
})
.catch(error => {
logger.debug("AuthStorage usage fetch failed", {
provider,
error: String(error),
});
return null;
}),
);
}
}
if (tasks.length === 0) return [];
const results = await Promise.all(tasks);
const reports = results.filter((report): report is UsageReport => report !== null);
const deduped = this.#dedupeUsageReports(reports);
this.#usageLogger?.debug("Usage fetch resolved", {
reports: deduped.map(report => {
const accountLabel =
this.#getUsageReportMetadataValue(report, "email") ??
this.#getUsageReportMetadataValue(report, "accountId") ??
this.#getUsageReportMetadataValue(report, "account") ??
this.#getUsageReportMetadataValue(report, "user") ??
this.#getUsageReportMetadataValue(report, "username") ??
this.#getUsageReportScopeAccountId(report);
return {
provider: report.provider,
limits: report.limits.length,
account: accountLabel,
};
}),
});
return deduped;
}
/**
* Marks the current session's credential as temporarily blocked due to usage limits.
* Uses usage reports to determine accurate reset time when available.
* Returns true if a credential was blocked, enabling automatic fallback to the next credential.
*/
async markUsageLimitReached(
provider: string,
sessionId: string | undefined,
options?: { retryAfterMs?: number; baseUrl?: string },
): Promise<boolean> {
const sessionCredential = this.#getSessionCredential(provider, sessionId);
if (!sessionCredential) return false;
const providerKey = this.#getProviderTypeKey(provider, sessionCredential.type);
const now = this.#usageNow();
let blockedUntil = now + (options?.retryAfterMs ?? AuthStorage.#defaultBackoffMs);
if (provider === "openai-codex" && sessionCredential.type === "oauth") {
const credential = this.#getCredentialsForProvider(provider)[sessionCredential.index];
if (credential?.type === "oauth") {
const report = await this.#getUsageReport(provider, credential, options);
if (report && this.#isUsageLimitReached(report)) {
const resetAtMs = this.#getUsageResetAtMs(report, this.#usageNow());
if (resetAtMs && resetAtMs > blockedUntil) {
blockedUntil = resetAtMs;
}
}
}
}
this.#markCredentialBlocked(providerKey, sessionCredential.index, blockedUntil);
const remainingCredentials = this.#getCredentialsForProvider(provider)
.map((credential, index) => ({ credential, index }))
.filter(
(entry): entry is { credential: AuthCredential; index: number } =>
entry.credential.type === sessionCredential.type && entry.index !== sessionCredential.index,
);
return remainingCredentials.some(candidate => !this.#isCredentialBlocked(providerKey, candidate.index));
}
/**
* Resolves an OAuth API key, 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.
*/
async #resolveOAuthApiKey(
provider: string,
sessionId?: string,
options?: { baseUrl?: string },
): Promise<string | undefined> {
const credentials = this.#getCredentialsForProvider(provider)
.map((credential, index) => ({ credential, index }))
.filter((entry): entry is { credential: OAuthCredential; index: number } => entry.credential.type === "oauth");
if (credentials.length === 0) return undefined;
const providerKey = this.#getProviderTypeKey(provider, "oauth");
const order = this.#getCredentialOrder(providerKey, sessionId, credentials.length);
const fallback = credentials[order[0]];
const checkUsage = provider === "openai-codex" && credentials.length > 1;
for (const idx of order) {
const selection = credentials[idx];
const apiKey = await this.#tryOAuthCredential(
provider,
selection,
providerKey,
sessionId,
options,
checkUsage,
false,
);
if (apiKey) return apiKey;
}
if (fallback && this.#isCredentialBlocked(providerKey, fallback.index)) {
return this.#tryOAuthCredential(provider, fallback, providerKey, sessionId, options, checkUsage, true);
}
return undefined;
}
/** Attempts to use a single OAuth credential, checking usage and refreshing token. */
async #tryOAuthCredential(
provider: string,
selection: { credential: OAuthCredential; index: number },
providerKey: string,
sessionId: string | undefined,
options: { baseUrl?: string } | undefined,
checkUsage: boolean,
allowBlocked: boolean,
): Promise<string | undefined> {
if (!allowBlocked && this.#isCredentialBlocked(providerKey, selection.index)) {
return undefined;
}
let usage: UsageReport | null = null;
let usageChecked = false;
if (checkUsage && !allowBlocked) {
usage = await this.#getUsageReport(provider, selection.credential, options);
usageChecked = true;
if (usage && this.#isUsageLimitReached(usage)) {
const resetAtMs = this.#getUsageResetAtMs(usage, this.#usageNow());
this.#markCredentialBlocked(
providerKey,
selection.index,
resetAtMs ?? this.#usageNow() + AuthStorage.#defaultBackoffMs,
);
return undefined;
}
}
try {
let result: { newCredentials: OAuthCredentials; apiKey: string } | null;
const customProvider = getOAuthProvider(provider);
if (customProvider) {
let refreshedCredentials: OAuthCredentials = selection.credential;
if (Date.now() >= refreshedCredentials.expires) {
if (!customProvider.refreshToken) {
throw new Error(`OAuth provider "${provider}" does not support token refresh`);
}
refreshedCredentials = await customProvider.refreshToken(refreshedCredentials);
}
const apiKey = customProvider.getApiKey
? customProvider.getApiKey(refreshedCredentials)
: refreshedCredentials.access;
result = { newCredentials: refreshedCredentials, apiKey };
} else {
const oauthCreds: Record<string, OAuthCredentials> = {
[provider]: selection.credential,
};
result = await getOAuthApiKey(provider as OAuthProvider, oauthCreds);
}
if (!result) return undefined;
const updated: OAuthCredential = {
type: "oauth",
access: result.newCredentials.access,
refresh: result.newCredentials.refresh,
expires: result.newCredentials.expires,
accountId: result.newCredentials.accountId ?? selection.credential.accountId,
email: result.newCredentials.email ?? selection.credential.email,
projectId: result.newCredentials.projectId ?? selection.credential.projectId,
enterpriseUrl: result.newCredentials.enterpriseUrl ?? selection.credential.enterpriseUrl,
};
this.#replaceCredentialAt(provider, selection.index, updated);
if (checkUsage && !allowBlocked) {
const sameAccount = selection.credential.accountId === updated.accountId;
if (!usageChecked || !sameAccount) {
usage = await this.#getUsageReport(provider, updated, options);
}
if (usage && this.#isUsageLimitReached(usage)) {
const resetAtMs = this.#getUsageResetAtMs(usage, this.#usageNow());
this.#markCredentialBlocked(
providerKey,
selection.index,
resetAtMs ?? this.#usageNow() + AuthStorage.#defaultBackoffMs,
);
return undefined;
}
}
this.#recordSessionCredential(provider, sessionId, "oauth", selection.index);
return result.apiKey;
} catch (error) {
const errorMsg = String(error);
// Only remove credentials for definitive auth failures
// Keep credentials for transient errors (network, 5xx) and block temporarily
const isDefinitiveFailure =
/invalid_grant|invalid_token|revoked|unauthorized|expired.*refresh|refresh.*expired/i.test(errorMsg) ||
(/401|403/.test(errorMsg) && !/timeout|network|fetch failed|ECONNREFUSED/i.test(errorMsg));
logger.warn("OAuth token refresh failed", {
provider,
index: selection.index,
error: errorMsg,
isDefinitiveFailure,
});
if (isDefinitiveFailure) {
// Permanently remove invalid credentials
this.#removeCredentialAt(provider, selection.index);
if (this.#getCredentialsForProvider(provider).some(credential => credential.type === "oauth")) {
return this.getApiKey(provider, sessionId, options);
}
} else {
// Block temporarily for transient failures (5 minutes)
this.#markCredentialBlocked(providerKey, selection.index, this.#usageNow() + 5 * 60 * 1000);
}
}
return undefined;
}
/**
* Get API key for a provider.
* Priority:
* 1. Runtime override (CLI --api-key)
* 2. API key from agent.db
* 3. OAuth token from agent.db (auto-refreshed)
* 4. Environment variable
* 5. Fallback resolver (models.json custom providers)
*/
async getApiKey(provider: string, sessionId?: string, options?: { baseUrl?: string }): Promise<string | undefined> {
// Runtime override takes highest priority
const runtimeKey = this.#runtimeOverrides.get(provider);
if (runtimeKey) {
return runtimeKey;
}
const apiKeySelection = this.#selectCredentialByType(provider, "api_key", sessionId);
if (apiKeySelection) {
this.#recordSessionCredential(provider, sessionId, "api_key", apiKeySelection.index);
return resolveConfigValue(apiKeySelection.credential.key);
}
const oauthKey = await this.#resolveOAuthApiKey(provider, sessionId, options);
if (oauthKey) {
return oauthKey;
}
// Fall back to environment variable
const envKey = getEnvApiKey(provider);
if (envKey) return envKey;
// Fall back to custom resolver (e.g., models.json custom providers)
return this.#fallbackResolver?.(provider) ?? undefined;
}
}