keycycle 0.1.11__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 +8 -0
- keycycle/adapters/__init__.py +3 -0
- keycycle/adapters/openai_adapter.py +266 -0
- keycycle/config/__init__.py +25 -0
- keycycle/config/constants.py +47 -0
- keycycle/config/dataclasses.py +190 -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 +55 -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/openrouter.yaml +4 -0
- keycycle/config/models/openrouter_models.yaml +97 -0
- keycycle/key_rotation/__init__.py +7 -0
- keycycle/key_rotation/rotating_mixin.py +226 -0
- keycycle/key_rotation/rotation_manager.py +184 -0
- keycycle/multi_provider_wrapper.py +518 -0
- keycycle/usage/__init__.py +6 -0
- keycycle/usage/db_logic.py +95 -0
- keycycle/usage/usage_logger.py +85 -0
- keycycle/utils.py +26 -0
- keycycle-0.1.11.dist-info/METADATA +35 -0
- keycycle-0.1.11.dist-info/RECORD +27 -0
- keycycle-0.1.11.dist-info/WHEEL +5 -0
- keycycle-0.1.11.dist-info/top_level.txt +1 -0
keycycle/__init__.py
ADDED
|
@@ -0,0 +1,266 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import time
|
|
3
|
+
from typing import List, Any, Optional, Union, Callable, Generator, AsyncGenerator
|
|
4
|
+
|
|
5
|
+
from ..config.dataclasses import KeyUsage, RateLimits
|
|
6
|
+
from ..key_rotation.rotation_manager import RotatingKeyManager
|
|
7
|
+
|
|
8
|
+
logger = logging.getLogger(__name__)
|
|
9
|
+
|
|
10
|
+
try:
|
|
11
|
+
import openai
|
|
12
|
+
from openai import OpenAI, AsyncOpenAI, RateLimitError, APIError
|
|
13
|
+
HAS_OPENAI = True
|
|
14
|
+
except ImportError:
|
|
15
|
+
HAS_OPENAI = False
|
|
16
|
+
# Generic placeholders
|
|
17
|
+
OpenAI = object
|
|
18
|
+
AsyncOpenAI = object
|
|
19
|
+
RateLimitError = Exception
|
|
20
|
+
APIError = Exception
|
|
21
|
+
|
|
22
|
+
PROVIDER_BASE_URLS = {
|
|
23
|
+
"openai": "https://api.openai.com/v1",
|
|
24
|
+
"openrouter": "https://openrouter.ai/api/v1",
|
|
25
|
+
"gemini": "https://generativelanguage.googleapis.com/v1beta/openai/",
|
|
26
|
+
"cerebras": "https://api.cerebras.ai/v1",
|
|
27
|
+
"groq": "https://api.groq.com/openai/v1",
|
|
28
|
+
"cohere": "https://api.cohere.ai/compatibility/v1",
|
|
29
|
+
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
class BaseRotatingClient:
|
|
33
|
+
def __init__(self,
|
|
34
|
+
manager: RotatingKeyManager,
|
|
35
|
+
limit_resolver: Callable[[str], RateLimits],
|
|
36
|
+
default_model: str,
|
|
37
|
+
estimated_tokens: int = 1000,
|
|
38
|
+
max_retries: int = 5,
|
|
39
|
+
base_url: Optional[str] = None,
|
|
40
|
+
provider: Optional[str] = None,
|
|
41
|
+
client_kwargs: dict = None
|
|
42
|
+
):
|
|
43
|
+
"""
|
|
44
|
+
Initialize the rotating client.
|
|
45
|
+
|
|
46
|
+
Args:
|
|
47
|
+
manager: The key rotation manager
|
|
48
|
+
limit_resolver: Function to resolve rate limits for a model
|
|
49
|
+
default_model: Default model to use
|
|
50
|
+
estimated_tokens: Estimated tokens per request
|
|
51
|
+
max_retries: Maximum number of retries on rate limit
|
|
52
|
+
base_url: Base URL for the API (takes precedence over provider)
|
|
53
|
+
provider: Provider name (openai, openrouter, gemini, cerebras, groq)
|
|
54
|
+
client_kwargs: Additional kwargs to pass to OpenAI client
|
|
55
|
+
"""
|
|
56
|
+
|
|
57
|
+
if not HAS_OPENAI:
|
|
58
|
+
raise ImportError("The 'openai' library is required. Install with `pip install openai`.")
|
|
59
|
+
|
|
60
|
+
self.manager = manager
|
|
61
|
+
self.limit_resolver = limit_resolver
|
|
62
|
+
self.default_model = default_model
|
|
63
|
+
self.estimated_tokens = estimated_tokens
|
|
64
|
+
self.max_retries = max_retries
|
|
65
|
+
self.client_kwargs = client_kwargs or {}
|
|
66
|
+
|
|
67
|
+
if base_url:
|
|
68
|
+
self.base_url = base_url
|
|
69
|
+
elif provider:
|
|
70
|
+
provider_lower = provider.lower()
|
|
71
|
+
if provider_lower not in PROVIDER_BASE_URLS:
|
|
72
|
+
raise ValueError(
|
|
73
|
+
f"Unknown provider: {provider}. "
|
|
74
|
+
f"Valid providers: {', '.join(PROVIDER_BASE_URLS.keys())}"
|
|
75
|
+
)
|
|
76
|
+
self.base_url = PROVIDER_BASE_URLS[provider_lower]
|
|
77
|
+
else:
|
|
78
|
+
self.base_url = None # Will use OpenAI's default
|
|
79
|
+
|
|
80
|
+
if self.base_url:
|
|
81
|
+
self.client_kwargs['base_url'] = self.base_url
|
|
82
|
+
|
|
83
|
+
def _is_rate_limit_error(self, e: Exception) -> bool:
|
|
84
|
+
if isinstance(e, RateLimitError):
|
|
85
|
+
return True
|
|
86
|
+
err_str = str(e).lower()
|
|
87
|
+
flagged_strs = ["429", "too many requests", "rate limit", "resource exhausted", "traffic", "rate-limited"]
|
|
88
|
+
if any(i in err_str for i in
|
|
89
|
+
flagged_strs):
|
|
90
|
+
return True
|
|
91
|
+
|
|
92
|
+
if hasattr(e, "status_code") and e.status_code == 429:
|
|
93
|
+
return True
|
|
94
|
+
body = getattr(e, "body", None) or getattr(e, "response", None)
|
|
95
|
+
if body:
|
|
96
|
+
body_str = str(body).lower()
|
|
97
|
+
if any(i in body_str for i in flagged_strs):
|
|
98
|
+
return True
|
|
99
|
+
return False
|
|
100
|
+
|
|
101
|
+
def _record_usage(self, key_usage: KeyUsage, model_id: str, actual_tokens: int):
|
|
102
|
+
self.manager.record_usage(
|
|
103
|
+
key_obj=key_usage,
|
|
104
|
+
model_id=model_id,
|
|
105
|
+
actual_tokens=actual_tokens,
|
|
106
|
+
estimated_tokens=self.estimated_tokens
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
def _extract_usage(self, response: Any) -> int:
|
|
110
|
+
try:
|
|
111
|
+
if hasattr(response, 'usage') and response.usage:
|
|
112
|
+
return response.usage.total_tokens
|
|
113
|
+
except:
|
|
114
|
+
pass
|
|
115
|
+
return 0
|
|
116
|
+
|
|
117
|
+
# --- SYNC IMPLEMENTATION ---
|
|
118
|
+
|
|
119
|
+
class RotatingOpenAIClient(BaseRotatingClient):
|
|
120
|
+
def _get_fresh_client(self, api_key: str):
|
|
121
|
+
return OpenAI(api_key=api_key, **self.client_kwargs)
|
|
122
|
+
|
|
123
|
+
def __getattr__(self, name):
|
|
124
|
+
return SyncProxyHelper(self, [name])
|
|
125
|
+
|
|
126
|
+
def _execute(self, path: List[str], args, kwargs):
|
|
127
|
+
model_id = kwargs.get('model', self.default_model)
|
|
128
|
+
if 'model' not in kwargs:
|
|
129
|
+
kwargs['model'] = model_id
|
|
130
|
+
limits = self.limit_resolver(model_id)
|
|
131
|
+
|
|
132
|
+
for attempt in range(self.max_retries + 1):
|
|
133
|
+
key_usage = self.manager.get_key(model_id, limits, self.estimated_tokens)
|
|
134
|
+
if not key_usage:
|
|
135
|
+
raise RuntimeError(f"No available keys for {model_id}")
|
|
136
|
+
|
|
137
|
+
try:
|
|
138
|
+
real_client = self._get_fresh_client(key_usage.api_key)
|
|
139
|
+
|
|
140
|
+
target = real_client
|
|
141
|
+
for p in path:
|
|
142
|
+
target = getattr(target, p)
|
|
143
|
+
|
|
144
|
+
result = target(*args, **kwargs)
|
|
145
|
+
|
|
146
|
+
if kwargs.get('stream', False):
|
|
147
|
+
return self._wrap_stream(result, key_usage, model_id)
|
|
148
|
+
|
|
149
|
+
self._record_usage(key_usage, model_id, self._extract_usage(result))
|
|
150
|
+
return result
|
|
151
|
+
|
|
152
|
+
except Exception as e:
|
|
153
|
+
if self._is_rate_limit_error(e) and attempt < self.max_retries:
|
|
154
|
+
logger.warning(f"429/RateLimit hit for {model_id} on key ...{key_usage.api_key[-8:]}. Rotating. (Attempt {attempt + 1}/{self.max_retries + 1})")
|
|
155
|
+
key_usage.trigger_cooldown()
|
|
156
|
+
self.manager.force_rotate_index()
|
|
157
|
+
time.sleep(0.5)
|
|
158
|
+
continue
|
|
159
|
+
self._record_usage(key_usage, model_id, 0)
|
|
160
|
+
raise e
|
|
161
|
+
|
|
162
|
+
def _wrap_stream(self, generator: Generator, key_usage: KeyUsage, model_id: str):
|
|
163
|
+
accumulated_tokens = 0
|
|
164
|
+
try:
|
|
165
|
+
for chunk in generator:
|
|
166
|
+
if hasattr(chunk, 'usage') and chunk.usage:
|
|
167
|
+
accumulated_tokens = chunk.usage.total_tokens
|
|
168
|
+
yield chunk
|
|
169
|
+
except Exception as e:
|
|
170
|
+
if self._is_rate_limit_error(e):
|
|
171
|
+
logger.warning(f"Rate limit hit during streaming for {model_id} on key ...{key_usage.api_key[-8:]}.")
|
|
172
|
+
key_usage.trigger_cooldown()
|
|
173
|
+
self.manager.force_rotate_index()
|
|
174
|
+
raise
|
|
175
|
+
finally:
|
|
176
|
+
self._record_usage(key_usage, model_id, accumulated_tokens)
|
|
177
|
+
|
|
178
|
+
class SyncProxyHelper:
|
|
179
|
+
def __init__(self, client: RotatingOpenAIClient, path: List[str]):
|
|
180
|
+
self.client = client
|
|
181
|
+
self.path = path
|
|
182
|
+
|
|
183
|
+
def __getattr__(self, name):
|
|
184
|
+
return SyncProxyHelper(self.client, self.path + [name])
|
|
185
|
+
|
|
186
|
+
def __call__(self, *args, **kwargs):
|
|
187
|
+
return self.client._execute(self.path, args, kwargs)
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
# --- ASYNC IMPLEMENTATION ---
|
|
191
|
+
|
|
192
|
+
class RotatingAsyncOpenAIClient(BaseRotatingClient):
|
|
193
|
+
def _get_fresh_client(self, api_key: str):
|
|
194
|
+
return AsyncOpenAI(api_key=api_key, **self.client_kwargs)
|
|
195
|
+
|
|
196
|
+
def __getattr__(self, name):
|
|
197
|
+
return AsyncProxyHelper(self, [name])
|
|
198
|
+
|
|
199
|
+
async def _execute(self, path: List[str], args, kwargs: dict):
|
|
200
|
+
model_id = kwargs.get('model', self.default_model)
|
|
201
|
+
if 'model' not in kwargs:
|
|
202
|
+
kwargs['model'] = model_id
|
|
203
|
+
limits = self.limit_resolver(model_id)
|
|
204
|
+
|
|
205
|
+
if kwargs.get('stream', False) and 'stream_options' not in kwargs:
|
|
206
|
+
kwargs['stream_options'] = {"include_usage": True}
|
|
207
|
+
|
|
208
|
+
if kwargs.get('stream', False) and 'stream_options' not in kwargs:
|
|
209
|
+
kwargs['stream_options'] = {"include_usage": True}
|
|
210
|
+
|
|
211
|
+
for attempt in range(self.max_retries + 1):
|
|
212
|
+
key_usage = self.manager.get_key(model_id, limits, self.estimated_tokens)
|
|
213
|
+
if not key_usage:
|
|
214
|
+
raise RuntimeError(f"No available keys for {model_id}")
|
|
215
|
+
|
|
216
|
+
try:
|
|
217
|
+
real_client = self._get_fresh_client(key_usage.api_key)
|
|
218
|
+
|
|
219
|
+
target = real_client
|
|
220
|
+
for p in path:
|
|
221
|
+
target = getattr(target, p)
|
|
222
|
+
|
|
223
|
+
result = await target(*args, **kwargs)
|
|
224
|
+
|
|
225
|
+
if kwargs.get('stream', False):
|
|
226
|
+
return self._wrap_stream(result, key_usage, model_id)
|
|
227
|
+
|
|
228
|
+
self._record_usage(key_usage, model_id, self._extract_usage(result))
|
|
229
|
+
return result
|
|
230
|
+
|
|
231
|
+
except Exception as e:
|
|
232
|
+
if self._is_rate_limit_error(e) and attempt < self.max_retries:
|
|
233
|
+
logger.warning(f"429/RateLimit hit for {model_id} on key ...{key_usage.api_key[-8:]}. Rotating. (Attempt {attempt + 1}/{self.max_retries + 1})")
|
|
234
|
+
key_usage.trigger_cooldown()
|
|
235
|
+
self.manager.force_rotate_index()
|
|
236
|
+
time.sleep(0.5)
|
|
237
|
+
continue
|
|
238
|
+
self._record_usage(key_usage, model_id, 0)
|
|
239
|
+
raise e
|
|
240
|
+
|
|
241
|
+
async def _wrap_stream(self, generator: AsyncGenerator, key_usage: KeyUsage, model_id: str):
|
|
242
|
+
accumulated_tokens = 0
|
|
243
|
+
try:
|
|
244
|
+
async for chunk in generator:
|
|
245
|
+
if hasattr(chunk, 'usage') and chunk.usage:
|
|
246
|
+
accumulated_tokens = chunk.usage.total_tokens
|
|
247
|
+
yield chunk
|
|
248
|
+
except Exception as e:
|
|
249
|
+
if self._is_rate_limit_error(e):
|
|
250
|
+
logger.warning(f"Rate limit hit during streaming for {model_id} on key ...{key_usage.api_key[-8:]}.")
|
|
251
|
+
key_usage.trigger_cooldown()
|
|
252
|
+
self.manager.force_rotate_index()
|
|
253
|
+
raise
|
|
254
|
+
finally:
|
|
255
|
+
self._record_usage(key_usage, model_id, accumulated_tokens)
|
|
256
|
+
|
|
257
|
+
class AsyncProxyHelper:
|
|
258
|
+
def __init__(self, client: RotatingAsyncOpenAIClient, path: List[str]):
|
|
259
|
+
self.client = client
|
|
260
|
+
self.path = path
|
|
261
|
+
|
|
262
|
+
def __getattr__(self, name):
|
|
263
|
+
return AsyncProxyHelper(self.client, self.path + [name])
|
|
264
|
+
|
|
265
|
+
async def __call__(self, *args, **kwargs):
|
|
266
|
+
return await self.client._execute(self.path, args, kwargs)
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
from .dataclasses import (
|
|
2
|
+
RateLimits,
|
|
3
|
+
UsageSnapshot,
|
|
4
|
+
UsageBucket,
|
|
5
|
+
KeyUsage,
|
|
6
|
+
KeyDetailedStats,
|
|
7
|
+
KeySummary,
|
|
8
|
+
GlobalStats,
|
|
9
|
+
ModelAggregatedStats,
|
|
10
|
+
)
|
|
11
|
+
from .enums import RateLimitStrategy
|
|
12
|
+
from .log_config import configure_logging
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
"RateLimits",
|
|
16
|
+
"UsageSnapshot",
|
|
17
|
+
"UsageBucket",
|
|
18
|
+
"KeyUsage",
|
|
19
|
+
"KeyDetailedStats",
|
|
20
|
+
"KeySummary",
|
|
21
|
+
"GlobalStats",
|
|
22
|
+
"ModelAggregatedStats",
|
|
23
|
+
"RateLimitStrategy",
|
|
24
|
+
"configure_logging",
|
|
25
|
+
]
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from .dataclasses import RateLimits
|
|
3
|
+
from .enums import RateLimitStrategy
|
|
4
|
+
from typing import Any, TypedDict
|
|
5
|
+
from .loader import load_rate_limits_from_yaml, load_openrouter_models
|
|
6
|
+
|
|
7
|
+
# Get directory of this file
|
|
8
|
+
CONFIG_DIR = os.path.dirname(os.path.abspath(__file__))
|
|
9
|
+
MODELS_DIR = os.path.join(CONFIG_DIR, 'models')
|
|
10
|
+
|
|
11
|
+
class ModelDict(TypedDict):
|
|
12
|
+
"""All the provider currrently supported"""
|
|
13
|
+
cerebras: Any
|
|
14
|
+
groq: Any
|
|
15
|
+
gemini: Any
|
|
16
|
+
openrouter: Any
|
|
17
|
+
cohere: Any
|
|
18
|
+
|
|
19
|
+
PROVIDER_STRATEGIES: ModelDict = {
|
|
20
|
+
'cerebras': RateLimitStrategy.PER_MODEL,
|
|
21
|
+
'groq': RateLimitStrategy.PER_MODEL,
|
|
22
|
+
'gemini': RateLimitStrategy.PER_MODEL,
|
|
23
|
+
'openrouter': RateLimitStrategy.GLOBAL,
|
|
24
|
+
'cohere': RateLimitStrategy.PER_MODEL
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
# Load Cohere tiers first to handle the env var logic
|
|
28
|
+
_cohere_tiers = load_rate_limits_from_yaml(os.path.join(MODELS_DIR, 'cohere.yaml'))
|
|
29
|
+
|
|
30
|
+
COHERE_TIERS = {
|
|
31
|
+
'free': _cohere_tiers['free'],
|
|
32
|
+
'pro': _cohere_tiers['pro'],
|
|
33
|
+
'enterprise': _cohere_tiers['enterprise'],
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
MODEL_LIMITS: ModelDict = {
|
|
37
|
+
'cerebras': load_rate_limits_from_yaml(os.path.join(MODELS_DIR, 'cerebras.yaml')),
|
|
38
|
+
'groq': load_rate_limits_from_yaml(os.path.join(MODELS_DIR, 'groq.yaml')),
|
|
39
|
+
'gemini': load_rate_limits_from_yaml(os.path.join(MODELS_DIR, 'gemini.yaml')),
|
|
40
|
+
'openrouter': load_rate_limits_from_yaml(os.path.join(MODELS_DIR, 'openrouter.yaml')),
|
|
41
|
+
'cohere':{
|
|
42
|
+
# Same for every model
|
|
43
|
+
'default': COHERE_TIERS.get(os.getenv('COHERE_TIER', 'free').lower(), COHERE_TIERS['free'])
|
|
44
|
+
}
|
|
45
|
+
}
|
|
46
|
+
|
|
47
|
+
OPENROUTER_MODELS = load_openrouter_models(os.path.join(MODELS_DIR, 'openrouter_models.yaml'))
|
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
from typing import Dict, List, Optional
|
|
2
|
+
from .enums import RateLimitStrategy
|
|
3
|
+
import time
|
|
4
|
+
from collections import deque, defaultdict
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
|
|
7
|
+
# --- CONFIGURATION DATA ---
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class RateLimits:
|
|
11
|
+
"""Rate limits for a provider"""
|
|
12
|
+
requests_per_minute: int
|
|
13
|
+
requests_per_hour: int
|
|
14
|
+
requests_per_day: int
|
|
15
|
+
tokens_per_minute: Optional[int] = None
|
|
16
|
+
tokens_per_hour: Optional[int] = None
|
|
17
|
+
tokens_per_day: Optional[int] = None
|
|
18
|
+
|
|
19
|
+
@dataclass
|
|
20
|
+
class UsageSnapshot:
|
|
21
|
+
"""Standardized view of usage counters"""
|
|
22
|
+
rpm: int = 0
|
|
23
|
+
rph: int = 0
|
|
24
|
+
rpd: int = 0
|
|
25
|
+
tpm: int = 0
|
|
26
|
+
tph: int = 0
|
|
27
|
+
tpd: int = 0
|
|
28
|
+
total_requests: int = 0
|
|
29
|
+
total_tokens: int = 0
|
|
30
|
+
|
|
31
|
+
def __add__(self, other):
|
|
32
|
+
"""Allow summing snapshots for aggregation"""
|
|
33
|
+
if not isinstance(other, UsageSnapshot): return NotImplemented
|
|
34
|
+
return UsageSnapshot(
|
|
35
|
+
self.rpm + other.rpm, self.rph + other.rph, self.rpd + other.rpd,
|
|
36
|
+
self.tpm + other.tpm, self.tph + other.tph, self.tpd + other.tpd,
|
|
37
|
+
self.total_requests + other.total_requests, self.total_tokens + other.total_tokens
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
# --- STATS DATA TRANSFER OBJECTS (DTOs) ---
|
|
41
|
+
|
|
42
|
+
@dataclass
|
|
43
|
+
class KeySummary:
|
|
44
|
+
index: int; suffix: str; snapshot: UsageSnapshot
|
|
45
|
+
@dataclass
|
|
46
|
+
class GlobalStats:
|
|
47
|
+
total: UsageSnapshot; keys: List[KeySummary]
|
|
48
|
+
@dataclass
|
|
49
|
+
class KeyDetailedStats:
|
|
50
|
+
index: int; suffix: str; total: UsageSnapshot; breakdown: Dict[str, UsageSnapshot]
|
|
51
|
+
@dataclass
|
|
52
|
+
class ModelAggregatedStats:
|
|
53
|
+
model_id: str; total: UsageSnapshot; keys: List[KeySummary]
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
# --- USAGE TRACKING ---
|
|
57
|
+
@dataclass
|
|
58
|
+
class UsageBucket:
|
|
59
|
+
"""Tracks counters for a SINGLE model context"""
|
|
60
|
+
requests_minute: deque[float] = field(default_factory=deque)
|
|
61
|
+
requests_hour: deque[float] = field(default_factory=deque)
|
|
62
|
+
requests_day: deque[float] = field(default_factory=deque)
|
|
63
|
+
|
|
64
|
+
tokens_minute: deque[tuple[float, int]] = field(default_factory=deque)
|
|
65
|
+
tokens_hour: deque[tuple[float, int]] = field(default_factory=deque)
|
|
66
|
+
tokens_day: deque[tuple[float, int]] = field(default_factory=deque)
|
|
67
|
+
|
|
68
|
+
total_requests: int = 0
|
|
69
|
+
total_tokens: int = 0
|
|
70
|
+
|
|
71
|
+
pending_tokens: int = 0
|
|
72
|
+
|
|
73
|
+
def clean(self):
|
|
74
|
+
"""Clean old entries based on current time"""
|
|
75
|
+
now = time.time()
|
|
76
|
+
cutoffs = (now - 60, now - 3600, now - 86400)
|
|
77
|
+
|
|
78
|
+
for d, cut in zip(
|
|
79
|
+
[self.requests_minute, self.requests_hour, self.requests_day], cutoffs
|
|
80
|
+
):
|
|
81
|
+
while d and d[0] <= cut: d.popleft()
|
|
82
|
+
|
|
83
|
+
for d, cut in zip(
|
|
84
|
+
[self.tokens_minute, self.tokens_hour, self.tokens_day], cutoffs
|
|
85
|
+
):
|
|
86
|
+
while d and d[0][0] <= cut: d.popleft()
|
|
87
|
+
|
|
88
|
+
def add(self, tokens: int, timestamp: float):
|
|
89
|
+
self.requests_minute.append(timestamp)
|
|
90
|
+
self.requests_hour.append(timestamp)
|
|
91
|
+
self.requests_day.append(timestamp)
|
|
92
|
+
self.total_requests += 1
|
|
93
|
+
|
|
94
|
+
if tokens > 0:
|
|
95
|
+
self.tokens_minute.append((timestamp, tokens))
|
|
96
|
+
self.tokens_hour.append((timestamp, tokens))
|
|
97
|
+
self.tokens_day.append((timestamp, tokens))
|
|
98
|
+
self.total_tokens += tokens
|
|
99
|
+
|
|
100
|
+
def check_limits(self, limits: RateLimits, estimated_tokens: int) -> bool:
|
|
101
|
+
self.clean()
|
|
102
|
+
if len(self.requests_minute) >= limits.requests_per_minute: return False
|
|
103
|
+
if len(self.requests_hour) >= limits.requests_per_hour: return False
|
|
104
|
+
if len(self.requests_day) >= limits.requests_per_day: return False
|
|
105
|
+
|
|
106
|
+
current_tpm = sum(t[1] for t in self.tokens_minute) + self.pending_tokens
|
|
107
|
+
current_tph = sum(t[1] for t in self.tokens_hour) + self.pending_tokens
|
|
108
|
+
current_tpd = sum(t[1] for t in self.tokens_day) + self.pending_tokens
|
|
109
|
+
|
|
110
|
+
if limits.tokens_per_minute and (current_tpm + estimated_tokens > limits.tokens_per_minute): return False
|
|
111
|
+
if limits.tokens_per_hour and (current_tph + estimated_tokens > limits.tokens_per_hour): return False
|
|
112
|
+
if limits.tokens_per_day and (current_tpd + estimated_tokens > limits.tokens_per_day): return False
|
|
113
|
+
|
|
114
|
+
return True
|
|
115
|
+
|
|
116
|
+
def reserve(self, tokens: int):
|
|
117
|
+
"""Lock in estimated tokens"""
|
|
118
|
+
self.pending_tokens += tokens
|
|
119
|
+
|
|
120
|
+
def commit(self, actual_tokens: int, reserved_tokens: int, timestamp: float):
|
|
121
|
+
"""Remove reservation and add actual usage"""
|
|
122
|
+
self.pending_tokens -= reserved_tokens
|
|
123
|
+
if self.pending_tokens < 0: self.pending_tokens = 0 # Safety floor
|
|
124
|
+
self.add(actual_tokens, timestamp)
|
|
125
|
+
|
|
126
|
+
def get_snapshot(self) -> UsageSnapshot:
|
|
127
|
+
"""Return current counts as a clean snapshot"""
|
|
128
|
+
self.clean()
|
|
129
|
+
return UsageSnapshot(
|
|
130
|
+
rpm=len(self.requests_minute),
|
|
131
|
+
rph=len(self.requests_hour),
|
|
132
|
+
rpd=len(self.requests_day),
|
|
133
|
+
tpm=sum(t[1] for t in self.tokens_minute),
|
|
134
|
+
tph=sum(t[1] for t in self.tokens_hour),
|
|
135
|
+
tpd=sum(t[1] for t in self.tokens_day),
|
|
136
|
+
total_requests=self.total_requests,
|
|
137
|
+
total_tokens=self.total_tokens
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
@dataclass
|
|
142
|
+
class KeyUsage:
|
|
143
|
+
"""Represents an API Key and holds multiple UsageBuckets (one per model)"""
|
|
144
|
+
api_key: str
|
|
145
|
+
strategy: RateLimitStrategy
|
|
146
|
+
buckets: Dict[str, UsageBucket] = field(default_factory=lambda: defaultdict(UsageBucket))
|
|
147
|
+
global_bucket: UsageBucket = field(default_factory=UsageBucket)
|
|
148
|
+
last_429: float = 0.0
|
|
149
|
+
|
|
150
|
+
def record_usage(self, model_id: str, tokens: int, timestamp: float = None):
|
|
151
|
+
ts = timestamp if timestamp else time.time()
|
|
152
|
+
self.buckets[model_id].add(tokens, ts)
|
|
153
|
+
if self.strategy == RateLimitStrategy.GLOBAL:
|
|
154
|
+
self.global_bucket.add(tokens, ts)
|
|
155
|
+
|
|
156
|
+
def can_use_model(self, model_id: str, limits: RateLimits, estimated_tokens: int = 1000) -> bool:
|
|
157
|
+
"""Check limits based on the provider's strategy"""
|
|
158
|
+
if self.strategy == RateLimitStrategy.GLOBAL:
|
|
159
|
+
return self.global_bucket.check_limits(limits, estimated_tokens)
|
|
160
|
+
else: # Per-Model Limits
|
|
161
|
+
return self.buckets[model_id].check_limits(limits, estimated_tokens)
|
|
162
|
+
|
|
163
|
+
def get_total_snapshot(self) -> UsageSnapshot:
|
|
164
|
+
if self.strategy == RateLimitStrategy.GLOBAL:
|
|
165
|
+
return self.global_bucket.get_snapshot()
|
|
166
|
+
total = UsageSnapshot()
|
|
167
|
+
for b in self.buckets.values():
|
|
168
|
+
total = total + b.get_snapshot()
|
|
169
|
+
return total
|
|
170
|
+
|
|
171
|
+
def reserve(self, model_id: str, tokens: int):
|
|
172
|
+
self.buckets[model_id].reserve(tokens)
|
|
173
|
+
if self.strategy == RateLimitStrategy.GLOBAL:
|
|
174
|
+
self.global_bucket.reserve(tokens)
|
|
175
|
+
|
|
176
|
+
def commit(self, model_id: str, actual_tokens: int, reserved_tokens: int, timestamp: float = None):
|
|
177
|
+
ts = timestamp if timestamp else time.time()
|
|
178
|
+
self.buckets[model_id].commit(actual_tokens, reserved_tokens, ts)
|
|
179
|
+
if self.strategy == RateLimitStrategy.GLOBAL:
|
|
180
|
+
self.global_bucket.commit(actual_tokens, reserved_tokens, ts)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def is_cooling_down(self, cooldown_seconds: int = 30) -> bool:
|
|
184
|
+
"""Returns True if the key is still in its 30s penalty box."""
|
|
185
|
+
if self.last_429 == 0: return False
|
|
186
|
+
return (time.time() - self.last_429) < cooldown_seconds
|
|
187
|
+
|
|
188
|
+
def trigger_cooldown(self):
|
|
189
|
+
"""Mark this key as rate-limited."""
|
|
190
|
+
self.last_429 = time.time()
|
keycycle/config/enums.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import yaml
|
|
3
|
+
from typing import Dict, Any, List, TypedDict, Optional
|
|
4
|
+
from .dataclasses import RateLimits
|
|
5
|
+
|
|
6
|
+
class RateLimitConfig(TypedDict):
|
|
7
|
+
requests_per_minute: int
|
|
8
|
+
requests_per_hour: int
|
|
9
|
+
requests_per_day: int
|
|
10
|
+
tokens_per_minute: Optional[int]
|
|
11
|
+
tokens_per_hour: Optional[int]
|
|
12
|
+
tokens_per_day: Optional[int]
|
|
13
|
+
|
|
14
|
+
class ModelConfig(TypedDict):
|
|
15
|
+
name: str
|
|
16
|
+
id: str
|
|
17
|
+
max_context_length: int
|
|
18
|
+
|
|
19
|
+
def load_yaml_config(file_path: str) -> Any:
|
|
20
|
+
"""Loads a YAML configuration file."""
|
|
21
|
+
if not os.path.exists(file_path):
|
|
22
|
+
raise FileNotFoundError(f"Configuration file not found: {file_path}")
|
|
23
|
+
|
|
24
|
+
with open(file_path, 'r', encoding='utf-8') as f:
|
|
25
|
+
return yaml.safe_load(f)
|
|
26
|
+
|
|
27
|
+
def load_rate_limits_from_yaml(file_path: str) -> Dict[str, RateLimits]:
|
|
28
|
+
"""Loads rate limits from a YAML file and converts them to RateLimits objects."""
|
|
29
|
+
data = load_yaml_config(file_path)
|
|
30
|
+
if not isinstance(data, dict):
|
|
31
|
+
raise ValueError(f"Invalid rate limit config in {file_path}, expected a dictionary.")
|
|
32
|
+
|
|
33
|
+
limits = {}
|
|
34
|
+
for model_id, config in data.items():
|
|
35
|
+
try:
|
|
36
|
+
limits[model_id] = RateLimits(
|
|
37
|
+
requests_per_minute=config['requests_per_minute'],
|
|
38
|
+
requests_per_hour=config['requests_per_hour'],
|
|
39
|
+
requests_per_day=config['requests_per_day'],
|
|
40
|
+
tokens_per_minute=config.get('tokens_per_minute'),
|
|
41
|
+
tokens_per_hour=config.get('tokens_per_hour'),
|
|
42
|
+
tokens_per_day=config.get('tokens_per_day'),
|
|
43
|
+
)
|
|
44
|
+
except KeyError as e:
|
|
45
|
+
raise ValueError(f"Missing required field {e} for model {model_id} in {file_path}")
|
|
46
|
+
|
|
47
|
+
return limits
|
|
48
|
+
|
|
49
|
+
def load_openrouter_models(file_path: str) -> List[ModelConfig]:
|
|
50
|
+
"""Loads the OpenRouter models list."""
|
|
51
|
+
data = load_yaml_config(file_path)
|
|
52
|
+
if not isinstance(data, list):
|
|
53
|
+
raise ValueError(f"Invalid OpenRouter config in {file_path}, expected a list.")
|
|
54
|
+
return data
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import logging.config
|
|
3
|
+
from logging.handlers import RotatingFileHandler
|
|
4
|
+
from rich.logging import RichHandler
|
|
5
|
+
|
|
6
|
+
LOGGING_CONFIG = {
|
|
7
|
+
"version": 1,
|
|
8
|
+
"disable_existing_loggers": False,
|
|
9
|
+
"formatters": {
|
|
10
|
+
"standard": {
|
|
11
|
+
"format": "%(asctime)s | %(levelname)-8s | %(name)s:%(lineno)d | %(message)s",
|
|
12
|
+
"datefmt": "%Y-%m-%d %H:%M:%S",
|
|
13
|
+
},
|
|
14
|
+
},
|
|
15
|
+
|
|
16
|
+
"handlers": {
|
|
17
|
+
"console": {
|
|
18
|
+
"class": "rich.logging.RichHandler",
|
|
19
|
+
"level": "DEBUG",
|
|
20
|
+
"rich_tracebacks": True, # Beautiful traceback formatting
|
|
21
|
+
"markup": True, # Allow "[bold]text[/]" inside log messages!
|
|
22
|
+
"show_path": False # Cleaner output (hides file path column)
|
|
23
|
+
},
|
|
24
|
+
"file": {
|
|
25
|
+
"class": "logging.handlers.RotatingFileHandler",
|
|
26
|
+
"level": "INFO",
|
|
27
|
+
"formatter": "standard",
|
|
28
|
+
"filename": "app.log",
|
|
29
|
+
"maxBytes": 10 * 1024 * 1024, # 10 MiB per file
|
|
30
|
+
"backupCount": 5,
|
|
31
|
+
"encoding": "utf-8",
|
|
32
|
+
},
|
|
33
|
+
},
|
|
34
|
+
|
|
35
|
+
"root": {
|
|
36
|
+
"handlers": ["console", "file"],
|
|
37
|
+
"level": "WARNING",
|
|
38
|
+
},
|
|
39
|
+
|
|
40
|
+
"loggers": {
|
|
41
|
+
"key_manager": {
|
|
42
|
+
"handlers": ["console", "file"],
|
|
43
|
+
"level": "DEBUG",
|
|
44
|
+
"propagate": False,
|
|
45
|
+
},
|
|
46
|
+
},
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
def configure_logging():
|
|
50
|
+
logging.config.dictConfig(LOGGING_CONFIG)
|
|
51
|
+
|
|
52
|
+
configure_logging()
|
|
53
|
+
default_logger = logging.getLogger(__name__)
|