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
« prev ^ index » next coverage.py v7.9.2, created at 2025-12-09 08:16 -0500
1"""
2Core rate limiter implementation using PyrateLimiter.
3"""
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
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)
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
39logger = logging.getLogger(__name__)
42@dataclass
43class LimiterState:
44 """Current state of the rate limiter."""
46 # Current limits (None if not yet discovered)
47 request_limit: int | None = None
48 token_limit: int | None = None
50 # Usage tracking
51 requests_used: int = 0
52 tokens_used: int = 0
54 # Timing
55 last_request_time: float = field(default_factory=time.time)
56 last_limit_update: float = field(default_factory=time.time)
58 # Rate limit info from last response
59 last_rate_limit_info: RateLimitInfo | None = None
61 # Adaptive behavior
62 consecutive_rate_limit_errors: int = 0
63 adaptive_backoff_factor: float = 1.0
66class ChatLimiter:
67 """
68 A Pythonic rate limiter for API calls supporting OpenAI, Anthropic, and OpenRouter.
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
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 )
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 """
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.
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")
143 # Override base_url if provided
144 if base_url:
145 self.config.base_url = base_url
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
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
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
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 )
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 )
178 # Initialize HTTP clients
179 self._init_http_clients(http_client, sync_http_client, **kwargs)
181 # Initialize rate limiters
182 self._init_rate_limiters()
184 # Context manager state
185 self._async_context_active = False
186 self._sync_context_active = False
188 # Logging configuration
189 self._print_rate_limit_info = False
190 self._print_request_initiation = False
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.
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
220 Returns:
221 Configured ChatLimiter instance
223 Raises:
224 ValueError: If provider cannot be determined from model name or API key not found
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!")
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!")
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
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 }
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
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 )
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}. "
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."
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"
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)
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)
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 )
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 )
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 )
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 }
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 )
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 )
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 }
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"
432 # Merge with user-provided headers
433 if "headers" in kwargs:
434 headers.update(kwargs["headers"])
435 kwargs["headers"] = headers
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 )
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 )
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
468 # Dispose existing limiters to prevent background leaker thread accumulation
469 self._dispose_rate_limiters()
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 )
479 # Request rate limiter
480 self.request_limiter = Limiter(
481 Rate(
482 effective_request_limit,
483 Duration.MINUTE,
484 )
485 )
487 # Token rate limiter
488 self.token_limiter = Limiter(
489 Rate(
490 effective_token_limit,
491 Duration.MINUTE,
492 )
493 )
495 # Store effective limits for logging
496 self._effective_request_limit = effective_request_limit
497 self._effective_token_limit = effective_token_limit
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 )
506 self._async_context_active = True
508 # Discover rate limits if supported
509 if self.config.supports_dynamic_limits:
510 await self._discover_rate_limits()
512 # Print rate limit information if enabled
513 if self._print_rate_limit_info:
514 self._print_rate_limit_info_details()
516 return self
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()
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 )
536 self._sync_context_active = True
538 # Discover rate limits if supported
539 if self.config.supports_dynamic_limits:
540 self._discover_rate_limits_sync()
542 # Print rate limit information if enabled
543 if self._print_rate_limit_info:
544 self._print_rate_limit_info_details()
546 return self
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()
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
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
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()
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}")
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 )
598 except Exception as e:
599 logger.warning(f"Failed to discover rate limits: {e}")
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()
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 )
615 except Exception as e:
616 logger.warning(f"Failed to discover rate limits: {e}")
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 )
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)
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)
665 if updated:
666 # Reinitialize rate limiters with new limits
667 self._init_rate_limiters()
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
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)
684 # Store the rate limit info
685 self.state.last_rate_limit_info = rate_limit_info
686 self.state.last_limit_update = time.time()
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
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"])
701 # Rough estimation: 1 token ≈ 4 characters
702 return len(text) // 4
704 return 0
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")
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 )
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()
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")
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)
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()
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 )
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 }
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
846 # High-level chat completion methods
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.
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
872 Returns:
873 ChatCompletionResponse with the completion result
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 )
892 # Get the appropriate adapter
893 adapter = get_adapter(self.provider)
895 # Format the request for the provider
896 formatted_request = adapter.format_request(request)
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)")
904 # Estimate tokens
905 estimated_tokens = self._estimate_tokens(formatted_request)
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
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 )
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 )
934 # Update our rate limits
935 if self.enable_adaptive_limits:
936 self._update_rate_limits(rate_limit_info)
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))
952 response.raise_for_status()
953 else:
954 # Reset consecutive errors on success
955 self.state.consecutive_rate_limit_errors = 0
957 # Raise for all non-2xx responses (do not silently succeed)
958 if response.status_code != 429:
959 response.raise_for_status()
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
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.
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
1019 Returns:
1020 ChatCompletionResponse with the completion result
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 )
1039 # Get the appropriate adapter
1040 adapter = get_adapter(self.provider)
1042 # Format the request for the provider
1043 formatted_request = adapter.format_request(request)
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)")
1051 # Estimate tokens
1052 estimated_tokens = self._estimate_tokens(formatted_request)
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
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 )
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 )
1081 # Update our rate limits
1082 if self.enable_adaptive_limits:
1083 self._update_rate_limits(rate_limit_info)
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))
1099 response.raise_for_status()
1100 else:
1101 # Reset consecutive errors on success
1102 self.state.consecutive_rate_limit_errors = 0
1104 # Raise for all non-2xx responses (do not silently succeed)
1105 if response.status_code != 429:
1106 response.raise_for_status()
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
1142 # Convenience methods for different message types
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.
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
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 )
1174 if response.choices:
1175 return response.choices[0].message.content
1176 return ""
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.
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
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 )
1208 if response.choices:
1209 return response.choices[0].message.content
1210 return ""
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
1216 def set_print_request_initiation(self, enabled: bool) -> None:
1217 """Set whether to print request initiation messages."""
1218 self._print_request_initiation = enabled
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}")
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)")
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)")
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)