Coverage for src/chat_limiter/limiter.py: 73%

477 statements  

« prev     ^ index     » next       coverage.py v7.9.2, created at 2025-12-09 08:16 -0500

1""" 

2Core rate limiter implementation using PyrateLimiter. 

3""" 

4 

5import asyncio 

6from collections.abc import AsyncIterator, Iterator 

7from contextlib import asynccontextmanager, contextmanager 

8from dataclasses import dataclass, field 

9import logging 

10import time 

11from typing import Any 

12 

13import httpx 

14from pyrate_limiter import Duration, Limiter, Rate 

15from tenacity import ( 

16 retry, 

17 retry_if_exception_type, 

18 stop_after_attempt, 

19 wait_exponential, 

20) 

21 

22from .adapters import get_adapter 

23from .providers import ( 

24 Provider, 

25 ProviderConfig, 

26 RateLimitInfo, 

27 detect_provider_from_url, 

28 get_provider_config, 

29) 

30from .types import ( 

31 ChatCompletionRequest, 

32 ChatCompletionResponse, 

33 Message, 

34 MessageRole, 

35 detect_provider_from_model, 

36) 

37from .models import detect_provider_from_model_sync 

38 

39logger = logging.getLogger(__name__) 

40 

41 

42@dataclass 

43class LimiterState: 

44 """Current state of the rate limiter.""" 

45 

46 # Current limits (None if not yet discovered) 

47 request_limit: int | None = None 

48 token_limit: int | None = None 

49 

50 # Usage tracking 

51 requests_used: int = 0 

52 tokens_used: int = 0 

53 

54 # Timing 

55 last_request_time: float = field(default_factory=time.time) 

56 last_limit_update: float = field(default_factory=time.time) 

57 

58 # Rate limit info from last response 

59 last_rate_limit_info: RateLimitInfo | None = None 

60 

61 # Adaptive behavior 

62 consecutive_rate_limit_errors: int = 0 

63 adaptive_backoff_factor: float = 1.0 

64 

65 

66class ChatLimiter: 

67 """ 

68 A Pythonic rate limiter for API calls supporting OpenAI, Anthropic, and OpenRouter. 

69 

70 Features: 

71 - Automatic rate limit discovery and adaptation 

72 - Sync and async support with context managers 

73 - Intelligent retry logic with exponential backoff 

74 - Token and request rate limiting 

75 - Provider-specific optimizations 

76 

77 Example: 

78 # High-level interface (recommended) 

79 async with ChatLimiter.for_model("gpt-4o", api_key="sk-...") as limiter: 

80 response = await limiter.chat_completion( 

81 model="gpt-4o", 

82 messages=[Message(role=MessageRole.USER, content="Hello!")] 

83 ) 

84 

85 # Low-level interface (for advanced users) 

86 async with ChatLimiter(provider=Provider.OPENAI, api_key="sk-...") as limiter: 

87 response = await limiter.request("POST", "/chat/completions", json=data) 

88 """ 

89 

90 def __init__( 

91 self, 

92 provider: Provider | None = None, 

93 api_key: str | None = None, 

94 base_url: str | None = None, 

95 config: ProviderConfig | None = None, 

96 http_client: httpx.AsyncClient | None = None, 

97 sync_http_client: httpx.Client | None = None, 

98 enable_adaptive_limits: bool = True, 

99 enable_token_estimation: bool = True, 

100 request_limit: int | None = None, 

101 token_limit: int | None = None, 

102 max_retries: int | None = None, 

103 base_backoff: float | None = None, 

104 timeout: float | None = None, 

105 **kwargs: Any, 

106 ): 

107 """ 

108 Initialize the ChatLimiter. 

109 

110 Args: 

111 provider: The API provider (OpenAI, Anthropic, OpenRouter) 

112 api_key: API key for authentication 

113 base_url: Base URL for API requests 

114 config: Custom provider configuration 

115 http_client: Custom async HTTP client 

116 sync_http_client: Custom sync HTTP client 

117 enable_adaptive_limits: Enable adaptive rate limit adjustment 

118 enable_token_estimation: Enable token usage estimation 

119 request_limit: Override request limit (if not provided, must be discovered from API) 

120 token_limit: Override token limit (if not provided, must be discovered from API) 

121 max_retries: Override max retries (defaults to 3 if not provided) 

122 base_backoff: Override base backoff (defaults to 1.0 if not provided) 

123 timeout: HTTP request timeout in seconds (defaults to 120.0 for better reliability) 

124 **kwargs: Additional arguments passed to HTTP clients 

125 """ 

126 # Determine provider and config 

127 if config: 

128 self.config = config 

129 self.provider = config.provider 

130 elif provider: 

131 self.provider = provider 

132 self.config = get_provider_config(provider) 

133 elif base_url: 

134 detected_provider = detect_provider_from_url(base_url) 

135 if detected_provider: 

136 self.provider = detected_provider 

137 self.config = get_provider_config(detected_provider) 

138 else: 

139 raise ValueError(f"Could not detect provider from URL: {base_url}") 

140 else: 

141 raise ValueError("Must provide either provider, config, or base_url") 

142 

143 # Override base_url if provided 

144 if base_url: 

145 self.config.base_url = base_url 

146 

147 # Store configuration 

148 self.api_key = api_key 

149 self.enable_adaptive_limits = enable_adaptive_limits 

150 self.enable_token_estimation = enable_token_estimation 

151 

152 # Store user-provided overrides 

153 self._user_request_limit = request_limit 

154 self._user_token_limit = token_limit 

155 self._user_max_retries = max_retries or 3 # Default to 3 if not provided 

156 self._user_base_backoff = base_backoff or 1.0 # Default to 1.0 if not provided 

