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

413 statements  

« prev     ^ index     » next       coverage.py v7.9.2, created at 2025-10-02 12:23 +0100

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 # Determine API key 

304 if api_key is None: 

305 # Try to get from environment variables 

306 env_var_map = { 

307 "openai": "OPENAI_API_KEY", 

308 "anthropic": "ANTHROPIC_API_KEY", 

309 "openrouter": "OPENROUTER_API_KEY", 

310 } 

311 

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

313 if env_var_name: 

314 api_key = os.getenv(env_var_name) 

315 if not api_key: 

316 raise ValueError( 

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

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

319 ) 

320 else: 

321 raise ValueError( 

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

323 ) 

324 

325 return cls( 

326 provider=provider_enum, 

327 api_key=api_key, 

328 request_limit=request_limit, 

329 token_limit=token_limit, 

330 max_retries=max_retries, 

331 base_backoff=base_backoff, 

332 timeout=timeout, 

333 **kwargs, 

334 ) 

335 

336 def _init_http_clients( 

337 self, 

338 http_client: httpx.AsyncClient | None, 

339 sync_http_client: httpx.Client | None, 

340 **kwargs: Any, 

341 ) -> None: 

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

343 # Prepare headers 

344 headers = { 

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

346 } 

347 

348 # Add provider-specific headers 

349 if self.api_key: 

350 if self.provider == Provider.OPENAI: 

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

352 elif self.provider == Provider.ANTHROPIC: 

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

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

355 elif self.provider == Provider.OPENROUTER: 

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

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

358 

359 # Merge with user-provided headers 

360 if "headers" in kwargs: 

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

362 kwargs["headers"] = headers 

363 

364 # Initialize clients 

365 if http_client: 

366 self.async_client = http_client 

367 else: 

368 self.async_client = httpx.AsyncClient( 

369 base_url=self.config.base_url, 

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

371 **kwargs, 

372 ) 

373 

374 if sync_http_client: 

375 self.sync_client = sync_http_client 

376 else: 

377 self.sync_client = httpx.Client( 

378 base_url=self.config.base_url, 

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

380 **kwargs, 

381 ) 

382 

383 def _init_rate_limiters(self) -> None: 

384 """Initialize PyrateLimiter instances.""" 

385 # Only initialize if we have limits 

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

387 # Cannot initialize rate limiters without limits 

388 # This will be called again after limits are discovered 

389 self.request_limiter = None 

390 self.token_limiter = None 

391 self._effective_request_limit = None 

392 self._effective_token_limit = None 

393 return 

394 

395 # Dispose existing limiters to prevent background leaker thread accumulation 

396 self._dispose_rate_limiters() 

397 

398 # Calculate effective limits with buffer 

399 effective_request_limit = int( 

400 self.state.request_limit * self.config.request_buffer_ratio 

401 ) 

402 effective_token_limit = int( 

403 self.state.token_limit * self.config.token_buffer_ratio 

404 ) 

405 

406 # Request rate limiter 

407 self.request_limiter = Limiter( 

408 Rate( 

409 effective_request_limit, 

410 Duration.MINUTE, 

411 ) 

412 ) 

413 

414 # Token rate limiter 

415 self.token_limiter = Limiter( 

416 Rate( 

417 effective_token_limit, 

418 Duration.MINUTE, 

419 ) 

420 ) 

421 

422 # Store effective limits for logging 

423 self._effective_request_limit = effective_request_limit 

424 self._effective_token_limit = effective_token_limit 

425 

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

427 """Async context manager entry.""" 

428 if self._async_context_active: 

429 raise RuntimeError( 

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

431 ) 

432 

433 self._async_context_active = True 

434 

435 # Discover rate limits if supported 

436 if self.config.supports_dynamic_limits: 

437 await self._discover_rate_limits() 

438 

439 # Print rate limit information if enabled 

440 if self._print_rate_limit_info: 

441 self._print_rate_limit_info_details() 

442 

443 return self 

444 

445 async def __aexit__( 

446 self, 

447 exc_type: type[BaseException] | None, 

448 exc_val: BaseException | None, 

449 exc_tb: object, 

450 ) -> None: 

451 """Async context manager exit.""" 

452 self._async_context_active = False 

453 self._dispose_rate_limiters() 

454 await self.async_client.aclose() 

455 

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

457 """Sync context manager entry.""" 

458 if self._sync_context_active: 

459 raise RuntimeError( 

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

461 ) 

462 

463 self._sync_context_active = True 

464 

