keycycle 0.3.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- keycycle/__init__.py +48 -0
- keycycle/adapters/__init__.py +22 -0
- keycycle/adapters/generic_adapter.py +663 -0
- keycycle/adapters/openai_adapter.py +381 -0
- keycycle/config/__init__.py +25 -0
- keycycle/config/constants.py +35 -0
- keycycle/config/dataclasses.py +225 -0
- keycycle/config/enums.py +5 -0
- keycycle/config/loader.py +54 -0
- keycycle/config/log_config.py +53 -0
- keycycle/config/models/cerebras.yaml +36 -0
- keycycle/config/models/cohere.yaml +14 -0
- keycycle/config/models/gemini.yaml +69 -0
- keycycle/config/models/groq.yaml +161 -0
- keycycle/config/models/mistral.yaml +13 -0
- keycycle/config/models/moonshot.yaml +59 -0
- keycycle/config/models/openrouter.yaml +4 -0
- keycycle/config/models/openrouter_models.yaml +97 -0
- keycycle/config/models.py +53 -0
- keycycle/core/__init__.py +40 -0
- keycycle/core/backoff.py +79 -0
- keycycle/core/exceptions.py +94 -0
- keycycle/core/utils.py +480 -0
- keycycle/key_rotation/__init__.py +7 -0
- keycycle/key_rotation/rotating_mixin.py +392 -0
- keycycle/key_rotation/rotation_manager.py +231 -0
- keycycle/legacy_multi_provider_wrapper.py +734 -0
- keycycle/multi_client_wrapper.py +372 -0
- keycycle/py.typed +0 -0
- keycycle/usage/__init__.py +6 -0
- keycycle/usage/db_logic.py +98 -0
- keycycle/usage/usage_logger.py +89 -0
- keycycle/utils.py +27 -0
- keycycle-0.3.2.dist-info/METADATA +44 -0
- keycycle-0.3.2.dist-info/RECORD +37 -0
- keycycle-0.3.2.dist-info/WHEEL +5 -0
- keycycle-0.3.2.dist-info/top_level.txt +1 -0
keycycle/core/backoff.py
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
"""Exponential backoff utilities for key rotation."""
|
|
2
|
+
import random
|
|
3
|
+
import time
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from typing import Optional
|
|
6
|
+
|
|
7
|
+
from ..config.constants import DEFAULT_POLL_INTERVAL, MIN_POLL_INTERVAL, MAX_POLL_INTERVAL
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@dataclass
|
|
11
|
+
class BackoffConfig:
|
|
12
|
+
"""Configuration for exponential backoff."""
|
|
13
|
+
initial_interval: float = DEFAULT_POLL_INTERVAL
|
|
14
|
+
max_interval: float = MAX_POLL_INTERVAL
|
|
15
|
+
multiplier: float = 2.0
|
|
16
|
+
jitter: float = 0.1 # 10% jitter
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class ExponentialBackoff:
|
|
20
|
+
"""
|
|
21
|
+
Implements exponential backoff with jitter for key polling.
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
backoff = ExponentialBackoff()
|
|
25
|
+
for attempt in range(max_attempts):
|
|
26
|
+
key = try_get_key()
|
|
27
|
+
if key:
|
|
28
|
+
return key
|
|
29
|
+
backoff.wait()
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
def __init__(self, config: Optional[BackoffConfig] = None):
|
|
33
|
+
self.config = config or BackoffConfig()
|
|
34
|
+
self._attempt = 0
|
|
35
|
+
self._current_interval = self.config.initial_interval
|
|
36
|
+
|
|
37
|
+
def reset(self) -> None:
|
|
38
|
+
"""Reset backoff state to initial values."""
|
|
39
|
+
self._attempt = 0
|
|
40
|
+
self._current_interval = self.config.initial_interval
|
|
41
|
+
|
|
42
|
+
def get_next_interval(self) -> float:
|
|
43
|
+
"""
|
|
44
|
+
Get the next backoff interval without waiting.
|
|
45
|
+
|
|
46
|
+
Returns:
|
|
47
|
+
The calculated interval with jitter applied
|
|
48
|
+
"""
|
|
49
|
+
# Calculate base interval
|
|
50
|
+
interval = min(self._current_interval, self.config.max_interval)
|
|
51
|
+
|
|
52
|
+
# Apply jitter (+-jitter%)
|
|
53
|
+
jitter_range = interval * self.config.jitter
|
|
54
|
+
interval += random.uniform(-jitter_range, jitter_range)
|
|
55
|
+
|
|
56
|
+
# Ensure minimum
|
|
57
|
+
interval = max(interval, MIN_POLL_INTERVAL)
|
|
58
|
+
|
|
59
|
+
# Update for next call
|
|
60
|
+
self._current_interval *= self.config.multiplier
|
|
61
|
+
self._attempt += 1
|
|
62
|
+
|
|
63
|
+
return interval
|
|
64
|
+
|
|
65
|
+
def wait(self) -> float:
|
|
66
|
+
"""
|
|
67
|
+
Wait for the next backoff interval.
|
|
68
|
+
|
|
69
|
+
Returns:
|
|
70
|
+
The actual time waited
|
|
71
|
+
"""
|
|
72
|
+
interval = self.get_next_interval()
|
|
73
|
+
time.sleep(interval)
|
|
74
|
+
return interval
|
|
75
|
+
|
|
76
|
+
@property
|
|
77
|
+
def attempt(self) -> int:
|
|
78
|
+
"""Current attempt number (0-indexed)."""
|
|
79
|
+
return self._attempt
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
"""Custom exception hierarchy for keycycle."""
|
|
2
|
+
from typing import Any, Optional
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class KeycycleError(Exception):
|
|
6
|
+
"""Base exception for all keycycle errors."""
|
|
7
|
+
pass
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class KeycycleKeyError(KeycycleError):
|
|
11
|
+
"""Base for key-related errors."""
|
|
12
|
+
pass
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class NoAvailableKeyError(KeycycleKeyError):
|
|
16
|
+
"""Raised when no API keys are available for use."""
|
|
17
|
+
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
provider: str,
|
|
21
|
+
model_id: str,
|
|
22
|
+
wait: bool,
|
|
23
|
+
timeout: float,
|
|
24
|
+
total_keys: int = 0,
|
|
25
|
+
cooling_down: int = 0
|
|
26
|
+
):
|
|
27
|
+
self.provider = provider
|
|
28
|
+
self.model_id = model_id
|
|
29
|
+
self.wait = wait
|
|
30
|
+
self.timeout = timeout
|
|
31
|
+
self.total_keys = total_keys
|
|
32
|
+
self.cooling_down = cooling_down
|
|
33
|
+
|
|
34
|
+
if wait:
|
|
35
|
+
msg = f"Timeout: No available API keys for {provider}/{model_id} after {timeout}s"
|
|
36
|
+
else:
|
|
37
|
+
msg = f"No available API keys for {provider}/{model_id} (wait=False)"
|
|
38
|
+
|
|
39
|
+
if total_keys > 0:
|
|
40
|
+
msg += f" [{cooling_down}/{total_keys} keys cooling down]"
|
|
41
|
+
|
|
42
|
+
super().__init__(msg)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class KeyNotFoundError(KeycycleKeyError):
|
|
46
|
+
"""Raised when a specific key identifier cannot be found."""
|
|
47
|
+
|
|
48
|
+
def __init__(self, identifier: Any):
|
|
49
|
+
self.identifier = identifier
|
|
50
|
+
super().__init__(f"Key with identifier '{identifier}' not found.")
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class InvalidKeyError(KeycycleKeyError):
|
|
54
|
+
"""Raised when an API key is invalid (401/403 response)."""
|
|
55
|
+
|
|
56
|
+
def __init__(self, key_suffix: str, status_code: int):
|
|
57
|
+
self.key_suffix = key_suffix
|
|
58
|
+
self.status_code = status_code
|
|
59
|
+
super().__init__(f"API key ...{key_suffix} is invalid (HTTP {status_code})")
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
class RateLimitError(KeycycleError):
|
|
63
|
+
"""Raised when rate limit is hit and retries exhausted."""
|
|
64
|
+
|
|
65
|
+
def __init__(self, provider: str, model_id: str, attempts: int):
|
|
66
|
+
self.provider = provider
|
|
67
|
+
self.model_id = model_id
|
|
68
|
+
self.attempts = attempts
|
|
69
|
+
super().__init__(
|
|
70
|
+
f"Rate limit exceeded for {provider}/{model_id} after {attempts} attempts"
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class ConfigurationError(KeycycleError):
|
|
75
|
+
"""Raised for configuration-related errors."""
|
|
76
|
+
pass
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class MissingEnvironmentVariableError(ConfigurationError):
|
|
80
|
+
"""Raised when a required environment variable is missing."""
|
|
81
|
+
|
|
82
|
+
def __init__(self, var_name: str):
|
|
83
|
+
self.var_name = var_name
|
|
84
|
+
super().__init__(f"Environment variable '{var_name}' not found.")
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class InvalidConfigurationError(ConfigurationError):
|
|
88
|
+
"""Raised when configuration values are invalid."""
|
|
89
|
+
|
|
90
|
+
def __init__(self, config_name: str, value: Any, expected: str):
|
|
91
|
+
self.config_name = config_name
|
|
92
|
+
self.value = value
|
|
93
|
+
self.expected = expected
|
|
94
|
+
super().__init__(f"'{config_name}' must be {expected}, got: {value}")
|
keycycle/core/utils.py
ADDED
|
@@ -0,0 +1,480 @@
|
|
|
1
|
+
"""Shared utility functions for keycycle."""
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import re
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any, Callable, Dict, FrozenSet, List, Optional, Tuple, Union
|
|
7
|
+
|
|
8
|
+
from dotenv import load_dotenv
|
|
9
|
+
|
|
10
|
+
from ..config.constants import KEY_SUFFIX_LENGTH
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
# Type alias for key entries: either a string API key or a dict with params
|
|
14
|
+
KeyEntry = Union[str, Dict[str, Any]]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def normalize_key_entry(entry: KeyEntry, api_key_param: str = "api_key") -> Tuple[str, Dict[str, Any]]:
|
|
18
|
+
"""
|
|
19
|
+
Normalize a key entry to (primary_key, params_dict).
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
entry: Either a string API key or a dict containing api_key and other params
|
|
23
|
+
api_key_param: Name of the API key parameter in the dict (default: "api_key")
|
|
24
|
+
|
|
25
|
+
Returns:
|
|
26
|
+
Tuple of (primary_key, params_dict) where params_dict always contains the api_key_param
|
|
27
|
+
|
|
28
|
+
Raises:
|
|
29
|
+
ValueError: If entry is a dict but doesn't contain api_key_param
|
|
30
|
+
TypeError: If entry is neither str nor dict
|
|
31
|
+
"""
|
|
32
|
+
if isinstance(entry, str):
|
|
33
|
+
return entry, {api_key_param: entry}
|
|
34
|
+
if isinstance(entry, dict):
|
|
35
|
+
if api_key_param not in entry:
|
|
36
|
+
raise ValueError(f"Key dict must contain '{api_key_param}'")
|
|
37
|
+
return entry[api_key_param], dict(entry)
|
|
38
|
+
raise TypeError(f"Key entry must be str or dict, got {type(entry).__name__}")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def normalize_key_entries(
|
|
42
|
+
entries: List[KeyEntry],
|
|
43
|
+
api_key_param: str = "api_key"
|
|
44
|
+
) -> List[Tuple[str, Dict[str, Any]]]:
|
|
45
|
+
"""
|
|
46
|
+
Normalize a list of key entries.
|
|
47
|
+
|
|
48
|
+
Args:
|
|
49
|
+
entries: List of key entries (strings or dicts)
|
|
50
|
+
api_key_param: Name of the API key parameter
|
|
51
|
+
|
|
52
|
+
Returns:
|
|
53
|
+
List of (primary_key, params_dict) tuples
|
|
54
|
+
"""
|
|
55
|
+
return [normalize_key_entry(e, api_key_param) for e in entries]
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
# Rate limit indicators for detection
|
|
59
|
+
RATE_LIMIT_INDICATORS: FrozenSet[str] = frozenset([
|
|
60
|
+
"429", "too many requests", "rate limit",
|
|
61
|
+
"resource exhausted", "traffic", "rate-limited"
|
|
62
|
+
])
|
|
63
|
+
|
|
64
|
+
# Temporary/transient rate limit indicators - retry with SAME key
|
|
65
|
+
TEMPORARY_RATE_LIMIT_INDICATORS: FrozenSet[str] = frozenset([
|
|
66
|
+
"temporarily rate-limited",
|
|
67
|
+
"temporarily unavailable",
|
|
68
|
+
"high traffic",
|
|
69
|
+
"retry shortly",
|
|
70
|
+
"please retry",
|
|
71
|
+
"try again shortly",
|
|
72
|
+
"rate-limited upstream",
|
|
73
|
+
"experiencing high load",
|
|
74
|
+
])
|
|
75
|
+
|
|
76
|
+
# Hard/quota-based rate limit indicators - rotate to different key
|
|
77
|
+
HARD_RATE_LIMIT_INDICATORS: FrozenSet[str] = frozenset([
|
|
78
|
+
"per-day",
|
|
79
|
+
"per-hour",
|
|
80
|
+
"per-minute",
|
|
81
|
+
"quota exceeded",
|
|
82
|
+
"daily limit",
|
|
83
|
+
"hourly limit",
|
|
84
|
+
"free-models-per-day",
|
|
85
|
+
"x-ratelimit-remaining: 0",
|
|
86
|
+
])
|
|
87
|
+
|
|
88
|
+
# Auth error indicators
|
|
89
|
+
AUTH_STATUS_CODES = {401, 403}
|
|
90
|
+
AUTH_INDICATORS: FrozenSet[str] = frozenset([
|
|
91
|
+
"401", "403", "unauthorized", "forbidden",
|
|
92
|
+
"invalid api key", "invalid_api_key", "expired"
|
|
93
|
+
])
|
|
94
|
+
|
|
95
|
+
# Payment required indicators (e.g. free-tier quota/credits exhausted)
|
|
96
|
+
PAYMENT_REQUIRED_STATUS_CODES = {402}
|
|
97
|
+
PAYMENT_REQUIRED_INDICATORS: FrozenSet[str] = frozenset([
|
|
98
|
+
"402", "payment required", "payment_required",
|
|
99
|
+
])
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def get_key_suffix(api_key: str, length: int = KEY_SUFFIX_LENGTH) -> str:
|
|
103
|
+
"""
|
|
104
|
+
Extract the suffix of an API key for logging and identification.
|
|
105
|
+
|
|
106
|
+
Args:
|
|
107
|
+
api_key: The full API key
|
|
108
|
+
length: Number of characters to extract (default: 8)
|
|
109
|
+
|
|
110
|
+
Returns:
|
|
111
|
+
The last `length` characters, or the full key if shorter
|
|
112
|
+
"""
|
|
113
|
+
return api_key[-length:] if len(api_key) > length else api_key
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def is_rate_limit_error(e: Exception) -> bool:
|
|
117
|
+
"""
|
|
118
|
+
Heuristic to detect rate limit errors across different providers.
|
|
119
|
+
|
|
120
|
+
Checks:
|
|
121
|
+
1. Exception message for rate limit keywords
|
|
122
|
+
2. status_code attribute for 429
|
|
123
|
+
3. body/response attribute for rate limit indicators
|
|
124
|
+
|
|
125
|
+
Args:
|
|
126
|
+
e: The exception to check
|
|
127
|
+
|
|
128
|
+
Returns:
|
|
129
|
+
True if this appears to be a rate limit error
|
|
130
|
+
"""
|
|
131
|
+
# Check string representation
|
|
132
|
+
err_str = str(e).lower()
|
|
133
|
+
if any(indicator in err_str for indicator in RATE_LIMIT_INDICATORS):
|
|
134
|
+
return True
|
|
135
|
+
|
|
136
|
+
# Check status_code attribute
|
|
137
|
+
if hasattr(e, "status_code") and e.status_code == 429:
|
|
138
|
+
return True
|
|
139
|
+
|
|
140
|
+
# Check body/response for embedded error info
|
|
141
|
+
body = getattr(e, "body", None) or getattr(e, "response", None)
|
|
142
|
+
if body:
|
|
143
|
+
body_str = str(body).lower()
|
|
144
|
+
if any(indicator in body_str for indicator in RATE_LIMIT_INDICATORS):
|
|
145
|
+
return True
|
|
146
|
+
|
|
147
|
+
return False
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def is_temporary_rate_limit_error(e: Exception) -> bool:
|
|
151
|
+
"""
|
|
152
|
+
Detect if a rate limit error is temporary (upstream congestion)
|
|
153
|
+
vs. a hard quota limit.
|
|
154
|
+
|
|
155
|
+
Temporary errors should be retried with the SAME key after a delay.
|
|
156
|
+
Hard limit errors should trigger key rotation.
|
|
157
|
+
|
|
158
|
+
Args:
|
|
159
|
+
e: The exception to check
|
|
160
|
+
|
|
161
|
+
Returns:
|
|
162
|
+
True if this is a temporary rate limit that should be retried
|
|
163
|
+
with the same key (no rotation needed)
|
|
164
|
+
"""
|
|
165
|
+
# First, must be a rate limit error at all
|
|
166
|
+
if not is_rate_limit_error(e):
|
|
167
|
+
return False
|
|
168
|
+
|
|
169
|
+
# Build combined error string from all available sources
|
|
170
|
+
err_str = str(e).lower()
|
|
171
|
+
|
|
172
|
+
body = getattr(e, "body", None) or getattr(e, "response", None)
|
|
173
|
+
if body:
|
|
174
|
+
err_str += " " + str(body).lower()
|
|
175
|
+
|
|
176
|
+
# Check for hard limit indicators first (these take precedence)
|
|
177
|
+
if any(indicator in err_str for indicator in HARD_RATE_LIMIT_INDICATORS):
|
|
178
|
+
return False
|
|
179
|
+
|
|
180
|
+
# Check for temporary indicators
|
|
181
|
+
if any(indicator in err_str for indicator in TEMPORARY_RATE_LIMIT_INDICATORS):
|
|
182
|
+
return True
|
|
183
|
+
|
|
184
|
+
return False
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def is_payment_required_error(e: Exception) -> bool:
|
|
188
|
+
"""
|
|
189
|
+
Detect HTTP 402 Payment Required errors (e.g. a Cerebras free-trial key
|
|
190
|
+
whose credits/quota have been exhausted).
|
|
191
|
+
|
|
192
|
+
Unlike rate limit errors, a 402 will not clear on its own after a short
|
|
193
|
+
cooldown - the key is exhausted for good (until credits are topped up),
|
|
194
|
+
so callers should treat the key as permanently dead for this process
|
|
195
|
+
rather than retrying it later.
|
|
196
|
+
|
|
197
|
+
Args:
|
|
198
|
+
e: The exception to check
|
|
199
|
+
|
|
200
|
+
Returns:
|
|
201
|
+
True if this appears to be a payment-required / quota-exhausted error
|
|
202
|
+
"""
|
|
203
|
+
# Check status_code attribute
|
|
204
|
+
if hasattr(e, "status_code") and e.status_code in PAYMENT_REQUIRED_STATUS_CODES:
|
|
205
|
+
return True
|
|
206
|
+
|
|
207
|
+
# Check string representation
|
|
208
|
+
err_str = str(e).lower()
|
|
209
|
+
if any(indicator in err_str for indicator in PAYMENT_REQUIRED_INDICATORS):
|
|
210
|
+
return True
|
|
211
|
+
|
|
212
|
+
# Check body/response for embedded error info
|
|
213
|
+
body = getattr(e, "body", None) or getattr(e, "response", None)
|
|
214
|
+
if body:
|
|
215
|
+
body_str = str(body).lower()
|
|
216
|
+
if any(indicator in body_str for indicator in PAYMENT_REQUIRED_INDICATORS):
|
|
217
|
+
return True
|
|
218
|
+
|
|
219
|
+
return False
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def is_auth_error(e: Exception) -> bool:
|
|
223
|
+
"""
|
|
224
|
+
Detect authentication/authorization errors (invalid/expired keys).
|
|
225
|
+
|
|
226
|
+
Args:
|
|
227
|
+
e: The exception to check
|
|
228
|
+
|
|
229
|
+
Returns:
|
|
230
|
+
True if this appears to be an auth error (401/403)
|
|
231
|
+
"""
|
|
232
|
+
# Check status_code
|
|
233
|
+
if hasattr(e, "status_code") and e.status_code in AUTH_STATUS_CODES:
|
|
234
|
+
return True
|
|
235
|
+
|
|
236
|
+
# Check string representation
|
|
237
|
+
err_str = str(e).lower()
|
|
238
|
+
if any(indicator in err_str for indicator in AUTH_INDICATORS):
|
|
239
|
+
return True
|
|
240
|
+
|
|
241
|
+
return False
|
|
242
|
+
|
|
243
|
+
|
|
244
|
+
def validate_api_key(api_key: str) -> bool:
|
|
245
|
+
"""
|
|
246
|
+
Basic validation of API key format.
|
|
247
|
+
|
|
248
|
+
Args:
|
|
249
|
+
api_key: The API key to validate
|
|
250
|
+
|
|
251
|
+
Returns:
|
|
252
|
+
True if the key appears to be valid format
|
|
253
|
+
|
|
254
|
+
Note:
|
|
255
|
+
This only checks format, not whether the key actually works.
|
|
256
|
+
"""
|
|
257
|
+
if not api_key or not isinstance(api_key, str):
|
|
258
|
+
return False
|
|
259
|
+
|
|
260
|
+
# Most API keys are at least 20 characters
|
|
261
|
+
if len(api_key) < 20:
|
|
262
|
+
return False
|
|
263
|
+
|
|
264
|
+
# Check for common placeholder patterns (exact matches or obvious placeholders)
|
|
265
|
+
placeholder_patterns = ["your_api_key", "placeholder", "your-api-key", "insert_key", "api_key_here"]
|
|
266
|
+
key_lower = api_key.lower()
|
|
267
|
+
if any(pattern in key_lower for pattern in placeholder_patterns):
|
|
268
|
+
return False
|
|
269
|
+
|
|
270
|
+
# Check if key is all x's or looks like a template (simplified with regex)
|
|
271
|
+
stripped = re.sub(r'[-_]|sk', '', key_lower)
|
|
272
|
+
if stripped and all(c == 'x' for c in stripped):
|
|
273
|
+
return False
|
|
274
|
+
|
|
275
|
+
return True
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def load_api_keys(
|
|
279
|
+
provider: str,
|
|
280
|
+
env_file: Optional[str] = None,
|
|
281
|
+
extra_params: Optional[List[str]] = None,
|
|
282
|
+
api_key_param: str = "api_key",
|
|
283
|
+
) -> List[KeyEntry]:
|
|
284
|
+
"""
|
|
285
|
+
Load API keys from environment variables.
|
|
286
|
+
|
|
287
|
+
Args:
|
|
288
|
+
provider: Provider name (e.g., 'twelvelabs', 'anthropic')
|
|
289
|
+
env_file: Path to .env file (optional)
|
|
290
|
+
extra_params: List of extra parameter names to load alongside API keys.
|
|
291
|
+
For each param, loads {PROVIDER}_{PARAM}_N environment variables.
|
|
292
|
+
api_key_param: Name for the API key in returned dicts (default: "api_key")
|
|
293
|
+
|
|
294
|
+
Returns:
|
|
295
|
+
List of key entries. If extra_params is provided, returns list of dicts
|
|
296
|
+
with api_key and extra params. Otherwise returns list of strings.
|
|
297
|
+
|
|
298
|
+
Environment variables format:
|
|
299
|
+
NUM_{PROVIDER}=N
|
|
300
|
+
{PROVIDER}_API_KEY_1, {PROVIDER}_API_KEY_2, ...
|
|
301
|
+
{PROVIDER}_{PARAM}_1, {PROVIDER}_{PARAM}_2, ... (for each extra_param)
|
|
302
|
+
|
|
303
|
+
Example:
|
|
304
|
+
>>> # With extra_params=["index_id"], loads:
|
|
305
|
+
>>> # TWELVELABS_API_KEY_1, TWELVELABS_INDEX_ID_1
|
|
306
|
+
>>> # TWELVELABS_API_KEY_2, TWELVELABS_INDEX_ID_2
|
|
307
|
+
>>> keys = load_api_keys("twelvelabs", extra_params=["index_id"])
|
|
308
|
+
>>> # Returns: [{"api_key": "key1", "index_id": "idx1"}, ...]
|
|
309
|
+
"""
|
|
310
|
+
if env_file:
|
|
311
|
+
env_path = Path(env_file).resolve()
|
|
312
|
+
else:
|
|
313
|
+
env_path = Path.cwd() / ".env"
|
|
314
|
+
load_dotenv(dotenv_path=env_path, override=True)
|
|
315
|
+
|
|
316
|
+
provider_upper = provider.upper()
|
|
317
|
+
num_keys_var = f"NUM_{provider_upper}"
|
|
318
|
+
num_keys_raw = os.getenv(num_keys_var)
|
|
319
|
+
if not num_keys_raw:
|
|
320
|
+
raise ValueError(f"Environment variable '{num_keys_var}' not found.")
|
|
321
|
+
try:
|
|
322
|
+
expected_count = int(num_keys_raw)
|
|
323
|
+
except ValueError:
|
|
324
|
+
raise ValueError(f"'{num_keys_var}' must be an integer, got: {num_keys_raw}")
|
|
325
|
+
|
|
326
|
+
# NUM_{PROVIDER} is AUTHORITATIVE (0.3.2, reverting 0.3.1's scan-all):
|
|
327
|
+
# the fleet convention is that the first NUM indices (sorted ascending)
|
|
328
|
+
# are the sanctioned live keys, and keys parked at higher indices beyond
|
|
329
|
+
# a numbering gap are dead or spare on purpose -- loading them would put
|
|
330
|
+
# known-bad keys back into rotation. Gaps INSIDE the sanctioned range
|
|
331
|
+
# are tolerated (a retired key's index can be missing); the loader takes
|
|
332
|
+
# the first NUM indices that actually exist and warns if it comes up
|
|
333
|
+
# short, but never loads past NUM keys.
|
|
334
|
+
key_pattern = re.compile(rf"^{re.escape(provider_upper)}_API_KEY_(\d+)$")
|
|
335
|
+
found: Dict[int, str] = {}
|
|
336
|
+
for env_name, value in os.environ.items():
|
|
337
|
+
match = key_pattern.match(env_name)
|
|
338
|
+
if match and value:
|
|
339
|
+
found[int(match.group(1))] = value
|
|
340
|
+
|
|
341
|
+
if not found:
|
|
342
|
+
raise ValueError(f"Missing API key: {provider_upper}_API_KEY_1")
|
|
343
|
+
|
|
344
|
+
all_indices = sorted(found)
|
|
345
|
+
indices = all_indices[:expected_count]
|
|
346
|
+
excluded = all_indices[expected_count:]
|
|
347
|
+
if excluded:
|
|
348
|
+
logging.getLogger(__name__).info(
|
|
349
|
+
"%s: %s=%d -- using the first %d key indices %s; excluding %d "
|
|
350
|
+
"parked key(s) at indices %s (dead/spare by convention).",
|
|
351
|
+
provider_upper, num_keys_var, expected_count, len(indices), indices,
|
|
352
|
+
len(excluded), excluded,
|
|
353
|
+
)
|
|
354
|
+
if len(indices) < expected_count:
|
|
355
|
+
logging.getLogger(__name__).warning(
|
|
356
|
+
"%s: only %d API key(s) found (indices %s) but %s=%d -- the "
|
|
357
|
+
"declared count is stale. Using the %d key(s) actually found.",
|
|
358
|
+
provider_upper, len(indices), indices, num_keys_var, expected_count,
|
|
359
|
+
len(indices),
|
|
360
|
+
)
|
|
361
|
+
|
|
362
|
+
keys: List[KeyEntry] = []
|
|
363
|
+
for i in indices:
|
|
364
|
+
key = found[i]
|
|
365
|
+
if extra_params:
|
|
366
|
+
key_dict: Dict[str, Any] = {api_key_param: key}
|
|
367
|
+
for param in extra_params:
|
|
368
|
+
val = os.getenv(f"{provider_upper}_{param.upper()}_{i}")
|
|
369
|
+
if val:
|
|
370
|
+
key_dict[param] = val
|
|
371
|
+
keys.append(key_dict)
|
|
372
|
+
else:
|
|
373
|
+
keys.append(key)
|
|
374
|
+
return keys
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
# Type alias for key limit overrides - can be imported from config.dataclasses
|
|
378
|
+
KeyLimitOverride = Any # Actually Union[RateLimits, Dict[str, RateLimits]] but avoiding circular import
|
|
379
|
+
|
|
380
|
+
|
|
381
|
+
def normalize_key_limits(
|
|
382
|
+
api_keys: List[str],
|
|
383
|
+
key_limits: Optional[Dict[Union[int, str], KeyLimitOverride]],
|
|
384
|
+
logger: Optional[logging.Logger] = None,
|
|
385
|
+
) -> Dict[str, KeyLimitOverride]:
|
|
386
|
+
"""
|
|
387
|
+
Normalize key_limits by converting all identifiers to key suffixes.
|
|
388
|
+
|
|
389
|
+
Accepts:
|
|
390
|
+
- Integer index (0-based): Maps to the key at that position
|
|
391
|
+
- String suffix: Matched against the last N characters of each key
|
|
392
|
+
- Full key string: Exact match
|
|
393
|
+
|
|
394
|
+
Args:
|
|
395
|
+
api_keys: List of primary API key strings
|
|
396
|
+
key_limits: Dict mapping identifiers to limit overrides
|
|
397
|
+
logger: Optional logger for warnings
|
|
398
|
+
|
|
399
|
+
Returns:
|
|
400
|
+
Dict mapping key suffixes to their limit overrides
|
|
401
|
+
"""
|
|
402
|
+
if not key_limits:
|
|
403
|
+
return {}
|
|
404
|
+
|
|
405
|
+
normalized: Dict[str, KeyLimitOverride] = {}
|
|
406
|
+
|
|
407
|
+
for identifier, limits in key_limits.items():
|
|
408
|
+
if isinstance(identifier, int):
|
|
409
|
+
# Index-based: get the key at that position
|
|
410
|
+
if 0 <= identifier < len(api_keys):
|
|
411
|
+
suffix = get_key_suffix(api_keys[identifier])
|
|
412
|
+
normalized[suffix] = limits
|
|
413
|
+
elif logger:
|
|
414
|
+
logger.warning(
|
|
415
|
+
"key_limits index %d is out of range (0-%d)",
|
|
416
|
+
identifier, len(api_keys) - 1
|
|
417
|
+
)
|
|
418
|
+
elif isinstance(identifier, str):
|
|
419
|
+
# String: try exact match first, then suffix match
|
|
420
|
+
matched = False
|
|
421
|
+
for key in api_keys:
|
|
422
|
+
if key == identifier or key.endswith(identifier):
|
|
423
|
+
suffix = get_key_suffix(key)
|
|
424
|
+
normalized[suffix] = limits
|
|
425
|
+
matched = True
|
|
426
|
+
break
|
|
427
|
+
if not matched and logger:
|
|
428
|
+
logger.warning(
|
|
429
|
+
"key_limits identifier '%s' did not match any API key",
|
|
430
|
+
identifier[-8:]
|
|
431
|
+
)
|
|
432
|
+
|
|
433
|
+
return normalized
|
|
434
|
+
|
|
435
|
+
|
|
436
|
+
def resolve_limits(
|
|
437
|
+
model_id: str,
|
|
438
|
+
default_model_id: str,
|
|
439
|
+
key_suffix: Optional[str],
|
|
440
|
+
key_limits: Dict[str, KeyLimitOverride],
|
|
441
|
+
model_limits: Dict[str, Dict[str, Any]],
|
|
442
|
+
provider: str,
|
|
443
|
+
default_limits_factory: Callable[[], Any],
|
|
444
|
+
) -> Any:
|
|
445
|
+
"""
|
|
446
|
+
Resolve rate limits for a model, optionally with per-key overrides.
|
|
447
|
+
|
|
448
|
+
Args:
|
|
449
|
+
model_id: The model identifier
|
|
450
|
+
default_model_id: Fallback model ID if model_id is None
|
|
451
|
+
key_suffix: Optional key suffix for per-key limit lookup
|
|
452
|
+
key_limits: Dict of key suffix -> limit overrides
|
|
453
|
+
model_limits: Dict of provider -> model -> RateLimits (e.g., MODEL_LIMITS)
|
|
454
|
+
provider: Provider name for looking up model limits
|
|
455
|
+
default_limits_factory: Callable that returns default RateLimits
|
|
456
|
+
|
|
457
|
+
Returns:
|
|
458
|
+
RateLimits for the model/key combination
|
|
459
|
+
"""
|
|
460
|
+
from ..config.dataclasses import RateLimits
|
|
461
|
+
|
|
462
|
+
mid = model_id or default_model_id
|
|
463
|
+
|
|
464
|
+
# Check key-specific overrides first
|
|
465
|
+
if key_suffix and key_limits:
|
|
466
|
+
override = key_limits.get(key_suffix)
|
|
467
|
+
if override:
|
|
468
|
+
if isinstance(override, RateLimits):
|
|
469
|
+
return override
|
|
470
|
+
elif isinstance(override, dict):
|
|
471
|
+
# Per-model limits for this key
|
|
472
|
+
if mid in override:
|
|
473
|
+
return override[mid]
|
|
474
|
+
# Check for __default__ fallback
|
|
475
|
+
if '__default__' in override:
|
|
476
|
+
return override['__default__']
|
|
477
|
+
|
|
478
|
+
# Fall back to provider/model defaults
|
|
479
|
+
provider_limits = model_limits.get(provider, {})
|
|
480
|
+
return provider_limits.get(mid, provider_limits.get('default', default_limits_factory()))
|