157 self._user_timeout = ( 

158 timeout or 120.0 

159 ) # Default to 120 seconds for better reliability 

160 

161 # Determine initial limits (user override, config default, or None for discovery) 

162 initial_request_limit = ( 

163 request_limit or self.config.default_request_limit or None 

164 ) 

165 initial_token_limit = token_limit or self.config.default_token_limit or None 

166 

167 # Initialize state - will be None if no defaults and no discovery yet 

168 self.state = LimiterState( 

169 request_limit=initial_request_limit, 

170 token_limit=initial_token_limit, 

171 ) 

172 

173 # Flag to track if we need to discover limits 

174 self._limits_discovered = ( 

175 initial_request_limit is not None and initial_token_limit is not None 

176 ) 

177 

178 # Initialize HTTP clients 

179 self._init_http_clients(http_client, sync_http_client, **kwargs) 

180 

181 # Initialize rate limiters 

182 self._init_rate_limiters() 

183 

184 # Context manager state 

185 self._async_context_active = False 

186 self._sync_context_active = False 

187 

188 # Logging configuration 

189 self._print_rate_limit_info = False 

190 self._print_request_initiation = False 

191 

192 @classmethod 

193 def for_model( 

194 cls, 

195 model: str, 

196 api_key: str | None = None, 

197 provider: str | Provider | None = None, 

198 use_dynamic_discovery: bool = True, 

199 request_limit: int | None = None, 

200 token_limit: int | None = None, 

201 max_retries: int | None = None, 

202 base_backoff: float | None = None, 

203 timeout: float | None = None, 

204 **kwargs: Any, 

205 ) -> "ChatLimiter": 

206 """ 

207 Create a ChatLimiter instance automatically detecting the provider from the model name. 

208 

209 Args: 

210 model: The model name (e.g., "gpt-4o", "claude-3-sonnet-20240229") 

211 api_key: API key for the provider. If None, will be read from environment variables 

212 (OPENAI_API_KEY, ANTHROPIC_API_KEY, OPENROUTER_API_KEY) 

213 provider: Override provider detection. Can be "openai", "anthropic", "openrouter", 

214 or Provider enum. If None, will be auto-detected from model name 

215 use_dynamic_discovery: Whether to query live APIs for model availability (default: True). 

216 Requires appropriate API keys to be available. Falls back to 

217 hardcoded model lists when disabled or when API calls fail. 

218 **kwargs: Additional arguments passed to ChatLimiter 

219 

220 Returns: 

221 Configured ChatLimiter instance 

222 

223 Raises: 

224 ValueError: If provider cannot be determined from model name or API key not found 

225 

226 Example: 

227 # Auto-detect provider with dynamic discovery (default behavior) 

228 async with ChatLimiter.for_model("gpt-4o") as limiter: 

229 response = await limiter.simple_chat("gpt-4o", "Hello!") 

230 

231 # Override provider detection 

232 async with ChatLimiter.for_model("custom-model", provider="openai") as limiter: 

233 response = await limiter.simple_chat("custom-model", "Hello!") 

234 

235 # Disable dynamic discovery to use only hardcoded model lists 

236 async with ChatLimiter.for_model("gpt-4o", use_dynamic_discovery=False) as limiter: 

237 response = await limiter.simple_chat("gpt-4o", "Hello!") 

238 """ 

239 import os 

240 

241 # Determine provider 

242 if provider is not None: 

243 # Use provided provider 

244 if isinstance(provider, str): 

245 provider_enum = Provider(provider) 

246 else: 

247 provider_enum = provider 

248 provider_name = provider_enum.value 

249 else: 

250 # Auto-detect from model name 

251 # If dynamic discovery is requested, we need to collect API keys first 

252 api_keys_for_discovery = {} 

253 if use_dynamic_discovery: 

254 # Collect available API keys from environment 

255 env_var_map = { 

256 "openai": "OPENAI_API_KEY", 

257 "anthropic": "ANTHROPIC_API_KEY", 

258 "openrouter": "OPENROUTER_API_KEY", 

259 } 

260 

261 for provider_key, env_var in env_var_map.items(): 

262 key_value = os.getenv(env_var) 

263 if key_value: 

264 api_keys_for_discovery[provider_key] = key_value 

265 

266 # Try dynamic discovery first to get more detailed information 

267 discovery_result = None 

268 detected_provider = detect_provider_from_model( 

269 model, use_dynamic_discovery, api_keys_for_discovery 

270 ) 

271 

272 if not detected_provider: 

273 discovery_msg = ( 

274 " with dynamic API discovery" if use_dynamic_discovery else "" 

275 ) 

276 error_msg = f"Could not determine provider from model '{model}'{discovery_msg}. " 

277 

278 # Add detailed information about available models if we have discovery results 

279 if discovery_result and discovery_result.get_total_models_found() > 0: 

280 error_msg += f"\n\nFound {discovery_result.get_total_models_found()} models across providers:\n" 

281 for ( 

282 provider_name, 

283 models, 

284 ) in discovery_result.get_all_models().items(): 

285 error_msg += f" {provider_name}: {len(models)} models\n" 

286 for example in sorted(list(models)): 

287 error_msg += f" - {example}\n" 

288 error_msg += "\nPlease check the model name or specify the provider explicitly using the 'provider' parameter." 

289 else: 

290 error_msg += "Please specify the provider explicitly using the 'provider' parameter." 

291 

292 # Add information about discovery errors if any 

293 if discovery_result and discovery_result.errors: 

294 error_msg += "\n\nDiscovery errors encountered:\n" 

295 for provider_name, error in discovery_result.errors.items(): 