465 # Discover rate limits if supported 

466 if self.config.supports_dynamic_limits: 

467 self._discover_rate_limits_sync() 

468 

469 # Print rate limit information if enabled 

470 if self._print_rate_limit_info: 

471 self._print_rate_limit_info_details() 

472 

473 return self 

474 

475 def __exit__( 

476 self, 

477 exc_type: type[BaseException] | None, 

478 exc_val: BaseException | None, 

479 exc_tb: object, 

480 ) -> None: 

481 """Sync context manager exit.""" 

482 self._sync_context_active = False 

483 self._dispose_rate_limiters() 

484 self.sync_client.close() 

485 

486 def _dispose_rate_limiters(self) -> None: 

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

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

489 if rl is not None: 

490 for bucket in rl.buckets(): 

491 rl.dispose(bucket) 

492 self.request_limiter = None 

493 self._effective_request_limit = None 

494 

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

496 if tl is not None: 

497 for bucket in tl.buckets(): 

498 tl.dispose(bucket) 

499 self.token_limiter = None 

500 self._effective_token_limit = None 

501 

502 async def _discover_rate_limits(self) -> None: 

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

504 try: 

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

506 # OpenRouter uses a special auth endpoint 

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

508 response.raise_for_status() 

509 

510 data = response.json() 

511 # Update limits based on response 

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

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

514 

515 else: 

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

517 if self._print_rate_limit_info: 

518 print( 

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

520 ) 

521 logger.info( 

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

523 ) 

524 

525 except Exception as e: 

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

527 

528 def _discover_rate_limits_sync(self) -> None: 

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

530 try: 

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

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

533 response.raise_for_status() 

534 

535 data = response.json() 

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

537 else: 

538 logger.info( 

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

540 ) 

541 

542 except Exception as e: 

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

544 

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

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

547 updated = False 

548 was_uninitialized = ( 

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

550 ) 

551 

552 # Update request limits 

553 if ( 

554 rate_limit_info.requests_limit 

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

556 ): 

557 old_limit = self.state.request_limit 

558 self.state.request_limit = rate_limit_info.requests_limit 

559 updated = True 

560 if was_uninitialized: 

561 message = ( 

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

563 ) 

564 if self._print_rate_limit_info: 

565 print(message) 

566 logger.info(message) 

567 else: 

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

569 if self._print_rate_limit_info: 

570 print(message) 

571 logger.info(message) 

572 

573 # Update token limits 

574 if ( 

575 rate_limit_info.tokens_limit 

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

577 ): 

578 old_limit = self.state.token_limit 

579 self.state.token_limit = rate_limit_info.tokens_limit 

580 updated = True 

581 if was_uninitialized: 

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

583 if self._print_rate_limit_info: 

584 print(message) 

585 logger.info(message) 

586 else: 

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

588 if self._print_rate_limit_info: 

589 print(message) 

590 logger.info(message) 

591 

592 if updated: 

593 # Reinitialize rate limiters with new limits 

594 self._init_rate_limiters() 

595 

596 # Update limits_discovered flag if both limits are now available 

597 if ( 

598 self.state.request_limit is not None 

599 and self.state.token_limit is not None 

600 ): 

601 self._limits_discovered = True 

602 

603 if was_uninitialized: 

604 message = "Rate limiters initialized after discovery" 

605 if self._print_rate_limit_info: 

606 print(message) 

607 # Print updated rate limit info after discovery 

608 self._print_rate_limit_info_details() 

609 logger.info(message) 

610 

611 # Store the rate limit info 

612 self.state.last_rate_limit_info = rate_limit_info 

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

614 

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

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

617 if not self.enable_token_estimation: 

618 return 0 

619 

620 # Simple token estimation 

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

622 if "messages" in request_data: 

623 text = "" 

624 for message in request_data["messages"]: 

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

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

627 

628 # Rough estimation: 1 token ≈ 4 characters 

629 return len(text) // 4 

630 

631 return 0 

632 

633 @asynccontextmanager 

634 async def _acquire_rate_limits( 

635 self, estimated_tokens: int = 0 

636 ) -> AsyncIterator[None]: 

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

638 # Check if rate limiters are initialized 

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

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

641 logger.info( 

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

643 ) 

644 else: 

645 # Wait for request rate limit 

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

647 

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

649 if ( 

650 estimated_tokens > 0 

651 and self.token_limiter is not None 

652 and self._effective_token_limit is not None 

653 ): 

654 # Check if request is too large for bucket capacity 

