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 ADDED
@@ -0,0 +1,8 @@
1
+ from .multi_provider_wrapper import RateLimits, MultiProviderWrapper, RotatingAsyncOpenAIClient, RotatingOpenAIClient
2
+ __all__ = [
3
+ "RateLimits",
4
+ "RotatingKeyManager",
5
+ "MultiProviderWrapper",
6
+ "RotatingAsyncOpenAIClient",
7
+ "RotatingOpenAIClient",
8
+ ]
@@ -0,0 +1,3 @@
1
+ from .openai_adapter import RotatingOpenAIClient, RotatingAsyncOpenAIClient
2
+
3
+ __all__ = ["RotatingOpenAIClient", "RotatingAsyncOpenAIClient"]
@@ -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()
@@ -0,0 +1,5 @@
1
+ from enum import Enum
2
+
3
+ class RateLimitStrategy(Enum):
4
+ PER_MODEL = "per_model" # Cerebras, Groq, Gemini
5
+ GLOBAL = "global" # OpenRouter (Shared limits across all models)
@@ -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__)