296 error_msg += f" {provider_name}: {error}\n" 

297 

298 raise ValueError(error_msg) 

299 assert detected_provider is not None # Help MyPy understand type narrowing 

300 provider_name = detected_provider 

301 provider_enum = Provider(provider_name) 

302 

303 # If a provider prefix was used, print discovered models when the base 

304 # model is not present in the provider's discovered set (diagnostics). 

305 if use_dynamic_discovery and "/" in model and api_keys_for_discovery: 

306 parts = model.split("/", 1) 

307 if len(parts) == 2: 

308 provider_prefix, base_model = parts 

309 discovery_result = detect_provider_from_model_sync(model, api_keys_for_discovery) 

310 

311 if provider_name == "openai": 

312 if discovery_result.openai_models is None: 

313 # Print discovery errors if available 

314 if discovery_result.errors: 

315 print("OpenAI discovery: no model list available. Errors:") 

316 for k, v in discovery_result.errors.items(): 

317 print(f" - {k}: {v}") 

318 elif base_model not in discovery_result.openai_models: 

319 print( 

320 f"OpenAI discovery summary: found={len(discovery_result.openai_models)} models, " 

321 f"contains('{base_model}')=False" 

322 ) 

323 models = sorted(list(discovery_result.openai_models)) 

324 print(f"OpenAI discovery: base model '{base_model}' not found. Listing {len(models)} discovered models:") 

325 for example in models[:20]: 

326 print(f" - {example}") 

327 else: 

328 print( 

329 f"OpenAI discovery summary: found={len(discovery_result.openai_models)} models, " 

330 f"contains('{base_model}')=True" 

331 ) 

332 

333 elif provider_name == "anthropic": 

334 if discovery_result.anthropic_models is None: 

335 if discovery_result.errors: 

336 print("Anthropic discovery: no model list available. Errors:") 

337 for k, v in discovery_result.errors.items(): 

338 print(f" - {k}: {v}") 

339 elif base_model not in discovery_result.anthropic_models: 

340 print( 

341 f"Anthropic discovery summary: found={len(discovery_result.anthropic_models)} models, " 

342 f"contains('{base_model}')=False" 

343 ) 

344 models = sorted(list(discovery_result.anthropic_models)) 

345 print(f"Anthropic discovery: base model '{base_model}' not found. Listing {len(models)} discovered models:") 

346 for example in models[:20]: 

347 print(f" - {example}") 

348 else: 

349 print( 

350 f"Anthropic discovery summary: found={len(discovery_result.anthropic_models)} models, " 

351 f"contains('{base_model}')=True" 

352 ) 

353 

354 elif provider_name == "openrouter": 

355 # For OpenRouter, the model string includes the provider prefix 

356 if discovery_result.openrouter_models is None: 

357 if discovery_result.errors: 

358 print("OpenRouter discovery: no model list available. Errors:") 

359 for k, v in discovery_result.errors.items(): 

360 print(f" - {k}: {v}") 

361 elif model not in discovery_result.openrouter_models: 

362 print( 

363 f"OpenRouter discovery summary: found={len(discovery_result.openrouter_models)} models, " 

364 f"contains('{model}')=False" 

365 ) 

366 models = sorted(list(discovery_result.openrouter_models)) 

367 print(f"OpenRouter discovery: model '{model}' not found. Listing {len(models)} discovered models:") 

368 for example in models[:20]: 

369 print(f" - {example}") 

370 else: 

371 print( 

372 f"OpenRouter discovery summary: found={len(discovery_result.openrouter_models)} models, " 

373 f"contains('{model}')=True" 

374 ) 

375 

376 # Determine API key 

377 if api_key is None: 

378 # Try to get from environment variables 

379 env_var_map = { 

380 "openai": "OPENAI_API_KEY", 

381 "anthropic": "ANTHROPIC_API_KEY", 

382 "openrouter": "OPENROUTER_API_KEY", 

383 } 

384 

385 env_var_name: str | None = env_var_map.get(provider_name) 

386 if env_var_name: 

387 api_key = os.getenv(env_var_name) 

388 if not api_key: 

389 raise ValueError( 

390 f"API key not provided and {env_var_name} environment variable not set. " 

391 f"Please provide api_key parameter or set {env_var_name} environment variable." 

392 ) 

393 else: 

394 raise ValueError( 

395 f"Unknown provider '{provider_name}'. Cannot determine environment variable for API key." 

396 ) 

397 

398 return cls( 

399 provider=provider_enum, 

400 api_key=api_key, 

401 request_limit=request_limit, 

402 token_limit=token_limit, 

403 max_retries=max_retries, 

404 base_backoff=base_backoff, 

405 timeout=timeout, 

406 **kwargs, 

407 ) 

408 

409 def _init_http_clients( 

410 self, 

411 http_client: httpx.AsyncClient | None, 

412 sync_http_client: httpx.Client | None, 

413 **kwargs: Any, 

414 ) -> None: 

415 """Initialize HTTP clients with proper headers.""" 

416 # Prepare headers 

417 headers = { 

418 "User-Agent": f"chat-limiter/0.1.0 ({self.provider.value})", 

419 } 

420 

421 # Add provider-specific headers 

422 if self.api_key: 

423 if self.provider == Provider.OPENAI: 

424 headers["Authorization"] = f"Bearer {self.api_key}" 

425 elif self.provider == Provider.ANTHROPIC: 

426 headers["x-api-key"] = self.api_key 

427 headers["anthropic-version"] = "2023-06-01" 

428 elif self.provider == Provider.OPENROUTER: 

429 headers["Authorization"] = f"Bearer {self.api_key}" 

430 headers["HTTP-Referer"] = "https://github.com/your-repo/chat-limiter" 

