math-solver / llm /key_pool.py
Cuong2004
Deploy API from GitHub Actions
0772b5a
Raw
History Blame Contribute Delete
6.36 kB
import time
import random
import logging
from typing import Dict, List, Optional, Tuple
from llm.key_state import KeyState, KeyMetadata, BaseKeyStateStore, RedisKeyStateStore, MemoryKeyStateStore, hash_key
from config.settings import settings
logger = logging.getLogger(__name__)
class APIKeyPool:
"""Multi-Provider API Key Pool with Round-Robin Selection and Automatic State Management."""
def __init__(self, state_store: Optional[BaseKeyStateStore] = None):
if state_store is not None:
self.state_store = state_store
else:
try:
self.state_store = RedisKeyStateStore(settings.redis_url)
except Exception:
self.state_store = MemoryKeyStateStore()
self._provider_keys: Dict[str, List[str]] = {}
self._provider_indices: Dict[str, int] = {}
self._init_from_settings()
def _init_from_settings(self) -> None:
creds = settings.get_provider_credentials()
for provider, prov_cred in creds.items():
self._provider_keys[provider] = list(prov_cred.keys)
self._provider_indices[provider] = 0
logger.info(f"[APIKeyPool] Initialized pool for '{provider}' with {len(prov_cred.keys)} keys")
def register_keys(self, provider: str, keys: List[str]) -> None:
if provider not in self._provider_keys:
self._provider_keys[provider] = []
self._provider_indices[provider] = 0
for k in keys:
if k and k not in self._provider_keys[provider]:
self._provider_keys[provider].append(k)
async def get_next_key(self, provider: str) -> Optional[Tuple[str, str]]:
"""
Returns (api_key, key_hash) for an AVAILABLE key using round-robin.
If all keys are on cooldown, returns the one with the earliest retry_at.
If no keys configured, returns None.
"""
keys = self._provider_keys.get(provider, [])
if not keys:
# Check fallback to gemini or any available
if provider == "google":
keys = self._provider_keys.get("gemini", [])
elif provider in ("openai", "openrouter"):
keys = self._provider_keys.get(provider, [])
if not keys:
return None
n = len(keys)
start_idx = self._provider_indices.get(provider, 0)
# 1. Round-robin search for AVAILABLE key
for offset in range(n):
idx = (start_idx + offset) % n
candidate_key = keys[idx]
k_hash = hash_key(candidate_key)
meta = await self.state_store.get_state(k_hash)
if meta.state == KeyState.AVAILABLE:
self._provider_indices[provider] = (idx + 1) % n
return candidate_key, k_hash
# 2. If all keys are on cooldown/exhausted, find the earliest cooldown recovery
best_candidate: Optional[Tuple[str, str, float]] = None
for candidate_key in keys:
k_hash = hash_key(candidate_key)
meta = await self.state_store.get_state(k_hash)
if meta.state == KeyState.COOLDOWN:
if best_candidate is None or meta.retry_at < best_candidate[2]:
best_candidate = (candidate_key, k_hash, meta.retry_at)
if best_candidate:
candidate_key, k_hash, retry_at = best_candidate
now = time.time()
wait_needed = max(0.0, retry_at - now)
logger.warning(
f"[APIKeyPool] All {provider} keys on cooldown. Key {k_hash} available in {wait_needed:.1f}s"
)
# If wait is very short (< 3s), use it
if wait_needed < 3.0:
return candidate_key, k_hash
# Return first non-disabled key as emergency attempt
for candidate_key in keys:
k_hash = hash_key(candidate_key)
meta = await self.state_store.get_state(k_hash)
if meta.state != KeyState.DISABLED:
return candidate_key, k_hash
return None
async def mark_success(self, key: str) -> None:
k_hash = hash_key(key)
meta = await self.state_store.get_state(k_hash)
meta.state = KeyState.AVAILABLE
meta.failure_count = 0
meta.last_error = None
await self.state_store.set_state(k_hash, meta)
async def mark_cooldown(self, key: str, retry_after: Optional[int] = None, error_msg: Optional[str] = None) -> None:
k_hash = hash_key(key)
meta = await self.state_store.get_state(k_hash)
meta.failure_count += 1
# Exponential backoff with jitter
base_cooldown = retry_after if retry_after else settings.llm_cooldown_seconds
factor = min(2 ** (meta.failure_count - 1), 8)
jitter = random.uniform(0.8, 1.2)
total_duration = int(base_cooldown * factor * jitter)
meta.state = KeyState.COOLDOWN
meta.retry_at = time.time() + total_duration
meta.last_error = error_msg
logger.warning(f"[APIKeyPool] Key {k_hash} moved to COOLDOWN for {total_duration}s (failures={meta.failure_count})")
await self.state_store.set_state(k_hash, meta, ttl_seconds=total_duration + 60)
async def mark_exhausted(self, key: str, error_msg: Optional[str] = None) -> None:
k_hash = hash_key(key)
meta = await self.state_store.get_state(k_hash)
meta.state = KeyState.EXHAUSTED
meta.last_error = error_msg
# Cooldown for 6 hours
meta.retry_at = time.time() + 21600
logger.error(f"[APIKeyPool] Key {k_hash} marked EXHAUSTED (daily quota): {error_msg}")
await self.state_store.set_state(k_hash, meta, ttl_seconds=21600)
async def mark_disabled(self, key: str, error_msg: Optional[str] = None) -> None:
k_hash = hash_key(key)
meta = await self.state_store.get_state(k_hash)
meta.state = KeyState.DISABLED
meta.last_error = error_msg
logger.error(f"[APIKeyPool] Key {k_hash} DISABLED permanently (Auth/Invalid): {error_msg}")
await self.state_store.set_state(k_hash, meta)
_GLOBAL_KEY_POOL: Optional[APIKeyPool] = None
def get_key_pool() -> APIKeyPool:
global _GLOBAL_KEY_POOL
if _GLOBAL_KEY_POOL is None:
_GLOBAL_KEY_POOL = APIKeyPool()
return _GLOBAL_KEY_POOL