File size: 6,359 Bytes
0772b5a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
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