431 

432 # Merge with user-provided headers 

433 if "headers" in kwargs: 

434 headers.update(kwargs["headers"]) 

435 kwargs["headers"] = headers 

436 

437 # Initialize clients 

438 if http_client: 

439 self.async_client = http_client 

440 else: 

441 self.async_client = httpx.AsyncClient( 

442 base_url=self.config.base_url, 

443 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout 

444 **kwargs, 

445 ) 

446 

447 if sync_http_client: 

448 self.sync_client = sync_http_client 

449 else: 

450 self.sync_client = httpx.Client( 

451 base_url=self.config.base_url, 

452 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout 

453 **kwargs, 

454 ) 

455 

456 def _init_rate_limiters(self) -> None: 

457 """Initialize PyrateLimiter instances.""" 

458 # Only initialize if we have limits 

459 if self.state.request_limit is None or self.state.token_limit is None: 

460 # Cannot initialize rate limiters without limits 

461 # This will be called again after limits are discovered 

462 self.request_limiter = None 

463 self.token_limiter = None 

464 self._effective_request_limit = None 

465 self._effective_token_limit = None 

466 return 

467 

468 # Dispose existing limiters to prevent background leaker thread accumulation 

469 self._dispose_rate_limiters() 

470 

471 # Calculate effective limits with buffer 

472 effective_request_limit = int( 

473 self.state.request_limit * self.config.request_buffer_ratio 

474 ) 

475 effective_token_limit = int( 

476 self.state.token_limit * self.config.token_buffer_ratio 

477 ) 

478 

479 # Request rate limiter 

480 self.request_limiter = Limiter( 

481 Rate( 

482 effective_request_limit, 

483 Duration.MINUTE, 

484 ) 

485 ) 

486 

487 # Token rate limiter 

488 self.token_limiter = Limiter( 

489 Rate( 

490 effective_token_limit, 

491 Duration.MINUTE, 

492 ) 

493 ) 

494 

495 # Store effective limits for logging 

496 self._effective_request_limit = effective_request_limit 

497 self._effective_token_limit = effective_token_limit 

498 

499 async def __aenter__(self) -> "ChatLimiter": 

500 """Async context manager entry.""" 

501 if self._async_context_active: 

502 raise RuntimeError( 

503 "ChatLimiter is already active as an async context manager" 

504 ) 

505 

506 self._async_context_active = True 

507 

508 # Discover rate limits if supported 

509 if self.config.supports_dynamic_limits: 

510 await self._discover_rate_limits() 

511 

512 # Print rate limit information if enabled 

513 if self._print_rate_limit_info: 

514 self._print_rate_limit_info_details() 

515 

516 return self 

517 

518 async def __aexit__( 

519 self, 

520 exc_type: type[BaseException] | None, 

521 exc_val: BaseException | None, 

522 exc_tb: object, 

523 ) -> None: 

524 """Async context manager exit.""" 

525 self._async_context_active = False 

526 self._dispose_rate_limiters() 

527 await self.async_client.aclose() 

528 

529 def __enter__(self) -> "ChatLimiter": 

530 """Sync context manager entry.""" 

531 if self._sync_context_active: 

532 raise RuntimeError( 

533 "ChatLimiter is already active as a sync context manager" 

534 ) 

535 

536 self._sync_context_active = True 

537 

538 # Discover rate limits if supported 

539 if self.config.supports_dynamic_limits: 

540 self._discover_rate_limits_sync() 

541 

542 # Print rate limit information if enabled 

543 if self._print_rate_limit_info: 

544 self._print_rate_limit_info_details() 

545 

546 return self 

547 

548 def __exit__( 

549 self, 

550 exc_type: type[BaseException] | None, 

551 exc_val: BaseException | None, 

552 exc_tb: object, 

553 ) -> None: 

554 """Sync context manager exit.""" 

555 self._sync_context_active = False 

556 self._dispose_rate_limiters() 

557 self.sync_client.close() 

558 

559 def _dispose_rate_limiters(self) -> None: 

560 """Dispose buckets from existing pyrate limiters to stop leaker threads.""" 

561 rl = getattr(self, "request_limiter", None) 

562 if rl is not None: 

563 for bucket in rl.buckets(): 

564 rl.dispose(bucket) 

565 self.request_limiter = None 

566 self._effective_request_limit = None 

567 

568 tl = getattr(self, "token_limiter", None) 

569 if tl is not None: 

570 for bucket in tl.buckets(): 

571 tl.dispose(bucket) 

572 self.token_limiter = None 

573 self._effective_token_limit = None 

574 

575 async def _discover_rate_limits(self) -> None: 

576 """Discover current rate limits from the API.""" 

577 try: 

578 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint: 

579 # OpenRouter uses a special auth endpoint 

580 response = await self.async_client.get(self.config.auth_endpoint) 

581 response.raise_for_status() 

582 

583 data = response.json() 

584 # Update limits based on response 

585 # This is a simplified version - actual implementation would parse the response 

586 logger.info(f"Discovered OpenRouter limits: {data}") 

587 

588 else: 

589 # For other providers, we'll discover limits on first request 

590 if self._print_rate_limit_info: 

591 print( 

592 f"Rate limit discovery will happen on first request for {self.provider.value}" 

593 ) 

594 logger.info( 

595 f"Rate limit discovery will happen on first request for {self.provider.value}" 

596 ) 

597 

598 except Exception as e: 

599 logger.warning(f"Failed to discover rate limits: {e}") 

600 

601 def _discover_rate_limits_sync(self) -> None: 

602 """Sync version of rate limit discovery.""" 

603 try: 

604 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint: 

605 response = self.sync_client.get(self.config.auth_endpoint) 