655 if estimated_tokens > self._effective_token_limit: 

656 # Log warning for large requests 

657 logger.warning( 

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

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

660 ) 

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

662 # Acquire tokens in chunks to avoid bucket overflow 

663 remaining_tokens = estimated_tokens 

664 while remaining_tokens > 0: 

665 chunk_size = min( 

666 remaining_tokens, self._effective_token_limit // 2 

667 ) 

668 await asyncio.to_thread( 

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

670 ) 

671 remaining_tokens -= chunk_size 

672 if remaining_tokens > 0: 

673 # Brief pause to let bucket refill 

674 await asyncio.sleep(0.1) 

675 else: 

676 # Normal acquisition for smaller requests 

677 await asyncio.to_thread( 

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

679 ) 

680 

681 try: 

682 yield 

683 finally: 

684 # Update usage tracking 

685 self.state.requests_used += 1 

686 self.state.tokens_used += estimated_tokens 

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

688 

689 @contextmanager 

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

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

692 # Check if rate limiters are initialized 

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

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

695 logger.info( 

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

697 ) 

698 else: 

699 # Wait for request rate limit 

700 self.request_limiter.try_acquire("request") 

701 

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

703 if ( 

704 estimated_tokens > 0 

705 and self.token_limiter is not None 

706 and self._effective_token_limit is not None 

707 ): 

708 # Check if request is too large for bucket capacity 

709 if estimated_tokens > self._effective_token_limit: 

710 # Log warning for large requests 

711 logger.warning( 

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

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

714 ) 

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

716 # Acquire tokens in chunks to avoid bucket overflow 

717 remaining_tokens = estimated_tokens 

718 while remaining_tokens > 0: 

719 chunk_size = min( 

720 remaining_tokens, self._effective_token_limit // 2 

721 ) 

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

723 remaining_tokens -= chunk_size 

724 if remaining_tokens > 0: 

725 # Brief pause to let bucket refill 

726 time.sleep(0.1) 

727 else: 

728 # Normal acquisition for smaller requests 

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

730 

731 try: 

732 yield 

733 finally: 

734 # Update usage tracking 

735 self.state.requests_used += 1 

736 self.state.tokens_used += estimated_tokens 

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

738 

739 def _get_retry_decorator(self) -> Any: 

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

741 return retry( 

742 stop=stop_after_attempt(self._user_max_retries), 

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

744 retry=retry_if_exception_type( 

745 ( 

746 httpx.HTTPStatusError, 

747 httpx.RequestError, 

748 httpx.ReadTimeout, 

749 httpx.ConnectTimeout, 

750 ) 

751 ), 

752 ) 

753 

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

755 """Get current rate limit information.""" 

756 return { 

757 "provider": self.provider.value, 

758 "request_limit": self.state.request_limit, 

759 "token_limit": self.state.token_limit, 

760 "requests_used": self.state.requests_used, 

761 "tokens_used": self.state.tokens_used, 

762 "last_request_time": self.state.last_request_time, 

763 "last_limit_update": self.state.last_limit_update, 

764 "consecutive_rate_limit_errors": self.state.consecutive_rate_limit_errors, 

765 } 

766 

767 def reset_usage_tracking(self) -> None: 

768 """Reset usage tracking counters.""" 

769 self.state.requests_used = 0 

770 self.state.tokens_used = 0 

771 self.state.consecutive_rate_limit_errors = 0 

772 

773 # High-level chat completion methods 

774 

775 async def chat_completion( 

776 self, 

777 model: str, 

778 messages: list[Message], 

779 max_tokens: int | None = None, 

780 temperature: float | None = None, 

781 top_p: float | None = None, 

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

783 stream: bool = False, 

784 **kwargs: Any, 

785 ) -> ChatCompletionResponse: 

786 """ 

787 Make a high-level chat completion request. 

788 

789 Args: 

790 model: The model to use for completion 

791 messages: List of messages in the conversation 

792 max_tokens: Maximum tokens to generate 

793 temperature: Sampling temperature 

794 top_p: Top-p sampling parameter 

795 stop: Stop sequences 

796 stream: Whether to stream the response 

797 **kwargs: Additional provider-specific parameters 

798 

799 Returns: 

800 ChatCompletionResponse with the completion result 

801 

802 Raises: 

803 ValueError: If provider cannot be determined from model 

804 httpx.HTTPStatusError: For HTTP error responses 

805 httpx.RequestError: For request errors 

806 """ 

807 # Create request object 

