6c884713af
Fixes #111
1401 lines
45 KiB
TypeScript
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;
|
|
}
|
|
}
|