606 response.raise_for_status() 

607 

608 data = response.json() 

609 logger.info(f"Discovered OpenRouter limits: {data}") 

610 else: 

611 logger.info( 

612 f"Rate limit discovery will happen on first request for {self.provider.value}" 

613 ) 

614 

615 except Exception as e: 

616 logger.warning(f"Failed to discover rate limits: {e}") 

617 

618 def _update_rate_limits(self, rate_limit_info: RateLimitInfo) -> None: 

619 """Update rate limits based on response headers.""" 

620 updated = False 

621 was_uninitialized = ( 

622 self.state.request_limit is None or self.state.token_limit is None 

623 ) 

624 

625 # Update request limits 

626 if ( 

627 rate_limit_info.requests_limit 

628 and rate_limit_info.requests_limit != self.state.request_limit 

629 ): 

630 old_limit = self.state.request_limit 

631 self.state.request_limit = rate_limit_info.requests_limit 

632 updated = True 

633 if was_uninitialized: 

634 message = ( 

635 f"Discovered request limit: {self.state.request_limit} req/min" 

636 ) 

637 if self._print_rate_limit_info: 

638 print(message) 

639 logger.info(message) 

640 else: 

641 message = f"Updated request limit: {old_limit} -> {self.state.request_limit} req/min" 

642 if self._print_rate_limit_info: 

643 print(message) 

644 logger.info(message) 

645 

646 # Update token limits 

647 if ( 

648 rate_limit_info.tokens_limit 

649 and rate_limit_info.tokens_limit != self.state.token_limit 

650 ): 

651 old_limit = self.state.token_limit 

652 self.state.token_limit = rate_limit_info.tokens_limit 

653 updated = True 

654 if was_uninitialized: 

655 message = f"Discovered token limit: {self.state.token_limit} tokens/min" 

656 if self._print_rate_limit_info: 

657 print(message) 

658 logger.info(message) 

659 else: 

660 message = f"Updated token limit: {old_limit} -> {self.state.token_limit} tokens/min" 

661 if self._print_rate_limit_info: 

662 print(message) 

663 logger.info(message) 

664 

665 if updated: 

666 # Reinitialize rate limiters with new limits 

667 self._init_rate_limiters() 

668 

669 # Update limits_discovered flag if both limits are now available 

670 if ( 

671 self.state.request_limit is not None 

672 and self.state.token_limit is not None 

673 ): 

674 self._limits_discovered = True 

675 

676 if was_uninitialized: 

677 message = "Rate limiters initialized after discovery" 

678 if self._print_rate_limit_info: 

679 print(message) 

680 # Print updated rate limit info after discovery 

681 self._print_rate_limit_info_details() 

682 logger.info(message) 

683 

684 # Store the rate limit info 

685 self.state.last_rate_limit_info = rate_limit_info 

686 self.state.last_limit_update = time.time() 

687 

688 def _estimate_tokens(self, request_data: dict[str, Any]) -> int: 

689 """Estimate token usage from request data.""" 

690 if not self.enable_token_estimation: 

691 return 0 

692 

693 # Simple token estimation 

694 # This is a placeholder - real implementation would use tiktoken or similar 

695 if "messages" in request_data: 

696 text = "" 

697 for message in request_data["messages"]: 

698 if isinstance(message, dict) and "content" in message: 

699 text += str(message["content"]) 

700 

701 # Rough estimation: 1 token ≈ 4 characters 

702 return len(text) // 4 

703 

704 return 0 

705 

706 @asynccontextmanager 

707 async def _acquire_rate_limits( 

708 self, estimated_tokens: int = 0 

709 ) -> AsyncIterator[None]: 

710 """Acquire rate limits before making a request.""" 

711 # Check if rate limiters are initialized 

712 if self.request_limiter is None or self.token_limiter is None: 

713 # Limits not yet discovered - this request will help discover them 

714 logger.info( 

715 "Rate limits not yet discovered, proceeding without rate limiting for discovery" 

716 ) 

717 else: 

718 # Wait for request rate limit 

719 await asyncio.to_thread(self.request_limiter.try_acquire, "request") 

720 

721 # Wait for token rate limit if we have token estimation and limiters are initialized 

722 if ( 

723 estimated_tokens > 0 

724 and self.token_limiter is not None 

725 and self._effective_token_limit is not None 

726 ): 

727 # Check if request is too large for bucket capacity 

728 if estimated_tokens > self._effective_token_limit: 

729 # Log warning for large requests 

730 logger.warning( 

731 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity " 

732 f"of {self._effective_token_limit} tokens. This may cause delays." 

733 ) 

734 # For very large requests, we'll split the acquisition 

735 # Acquire tokens in chunks to avoid bucket overflow 

736 remaining_tokens = estimated_tokens 

737 while remaining_tokens > 0: 

738 chunk_size = min( 

739 remaining_tokens, self._effective_token_limit // 2 

740 ) 

741 await asyncio.to_thread( 

742 self.token_limiter.try_acquire, "token", chunk_size 

743 ) 

744 remaining_tokens -= chunk_size 

745 if remaining_tokens > 0: 

746 # Brief pause to let bucket refill 

747 await asyncio.sleep(0.1) 

748 else: 

749 # Normal acquisition for smaller requests 

750 await asyncio.to_thread( 

751 self.token_limiter.try_acquire, "token", estimated_tokens 

752 ) 

753 

754 try: 

755 yield 

756 finally: 

757 # Update usage tracking 

758 self.state.requests_used += 1 

759 self.state.tokens_used += estimated_tokens 

760 self.state.last_request_time = time.time() 

761 

762 @contextmanager 

