diff --git a/.env.example b/.env.example index 0a222a0..df45908 100644 --- a/.env.example +++ b/.env.example @@ -5,3 +5,6 @@ TURSO_AUTH_TOKEN= # Optional: local server port. Defaults to 3003. PORT=3003 + +# Optional: rotate to the next provider account when upstream returns 429. +KLEIS_RATE_LIMIT_FAILOVER=1 diff --git a/README.md b/README.md index efe6bd9..3692ff6 100644 --- a/README.md +++ b/README.md @@ -45,6 +45,8 @@ ADMIN_TOKEN=replace-with-a-long-random-token CRON_SECRET=replace-with-a-long-random-token TURSO_CONNECTION_URL=libsql://..turso.io TURSO_AUTH_TOKEN= +# Optional: rotate to the next provider account when upstream returns 429. +KLEIS_RATE_LIMIT_FAILOVER=1 ``` ```sh diff --git a/src/domain/providers/provider-service.ts b/src/domain/providers/provider-service.ts index f477216..1d24c6f 100644 --- a/src/domain/providers/provider-service.ts +++ b/src/domain/providers/provider-service.ts @@ -5,8 +5,10 @@ import { findProviderAccountsByIds, findPrimaryProviderAccount, hasActiveProviderAccountRefreshLock, + listProviderAccounts, recordProviderAccountRefreshFailure, releaseProviderAccountRefreshLock, + setPrimaryProviderAccount, tryAcquireProviderAccountRefreshLock, updateProviderAccountTokens, upsertProviderAccount, @@ -335,6 +337,42 @@ const getPrimaryProviderAccount = async ( return refreshProviderAccount(database, account.id, now); }; +export const rotateProviderPrimaryAccount = async ( + database: Database, + provider: Provider, + currentAccountId: string, + now: number +): Promise => { + const accounts = (await listProviderAccounts(database)) + .filter((account) => account.provider === provider) + .sort((left, right) => right.createdAt - left.createdAt); + const currentIndex = accounts.findIndex( + (account) => account.id === currentAccountId + ); + if (currentIndex === -1 || !accounts[currentIndex]?.isPrimary) { + return null; + } + + const candidates = [ + ...accounts.slice(currentIndex + 1), + ...accounts.slice(0, currentIndex), + ]; + + for (const account of candidates) { + const refreshed = + account.expiresAt > now + ? account + : await refreshProviderAccount(database, account.id, now).catch( + () => null + ); + if (refreshed) { + return setPrimaryProviderAccount(database, refreshed.id, Date.now()); + } + } + + return null; +}; + export const getRoutableProviderAccount = async ( database: Database, provider: Provider, diff --git a/src/http/routes/proxy.ts b/src/http/routes/proxy.ts index 1483aa1..dcee020 100644 --- a/src/http/routes/proxy.ts +++ b/src/http/routes/proxy.ts @@ -6,7 +6,10 @@ import { recordRequestUsage, recordTokenUsage, } from "../../db/repositories/request-usage"; -import { getRoutableProviderAccount } from "../../domain/providers/provider-service"; +import { + getRoutableProviderAccount, + rotateProviderPrimaryAccount, +} from "../../domain/providers/provider-service"; import { prepareClaudeProxyRequest } from "../../providers/proxies/claude-proxy"; import { deriveCodexSessionId, @@ -371,6 +374,16 @@ const proxyRequest = async ( throw error; } + if ( + process.env.KLEIS_RATE_LIMIT_FAILOVER === "1" && + upstreamResponse.status === 429 && + !accountScopeIds?.length + ) { + runInBackground( + rotateProviderPrimaryAccount(db, route.provider, account.id, Date.now()) + ); + } + let responseToClient = upstreamResponse; if (responseTransformer) { try {