808 request = ChatCompletionRequest( 

809 model=model, 

810 messages=messages, 

811 max_tokens=max_tokens, 

812 temperature=temperature, 

813 top_p=top_p, 

814 stop=stop, 

815 stream=stream, 

816 **kwargs, 

817 ) 

818 

819 # Get the appropriate adapter 

820 adapter = get_adapter(self.provider) 

821 

822 # Format the request for the provider 

823 formatted_request = adapter.format_request(request) 

824 

825 # Make the HTTP request with rate limiting 

826 try: 

827 # Print request initiation if enabled 

828 if self._print_request_initiation: 

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

830 

831 # Estimate tokens 

832 estimated_tokens = self._estimate_tokens(formatted_request) 

833 

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

835 client = None 

836 close_client_after_use = False 

837 if self._async_context_active: 

838 client = self.async_client 

839 else: 

840 client = httpx.AsyncClient( 

841 base_url=self.config.base_url, 

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

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

844 ) 

845 close_client_after_use = True 

846 

847 try: 

848 # Acquire rate limits 

849 async with self._acquire_rate_limits(estimated_tokens): 

850 # Make the request 

851 response = await client.request( 

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

853 ) 

854 

855 # Extract rate limit info 

856 from .providers import extract_rate_limit_info 

857 rate_limit_info = extract_rate_limit_info( 

858 dict(response.headers), self.config 

859 ) 

860 

861 # Update our rate limits 

862 if self.enable_adaptive_limits: 

863 self._update_rate_limits(rate_limit_info) 

864 

865 # Handle rate limit errors 

866 if response.status_code == 429: 

867 self.state.consecutive_rate_limit_errors += 1 

868 if rate_limit_info.retry_after: 

869 import asyncio 

870 await asyncio.sleep(rate_limit_info.retry_after) 

871 else: 

872 # Exponential backoff 

873 import asyncio 

874 backoff = self.config.base_backoff * ( 

875 2**self.state.consecutive_rate_limit_errors 

876 ) 

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

878 

879 response.raise_for_status() 

880 else: 

881 # Reset consecutive errors on success 

882 self.state.consecutive_rate_limit_errors = 0 

883 

884 # Parse the response 

885 response_data = response.json() 

886 return adapter.parse_response(response_data, request) 

887 finally: 

888 if close_client_after_use: 

889 await client.aclose() 

890 

891 except Exception as e: 

892 # Handle errors and return error response 

893 error_response = ChatCompletionResponse( 

894 id="error", 

895 model=request.model, 

896 success=False, 

897 error_message=str(e), 

898 choices=[], 

899 usage=None, 

900 created=None, 

901 ) 

902 return error_response 

903 

904 def chat_completion_sync( 

905 self, 

906 model: str, 

907 messages: list[Message], 

908 max_tokens: int | None = None, 

909 temperature: float | None = None, 

910 top_p: float | None = None, 

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

912 stream: bool = False, 

913 **kwargs: Any, 

914 ) -> ChatCompletionResponse: 

915 """ 

916 Make a synchronous high-level chat completion request. 

917 

918 Args: 

919 model: The model to use for completion 

920 messages: List of messages in the conversation 

921 max_tokens: Maximum tokens to generate 

922 temperature: Sampling temperature 

923 top_p: Top-p sampling parameter 

924 stop: Stop sequences 

925 stream: Whether to stream the response 

926 **kwargs: Additional provider-specific parameters 

927 

928 Returns: 

929 ChatCompletionResponse with the completion result 

930 

931 Raises: 

932 ValueError: If provider cannot be determined from model 

933 httpx.HTTPStatusError: For HTTP error responses 

934 httpx.RequestError: For request errors 

935 """ 

936 # Create request object 

937 request = ChatCompletionRequest( 

938 model=model, 

939 messages=messages, 

940 max_tokens=max_tokens, 

941 temperature=temperature, 

942 top_p=top_p, 

943 stop=stop, 

944 stream=stream, 

945 **kwargs, 

946 ) 

947 

948 # Get the appropriate adapter 

949 adapter = get_adapter(self.provider) 

950 

951 # Format the request for the provider 

952 formatted_request = adapter.format_request(request) 

953 

954 # Make the HTTP request with rate limiting 

955 try: 

956 # Print request initiation if enabled 

957 if self._print_request_initiation: 

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

959 

960 # Estimate tokens 

961 estimated_tokens = self._estimate_tokens(formatted_request) 

962 

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

964 client = None 

965 close_client_after_use = False 