763 def _acquire_rate_limits_sync(self, estimated_tokens: int = 0) -> Iterator[None]: 

764 """Sync version of rate limit acquisition.""" 

765 # Check if rate limiters are initialized 

766 if self.request_limiter is None or self.token_limiter is None: 

767 # Limits not yet discovered - this request will help discover them 

768 logger.info( 

769 "Rate limits not yet discovered, proceeding without rate limiting for discovery" 

770 ) 

771 else: 

772 # Wait for request rate limit 

773 self.request_limiter.try_acquire("request") 

774 

775 # Wait for token rate limit if we have token estimation and limiters are initialized 

776 if ( 

777 estimated_tokens > 0 

778 and self.token_limiter is not None 

779 and self._effective_token_limit is not None 

780 ): 

781 # Check if request is too large for bucket capacity 

782 if estimated_tokens > self._effective_token_limit: 

783 # Log warning for large requests 

784 logger.warning( 

785 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity " 

786 f"of {self._effective_token_limit} tokens. This may cause delays." 

787 ) 

788 # For very large requests, we'll split the acquisition 

789 # Acquire tokens in chunks to avoid bucket overflow 

790 remaining_tokens = estimated_tokens 

791 while remaining_tokens > 0: 

792 chunk_size = min( 

793 remaining_tokens, self._effective_token_limit // 2 

794 ) 

795 self.token_limiter.try_acquire("token", chunk_size) 

796 remaining_tokens -= chunk_size 

797 if remaining_tokens > 0: 

798 # Brief pause to let bucket refill 

799 time.sleep(0.1) 

800 else: 

801 # Normal acquisition for smaller requests 

802 self.token_limiter.try_acquire("token", estimated_tokens) 

803 

804 try: 

805 yield 

806 finally: 

807 # Update usage tracking 

808 self.state.requests_used += 1 

809 self.state.tokens_used += estimated_tokens 

810 self.state.last_request_time = time.time() 

811 

812 def _get_retry_decorator(self) -> Any: 

813 """Get retry decorator with user-configured parameters.""" 

814 return retry( 

815 stop=stop_after_attempt(self._user_max_retries), 

816 wait=wait_exponential(multiplier=self._user_base_backoff, min=1, max=60), 

817 retry=retry_if_exception_type( 

818 ( 

819 httpx.HTTPStatusError, 

820 httpx.RequestError, 

821 httpx.ReadTimeout, 

822 httpx.ConnectTimeout, 

823 ) 

824 ), 

825 ) 

826 

827 def get_current_limits(self) -> dict[str, Any]: 

828 """Get current rate limit information.""" 

829 return { 

830 "provider": self.provider.value, 

831 "request_limit": self.state.request_limit, 

832 "token_limit": self.state.token_limit, 

833 "requests_used": self.state.requests_used, 

834 "tokens_used": self.state.tokens_used, 

835 "last_request_time": self.state.last_request_time, 

836 "last_limit_update": self.state.last_limit_update, 

837 "consecutive_rate_limit_errors": self.state.consecutive_rate_limit_errors, 

838 } 

839 

840 def reset_usage_tracking(self) -> None: 

841 """Reset usage tracking counters.""" 

842 self.state.requests_used = 0 

843 self.state.tokens_used = 0 

844 self.state.consecutive_rate_limit_errors = 0 

845 

846 # High-level chat completion methods 

847 

848 async def chat_completion( 

849 self, 

850 model: str, 

851 messages: list[Message], 

852 max_tokens: int | None = None, 

853 temperature: float | None = None, 

854 top_p: float | None = None, 

855 stop: str | list[str] | None = None, 

856 stream: bool = False, 

857 **kwargs: Any, 

858 ) -> ChatCompletionResponse: 

859 """ 

860 Make a high-level chat completion request. 

861 

862 Args: 

863 model: The model to use for completion 

864 messages: List of messages in the conversation 

865 max_tokens: Maximum tokens to generate 

866 temperature: Sampling temperature 

867 top_p: Top-p sampling parameter 

868 stop: Stop sequences 

869 stream: Whether to stream the response 

870 **kwargs: Additional provider-specific parameters 

871 

872 Returns: 

873 ChatCompletionResponse with the completion result 

874 

875 Raises: 

876 ValueError: If provider cannot be determined from model 

877 httpx.HTTPStatusError: For HTTP error responses 

878 httpx.RequestError: For request errors 

879 """ 

880 # Create request object 

881 request = ChatCompletionRequest( 

882 model=model, 

883 messages=messages, 

884 max_tokens=max_tokens, 

885 temperature=temperature, 

886 top_p=top_p, 

887 stop=stop, 

888 stream=stream, 

889 **kwargs, 

890 ) 

891 

892 # Get the appropriate adapter 

893 adapter = get_adapter(self.provider) 

894 

895 # Format the request for the provider 

896 formatted_request = adapter.format_request(request) 

897 

898 # Make the HTTP request with rate limiting 

899 try: 

900 # Print request initiation if enabled 

901 if self._print_request_initiation: 

902 print(f"Sending request for model {model} (attempt 1)") 

903 

904 # Estimate tokens 

905 estimated_tokens = self._estimate_tokens(formatted_request) 

906 

907 # Choose HTTP client: reuse within async context, per-call otherwise 

908 client = None 

909 close_client_after_use = False 

910 if self._async_context_active: 

911 client = self.async_client 

912 else: 

913 client = httpx.AsyncClient( 

914 base_url=self.config.base_url, 

915 timeout=httpx.Timeout(self._user_timeout), 

916 headers=dict(self.async_client.headers), 

917 ) 

918 close_client_after_use = True 

919 

920 try: 

921 # Acquire rate limits 

922 async with self._acquire_rate_limits(estimated_tokens): 

923 # Make the request 

924 response = await client.request( 

925 "POST", adapter.get_endpoint(), json=formatted_request 

926 ) 

927 

928 # Extract rate limit info 

929 from .providers import extract_rate_limit_info 

930 rate_limit_info = extract_rate_limit_info( 

931 dict(response.headers), self.config 

932 ) 

933 

934 # Update our rate limits 

935 if self.enable_adaptive_limits: 

936 self._update_rate_limits(rate_limit_info) 

937 

938 # Handle rate limit errors 

939 if response.status_code == 429: 

940 self.state.consecutive_rate_limit_errors += 1 

941 if rate_limit_info.retry_after: 

942 import asyncio 

943 await asyncio.sleep(rate_limit_info.retry_after) 

944 else: 

945 # Exponential backoff 

946 import asyncio 

947 backoff = self.config.base_backoff * ( 

948 2**self.state.consecutive_rate_limit_errors 

949 ) 

950 await asyncio.sleep(min(backoff, self.config.max_backoff)) 

951 

952 response.raise_for_status() 

953 else: 

954 # Reset consecutive errors on success 

955 self.state.consecutive_rate_limit_errors = 0 

956 

957 # Raise for all non-2xx responses (do not silently succeed) 

958 if response.status_code != 429: 

959 response.raise_for_status() 

960 

961 # Parse the response 

962 response_data = response.json() 

963 return adapter.parse_response(response_data, request) 

964 finally: 

965 if close_client_after_use: 

966 await client.aclose() 

967 except httpx.HTTPStatusError as e: 

968 body_text = "" 

969 try: 

970 body_text = e.response.text if e.response is not None else "" 

971 except Exception: 

972 body_text = "" 

973 error_response = ChatCompletionResponse( 

974 id="error", 

975 model=request.model, 

976 success=False, 

977 error_message=f"{str(e)} | body={body_text}", 

978 choices=[], 

979 usage=None, 

980 created=None, 

981 ) 

982 return error_response 

983 except Exception as e: 

984 error_response = ChatCompletionResponse( 

985 id="error", 

986 model=request.model, 

987 success=False, 

988 error_message=str(e), 

989 choices=[], 

990 usage=None, 

991 created=None, 

992 ) 

993 return error_response 

994 

995 def chat_completion_sync( 

996 self, 

997 model: str, 

998 messages: list[Message], 

999 max_tokens: int | None = None, 

1000 temperature: float | None = None, 

1001 top_p: float | None = None, 

1002 stop: str | list[str] | None = None, 

1003 stream: bool = False, 

1004 **kwargs: Any, 

1005 ) -> ChatCompletionResponse: 

1006 """ 

1007 Make a synchronous high-level chat completion request. 

1008 

1009 Args: 

1010 model: The model to use for completion 

1011 messages: List of messages in the conversation 

1012 max_tokens: Maximum tokens to generate 

1013 temperature: Sampling temperature 

1014 top_p: Top-p sampling parameter 

1015 stop: Stop sequences 

1016 stream: Whether to stream the response 

1017 **kwargs: Additional provider-specific parameters 

1018 

1019 Returns: 

1020 ChatCompletionResponse with the completion result 

1021 

1022 Raises: 

1023 ValueError: If provider cannot be determined from model 

1024 httpx.HTTPStatusError: For HTTP error responses 

1025 httpx.RequestError: For request errors 

1026 """ 

1027 # Create request object 

1028 request = ChatCompletionRequest( 

1029 model=model, 

1030 messages=messages, 

1031 max_tokens=max_tokens, 

1032 temperature=temperature, 

1033 top_p=top_p, 

1034 stop=stop, 

1035 stream=stream, 

1036 **kwargs, 

1037 ) 

1038 

1039 # Get the appropriate adapter 

1040 adapter = get_adapter(self.provider) 

1041 

1042 # Format the request for the provider 

1043 formatted_request = adapter.format_request(request) 

1044 

1045 # Make the HTTP request with rate limiting 

1046 try: 

1047 # Print request initiation if enabled 

1048 if self._print_request_initiation: 

1049 print(f"Sending request for model {model} (attempt 1)") 

1050 

1051 # Estimate tokens 

1052 estimated_tokens = self._estimate_tokens(formatted_request) 

1053 

1054 # Choose HTTP client: reuse within sync context, per-call otherwise 

1055 client = None 

1056 close_client_after_use = False 

1057 if self._sync_context_active: 

1058 client = self.sync_client 

1059 else: 

1060 client = httpx.Client( 

1061 base_url=self.config.base_url, 

1062 timeout=httpx.Timeout(self._user_timeout), 

1063 headers=dict(self.sync_client.headers), 

1064 ) 

1065 close_client_after_use = True 

1066 

1067 try: 

1068 # Acquire rate limits 

1069 with self._acquire_rate_limits_sync(estimated_tokens): 

1070 # Make the request 

1071 response = client.request( 

1072 "POST", adapter.get_endpoint(), json=formatted_request 

1073 ) 

1074 

1075 # Extract rate limit info 

1076 from .providers import extract_rate_limit_info 

1077 rate_limit_info = extract_rate_limit_info( 

1078 dict(response.headers), self.config 

1079 ) 

1080 

1081 # Update our rate limits 

1082 if self.enable_adaptive_limits: 

1083 self._update_rate_limits(rate_limit_info) 

1084 

1085 # Handle rate limit errors 

1086 if response.status_code == 429: 

1087 self.state.consecutive_rate_limit_errors += 1 

1088 if rate_limit_info.retry_after: 

1089 import time 

1090 time.sleep(rate_limit_info.retry_after) 