966 if self._sync_context_active: 

967 client = self.sync_client 

968 else: 

969 client = httpx.Client( 

970 base_url=self.config.base_url, 

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

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

973 ) 

974 close_client_after_use = True 

975 

976 try: 

977 # Acquire rate limits 

978 with self._acquire_rate_limits_sync(estimated_tokens): 

979 # Make the request 

980 response = client.request( 

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

982 ) 

983 

984 # Extract rate limit info 

985 from .providers import extract_rate_limit_info 

986 rate_limit_info = extract_rate_limit_info( 

987 dict(response.headers), self.config 

988 ) 

989 

990 # Update our rate limits 

991 if self.enable_adaptive_limits: 

992 self._update_rate_limits(rate_limit_info) 

993 

994 # Handle rate limit errors 

995 if response.status_code == 429: 

996 self.state.consecutive_rate_limit_errors += 1 

997 if rate_limit_info.retry_after: 

998 import time 

999 time.sleep(rate_limit_info.retry_after) 

1000 else: 

1001 # Exponential backoff 

1002 import time 

1003 backoff = self.config.base_backoff * ( 

1004 2**self.state.consecutive_rate_limit_errors 

1005 ) 

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

1007 

1008 response.raise_for_status() 

1009 else: 

1010 # Reset consecutive errors on success 

1011 self.state.consecutive_rate_limit_errors = 0 

1012 

1013 # Parse the response 

1014 response_data = response.json() 

1015 return adapter.parse_response(response_data, request) 

1016 finally: 

1017 if close_client_after_use: 

1018 client.close() 

1019 

1020 except Exception as e: 

1021 # Handle errors and return error response 

1022 error_response = ChatCompletionResponse( 

1023 id="error", 

1024 model=request.model, 

1025 success=False, 

1026 error_message=str(e), 

1027 choices=[], 

1028 usage=None, 

1029 created=None, 

1030 ) 

1031 return error_response 

1032 

1033 # Convenience methods for different message types 

1034 

1035 async def simple_chat( 

1036 self, 

1037 model: str, 

1038 prompt: str, 

1039 max_tokens: int | None = None, 

1040 temperature: float | None = None, 

1041 **kwargs: Any, 

1042 ) -> str: 

1043 """ 

1044 Simple chat completion that returns just the text response. 

1045 

1046 Args: 

1047 model: The model to use 

1048 prompt: The user prompt 

1049 max_tokens: Maximum tokens to generate 

1050 temperature: Sampling temperature 

1051 **kwargs: Additional parameters 

1052 

1053 Returns: 

1054 The text response from the model 

1055 """ 

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

1057 response = await self.chat_completion( 

1058 model=model, 

1059 messages=messages, 

1060 max_tokens=max_tokens, 

1061 temperature=temperature, 

1062 **kwargs, 

1063 ) 

1064 

1065 if response.choices: 

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

1067 return "" 

1068 

1069 def simple_chat_sync( 

1070 self, 

1071 model: str, 

1072 prompt: str, 

1073 max_tokens: int | None = None, 

1074 temperature: float | None = None, 

1075 **kwargs: Any, 

1076 ) -> str: 

1077 """ 

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

1079 

1080 Args: 

1081 model: The model to use 

1082 prompt: The user prompt 

1083 max_tokens: Maximum tokens to generate 

1084 temperature: Sampling temperature 

1085 **kwargs: Additional parameters 

1086 

1087 Returns: 

1088 The text response from the model 

1089 """ 

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

1091 response = self.chat_completion_sync( 

1092 model=model, 

1093 messages=messages, 

1094 max_tokens=max_tokens, 

1095 temperature=temperature, 

1096 **kwargs, 

1097 ) 

1098 

1099 if response.choices: 

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

1101 return "" 

1102 

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

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

1105 self._print_rate_limit_info = enabled 

1106 

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

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

1109 self._print_request_initiation = enabled 

1110 

1111 def _print_rate_limit_info_details(self) -> None: 

1112 """Print current rate limit configuration.""" 

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

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

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

1116 

1117 # Handle None values for limits 

1118 if self.state.request_limit is not None: 

1119 effective_req = self._effective_request_limit or "not calculated" 

1120 print( 

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

1122 ) 

1123 else: 

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

1125 

1126 if self.state.token_limit is not None: 

1127 effective_tok = self._effective_token_limit or "not calculated" 

1128 print( 

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

1130 ) 

1131 else: 

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

1133 

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

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

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

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

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

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

1140 print("=" * 50)