1091 else: 

1092 # Exponential backoff 

1093 import time 

1094 backoff = self.config.base_backoff * ( 

1095 2**self.state.consecutive_rate_limit_errors 

1096 ) 

1097 time.sleep(min(backoff, self.config.max_backoff)) 

1098 

1099 response.raise_for_status() 

1100 else: 

1101 # Reset consecutive errors on success 

1102 self.state.consecutive_rate_limit_errors = 0 

1103 

1104 # Raise for all non-2xx responses (do not silently succeed) 

1105 if response.status_code != 429: 

1106 response.raise_for_status() 

1107 

1108 # Parse the response 

1109 response_data = response.json() 

1110 return adapter.parse_response(response_data, request) 

1111 finally: 

1112 if close_client_after_use: 

1113 client.close() 

1114 except httpx.HTTPStatusError as e: 

1115 body_text = "" 

1116 try: 

1117 body_text = e.response.text if e.response is not None else "" 

1118 except Exception: 

1119 body_text = "" 

1120 error_response = ChatCompletionResponse( 

1121 id="error", 

1122 model=request.model, 

1123 success=False, 

1124 error_message=f"{str(e)} | body={body_text}", 

1125 choices=[], 

1126 usage=None, 

1127 created=None, 

1128 ) 

1129 return error_response 

1130 except Exception as e: 

1131 error_response = ChatCompletionResponse( 

1132 id="error", 

1133 model=request.model, 

1134 success=False, 

1135 error_message=str(e), 

1136 choices=[], 

1137 usage=None, 

1138 created=None, 

1139 ) 

1140 return error_response 

1141 

1142 # Convenience methods for different message types 

1143 

1144 async def simple_chat( 

1145 self, 

1146 model: str, 

1147 prompt: str, 

1148 max_tokens: int | None = None, 

1149 temperature: float | None = None, 

1150 **kwargs: Any, 

1151 ) -> str: 

1152 """ 

1153 Simple chat completion that returns just the text response. 

1154 

1155 Args: 

1156 model: The model to use 

1157 prompt: The user prompt 

1158 max_tokens: Maximum tokens to generate 

1159 temperature: Sampling temperature 

1160 **kwargs: Additional parameters 

1161 

1162 Returns: 

1163 The text response from the model 

1164 """ 

1165 messages = [Message(role=MessageRole.USER, content=prompt)] 

1166 response = await self.chat_completion( 

1167 model=model, 

1168 messages=messages, 

1169 max_tokens=max_tokens, 

1170 temperature=temperature, 

1171 **kwargs, 

1172 ) 

1173 

1174 if response.choices: 

1175 return response.choices[0].message.content 

1176 return "" 

1177 

1178 def simple_chat_sync( 

1179 self, 

1180 model: str, 

1181 prompt: str, 

1182 max_tokens: int | None = None, 

1183 temperature: float | None = None, 

1184 **kwargs: Any, 

1185 ) -> str: 

1186 """ 

1187 Simple synchronous chat completion that returns just the text response. 

1188 

1189 Args: 

1190 model: The model to use 

1191 prompt: The user prompt 

1192 max_tokens: Maximum tokens to generate 

1193 temperature: Sampling temperature 

1194 **kwargs: Additional parameters 

1195 

1196 Returns: 

1197 The text response from the model 

1198 """ 

1199 messages = [Message(role=MessageRole.USER, content=prompt)] 

1200 response = self.chat_completion_sync( 

1201 model=model, 

1202 messages=messages, 

1203 max_tokens=max_tokens, 

1204 temperature=temperature, 

1205 **kwargs, 

1206 ) 

1207 

1208 if response.choices: 

1209 return response.choices[0].message.content 

1210 return "" 

1211 

1212 def set_print_rate_limit_info(self, enabled: bool) -> None: 

1213 """Set whether to print rate limit information.""" 

1214 self._print_rate_limit_info = enabled 

1215 

1216 def set_print_request_initiation(self, enabled: bool) -> None: 

1217 """Set whether to print request initiation messages.""" 

1218 self._print_request_initiation = enabled 

1219 

1220 def _print_rate_limit_info_details(self) -> None: 

1221 """Print current rate limit configuration.""" 

1222 print(f"\n=== Rate Limit Configuration for {self.provider.value.title()} ===") 

1223 print(f"Provider: {self.provider.value}") 

1224 print(f"Base URL: {self.config.base_url}") 

1225 

1226 # Handle None values for limits 

1227 if self.state.request_limit is not None: 

1228 effective_req = self._effective_request_limit or "not calculated" 

1229 print( 

1230 f"Request Limit: {self.state.request_limit}/minute (effective: {effective_req}/minute)" 

1231 ) 

1232 else: 

1233 print("Request Limit: Not yet discovered (will be fetched from API)") 

1234 

1235 if self.state.token_limit is not None: 

1236 effective_tok = self._effective_token_limit or "not calculated" 

1237 print( 

1238 f"Token Limit: {self.state.token_limit}/minute (effective: {effective_tok}/minute)" 

1239 ) 

1240 else: 

1241 print("Token Limit: Not yet discovered (will be fetched from API)") 

1242 

1243 print(f"Request Buffer Ratio: {self.config.request_buffer_ratio}") 

1244 print(f"Token Buffer Ratio: {self.config.token_buffer_ratio}") 

1245 print(f"Adaptive Limits: {self.enable_adaptive_limits}") 

1246 print(f"Token Estimation: {self.enable_token_estimation}") 

1247 print(f"Dynamic Discovery: {self.config.supports_dynamic_limits}") 

1248 print(f"Limits Discovered: {self._limits_discovered}") 

1249 print("=" * 50)