Coverage for src/chat_limiter/limiter.py: 77%
386 statements
« prev ^ index » next coverage.py v7.9.2, created at 2025-09-01 14:16 +0100
« prev ^ index » next coverage.py v7.9.2, created at 2025-09-01 14:16 +0100
1"""
2Core rate limiter implementation using PyrateLimiter.
3"""
5import asyncio
6import logging
7import time
8from collections.abc import AsyncIterator, Iterator
9from contextlib import asynccontextmanager, contextmanager
10from dataclasses import dataclass, field
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)
38logger = logging.getLogger(__name__)
41@dataclass
42class LimiterState:
43 """Current state of the rate limiter."""
45 # Current limits (None if not yet discovered)
46 request_limit: int | None = None
47 token_limit: int | None = None
49 # Usage tracking
50 requests_used: int = 0
51 tokens_used: int = 0
53 # Timing
54 last_request_time: float = field(default_factory=time.time)
55 last_limit_update: float = field(default_factory=time.time)
57 # Rate limit info from last response
58 last_rate_limit_info: RateLimitInfo | None = None
60 # Adaptive behavior
61 consecutive_rate_limit_errors: int = 0
62 adaptive_backoff_factor: float = 1.0
65class ChatLimiter:
66 """
67 A Pythonic rate limiter for API calls supporting OpenAI, Anthropic, and OpenRouter.
69 Features:
70 - Automatic rate limit discovery and adaptation
71 - Sync and async support with context managers
72 - Intelligent retry logic with exponential backoff
73 - Token and request rate limiting
74 - Provider-specific optimizations
76 Example:
77 # High-level interface (recommended)
78 async with ChatLimiter.for_model("gpt-4o", api_key="sk-...") as limiter:
79 response = await limiter.chat_completion(
80 model="gpt-4o",
81 messages=[Message(role=MessageRole.USER, content="Hello!")]
82 )
84 # Low-level interface (for advanced users)
85 async with ChatLimiter(provider=Provider.OPENAI, api_key="sk-...") as limiter:
86 response = await limiter.request("POST", "/chat/completions", json=data)
87 """
89 def __init__(
90 self,
91 provider: Provider | None = None,
92 api_key: str | None = None,
93 base_url: str | None = None,
94 config: ProviderConfig | None = None,
95 http_client: httpx.AsyncClient | None = None,
96 sync_http_client: httpx.Client | None = None,
97 enable_adaptive_limits: bool = True,
98 enable_token_estimation: bool = True,
99 request_limit: int | None = None,
100 token_limit: int | None = None,
101 max_retries: int | None = None,
102 base_backoff: float | None = None,
103 timeout: float | None = None,
104 **kwargs: Any,
105 ):
106 """
107 Initialize the ChatLimiter.
109 Args:
110 provider: The API provider (OpenAI, Anthropic, OpenRouter)
111 api_key: API key for authentication
112 base_url: Base URL for API requests
113 config: Custom provider configuration
114 http_client: Custom async HTTP client
115 sync_http_client: Custom sync HTTP client
116 enable_adaptive_limits: Enable adaptive rate limit adjustment
117 enable_token_estimation: Enable token usage estimation
118 request_limit: Override request limit (if not provided, must be discovered from API)
119 token_limit: Override token limit (if not provided, must be discovered from API)
120 max_retries: Override max retries (defaults to 3 if not provided)
121 base_backoff: Override base backoff (defaults to 1.0 if not provided)
122 timeout: HTTP request timeout in seconds (defaults to 120.0 for better reliability)
123 **kwargs: Additional arguments passed to HTTP clients
124 """
125 # Determine provider and config
126 if config:
127 self.config = config
128 self.provider = config.provider
129 elif provider:
130 self.provider = provider
131 self.config = get_provider_config(provider)
132 elif base_url:
133 detected_provider = detect_provider_from_url(base_url)
134 if detected_provider:
135 self.provider = detected_provider
136 self.config = get_provider_config(detected_provider)
137 else:
138 raise ValueError(f"Could not detect provider from URL: {base_url}")
139 else:
140 raise ValueError("Must provide either provider, config, or base_url")
142 # Override base_url if provided
143 if base_url:
144 self.config.base_url = base_url
146 # Store configuration
147 self.api_key = api_key
148 self.enable_adaptive_limits = enable_adaptive_limits
149 self.enable_token_estimation = enable_token_estimation
151 # Store user-provided overrides
152 self._user_request_limit = request_limit
153 self._user_token_limit = token_limit
154 self._user_max_retries = max_retries or 3 # Default to 3 if not provided
155 self._user_base_backoff = base_backoff or 1.0 # Default to 1.0 if not provided
156 self._user_timeout = (
157 timeout or 120.0
158 ) # Default to 120 seconds for better reliability
160 # Determine initial limits (user override, config default, or None for discovery)
161 initial_request_limit = (
162 request_limit or self.config.default_request_limit or None
163 )
164 initial_token_limit = token_limit or self.config.default_token_limit or None
166 # Initialize state - will be None if no defaults and no discovery yet
167 self.state = LimiterState(
168 request_limit=initial_request_limit,
169 token_limit=initial_token_limit,
170 )
172 # Flag to track if we need to discover limits
173 self._limits_discovered = (
174 initial_request_limit is not None and initial_token_limit is not None
175 )
177 # Initialize HTTP clients
178 self._init_http_clients(http_client, sync_http_client, **kwargs)
180 # Initialize rate limiters
181 self._init_rate_limiters()
183 # Context manager state
184 self._async_context_active = False
185 self._sync_context_active = False
187 # Logging configuration
188 self._print_rate_limit_info = False
189 self._print_request_initiation = False
191 @classmethod
192 def for_model(
193 cls,
194 model: str,
195 api_key: str | None = None,
196 provider: str | Provider | None = None,
197 use_dynamic_discovery: bool = True,
198 request_limit: int | None = None,
199 token_limit: int | None = None,
200 max_retries: int | None = None,
201 base_backoff: float | None = None,
202 timeout: float | None = None,
203 **kwargs: Any,
204 ) -> "ChatLimiter":
205 """
206 Create a ChatLimiter instance automatically detecting the provider from the model name.
208 Args:
209 model: The model name (e.g., "gpt-4o", "claude-3-sonnet-20240229")
210 api_key: API key for the provider. If None, will be read from environment variables
211 (OPENAI_API_KEY, ANTHROPIC_API_KEY, OPENROUTER_API_KEY)
212 provider: Override provider detection. Can be "openai", "anthropic", "openrouter",
213 or Provider enum. If None, will be auto-detected from model name
214 use_dynamic_discovery: Whether to query live APIs for model availability (default: True).
215 Requires appropriate API keys to be available. Falls back to
216 hardcoded model lists when disabled or when API calls fail.
217 **kwargs: Additional arguments passed to ChatLimiter
219 Returns:
220 Configured ChatLimiter instance
222 Raises:
223 ValueError: If provider cannot be determined from model name or API key not found
225 Example:
226 # Auto-detect provider with dynamic discovery (default behavior)
227 async with ChatLimiter.for_model("gpt-4o") as limiter:
228 response = await limiter.simple_chat("gpt-4o", "Hello!")
230 # Override provider detection
231 async with ChatLimiter.for_model("custom-model", provider="openai") as limiter:
232 response = await limiter.simple_chat("custom-model", "Hello!")
234 # Disable dynamic discovery to use only hardcoded model lists
235 async with ChatLimiter.for_model("gpt-4o", use_dynamic_discovery=False) as limiter:
236 response = await limiter.simple_chat("gpt-4o", "Hello!")
237 """
238 import os
240 # Determine provider
241 if provider is not None:
242 # Use provided provider
243 if isinstance(provider, str):
244 provider_enum = Provider(provider)
245 else:
246 provider_enum = provider
247 provider_name = provider_enum.value
248 else:
249 # Auto-detect from model name
250 # If dynamic discovery is requested, we need to collect API keys first
251 api_keys_for_discovery = {}
252 if use_dynamic_discovery:
253 # Collect available API keys from environment
254 env_var_map = {
255 "openai": "OPENAI_API_KEY",
256 "anthropic": "ANTHROPIC_API_KEY",
257 "openrouter": "OPENROUTER_API_KEY",
258 }
260 for provider_key, env_var in env_var_map.items():
261 key_value = os.getenv(env_var)
262 if key_value:
263 api_keys_for_discovery[provider_key] = key_value
265 # Try dynamic discovery first to get more detailed information
266 discovery_result = None
267 if use_dynamic_discovery and api_keys_for_discovery:
268 from .models import detect_provider_from_model_sync
270 discovery_result = detect_provider_from_model_sync(
271 model, api_keys_for_discovery
272 )
273 detected_provider = discovery_result.found_provider
274 else:
275 detected_provider = detect_provider_from_model(
276 model, use_dynamic_discovery, api_keys_for_discovery
277 )
279 if not detected_provider:
280 discovery_msg = (
281 " with dynamic API discovery" if use_dynamic_discovery else ""
282 )
283 error_msg = f"Could not determine provider from model '{model}'{discovery_msg}. "
285 # Add detailed information about available models if we have discovery results
286 if discovery_result and discovery_result.get_total_models_found() > 0:
287 error_msg += f"\n\nFound {discovery_result.get_total_models_found()} models across providers:\n"
288 for (
289 provider_name,
290 models,
291 ) in discovery_result.get_all_models().items():
292 error_msg += f" {provider_name}: {len(models)} models\n"
293 for example in sorted(list(models)):
294 error_msg += f" - {example}\n"
295 error_msg += "\nPlease check the model name or specify the provider explicitly using the 'provider' parameter."
296 else:
297 error_msg += "Please specify the provider explicitly using the 'provider' parameter."
299 # Add information about discovery errors if any
300 if discovery_result and discovery_result.errors:
301 error_msg += "\n\nDiscovery errors encountered:\n"
302 for provider_name, error in discovery_result.errors.items():
303 error_msg += f" {provider_name}: {error}\n"
305 raise ValueError(error_msg)
306 assert detected_provider is not None # Help MyPy understand type narrowing
307 provider_name = detected_provider
308 provider_enum = Provider(provider_name)
310 # Determine API key
311 if api_key is None:
312 # Try to get from environment variables
313 env_var_map = {
314 "openai": "OPENAI_API_KEY",
315 "anthropic": "ANTHROPIC_API_KEY",
316 "openrouter": "OPENROUTER_API_KEY",
317 }
319 env_var_name: str | None = env_var_map.get(provider_name)
320 if env_var_name:
321 api_key = os.getenv(env_var_name)
322 if not api_key:
323 raise ValueError(
324 f"API key not provided and {env_var_name} environment variable not set. "
325 f"Please provide api_key parameter or set {env_var_name} environment variable."
326 )
327 else:
328 raise ValueError(
329 f"Unknown provider '{provider_name}'. Cannot determine environment variable for API key."
330 )
332 return cls(
333 provider=provider_enum,
334 api_key=api_key,
335 request_limit=request_limit,
336 token_limit=token_limit,
337 max_retries=max_retries,
338 base_backoff=base_backoff,
339 timeout=timeout,
340 **kwargs,
341 )
343 def _init_http_clients(
344 self,
345 http_client: httpx.AsyncClient | None,
346 sync_http_client: httpx.Client | None,
347 **kwargs: Any,
348 ) -> None:
349 """Initialize HTTP clients with proper headers."""
350 # Prepare headers
351 headers = {
352 "User-Agent": f"chat-limiter/0.1.0 ({self.provider.value})",
353 }
355 # Add provider-specific headers
356 if self.api_key:
357 if self.provider == Provider.OPENAI:
358 headers["Authorization"] = f"Bearer {self.api_key}"
359 elif self.provider == Provider.ANTHROPIC:
360 headers["x-api-key"] = self.api_key
361 headers["anthropic-version"] = "2023-06-01"
362 elif self.provider == Provider.OPENROUTER:
363 headers["Authorization"] = f"Bearer {self.api_key}"
364 headers["HTTP-Referer"] = "https://github.com/your-repo/chat-limiter"
366 # Merge with user-provided headers
367 if "headers" in kwargs:
368 headers.update(kwargs["headers"])
369 kwargs["headers"] = headers
371 # Initialize clients
372 if http_client:
373 self.async_client = http_client
374 else:
375 self.async_client = httpx.AsyncClient(
376 base_url=self.config.base_url,
377 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout
378 **kwargs,
379 )
381 if sync_http_client:
382 self.sync_client = sync_http_client
383 else:
384 self.sync_client = httpx.Client(
385 base_url=self.config.base_url,
386 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout
387 **kwargs,
388 )
390 def _init_rate_limiters(self) -> None:
391 """Initialize PyrateLimiter instances."""
392 # Only initialize if we have limits
393 if self.state.request_limit is None or self.state.token_limit is None:
394 # Cannot initialize rate limiters without limits
395 # This will be called again after limits are discovered
396 self.request_limiter = None
397 self.token_limiter = None
398 self._effective_request_limit = None
399 self._effective_token_limit = None
400 return
402 # Calculate effective limits with buffer
403 effective_request_limit = int(
404 self.state.request_limit * self.config.request_buffer_ratio
405 )
406 effective_token_limit = int(
407 self.state.token_limit * self.config.token_buffer_ratio
408 )
410 # Request rate limiter
411 self.request_limiter = Limiter(
412 Rate(
413 effective_request_limit,
414 Duration.MINUTE,
415 )
416 )
418 # Token rate limiter
419 self.token_limiter = Limiter(
420 Rate(
421 effective_token_limit,
422 Duration.MINUTE,
423 )
424 )
426 # Store effective limits for logging
427 self._effective_request_limit = effective_request_limit
428 self._effective_token_limit = effective_token_limit
430 async def __aenter__(self) -> "ChatLimiter":
431 """Async context manager entry."""
432 if self._async_context_active:
433 raise RuntimeError(
434 "ChatLimiter is already active as an async context manager"
435 )
437 self._async_context_active = True
439 # Discover rate limits if supported
440 if self.config.supports_dynamic_limits:
441 await self._discover_rate_limits()
443 # Print rate limit information if enabled
444 if self._print_rate_limit_info:
445 self._print_rate_limit_info_details()
447 return self
449 async def __aexit__(
450 self,
451 exc_type: type[BaseException] | None,
452 exc_val: BaseException | None,
453 exc_tb: object,
454 ) -> None:
455 """Async context manager exit."""
456 self._async_context_active = False
457 await self.async_client.aclose()
459 def __enter__(self) -> "ChatLimiter":
460 """Sync context manager entry."""
461 if self._sync_context_active:
462 raise RuntimeError(
463 "ChatLimiter is already active as a sync context manager"
464 )
466 self._sync_context_active = True
468 # Discover rate limits if supported
469 if self.config.supports_dynamic_limits:
470 self._discover_rate_limits_sync()
472 # Print rate limit information if enabled
473 if self._print_rate_limit_info:
474 self._print_rate_limit_info_details()
476 return self
478 def __exit__(
479 self,
480 exc_type: type[BaseException] | None,
481 exc_val: BaseException | None,
482 exc_tb: object,
483 ) -> None:
484 """Sync context manager exit."""
485 self._sync_context_active = False
486 self.sync_client.close()
488 async def _discover_rate_limits(self) -> None:
489 """Discover current rate limits from the API."""
490 try:
491 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint:
492 # OpenRouter uses a special auth endpoint
493 response = await self.async_client.get(self.config.auth_endpoint)
494 response.raise_for_status()
496 data = response.json()
497 # Update limits based on response
498 # This is a simplified version - actual implementation would parse the response
499 logger.info(f"Discovered OpenRouter limits: {data}")
501 else:
502 # For other providers, we'll discover limits on first request
503 if self._print_rate_limit_info:
504 print(
505 f"Rate limit discovery will happen on first request for {self.provider.value}"
506 )
507 logger.info(
508 f"Rate limit discovery will happen on first request for {self.provider.value}"
509 )
511 except Exception as e:
512 logger.warning(f"Failed to discover rate limits: {e}")
514 def _discover_rate_limits_sync(self) -> None:
515 """Sync version of rate limit discovery."""
516 try:
517 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint:
518 response = self.sync_client.get(self.config.auth_endpoint)
519 response.raise_for_status()
521 data = response.json()
522 logger.info(f"Discovered OpenRouter limits: {data}")
523 else:
524 logger.info(
525 f"Rate limit discovery will happen on first request for {self.provider.value}"
526 )
528 except Exception as e:
529 logger.warning(f"Failed to discover rate limits: {e}")
531 def _update_rate_limits(self, rate_limit_info: RateLimitInfo) -> None:
532 """Update rate limits based on response headers."""
533 updated = False
534 was_uninitialized = (
535 self.state.request_limit is None or self.state.token_limit is None
536 )
538 # Update request limits
539 if (
540 rate_limit_info.requests_limit
541 and rate_limit_info.requests_limit != self.state.request_limit
542 ):
543 old_limit = self.state.request_limit
544 self.state.request_limit = rate_limit_info.requests_limit
545 updated = True
546 if was_uninitialized:
547 message = (
548 f"Discovered request limit: {self.state.request_limit} req/min"
549 )
550 if self._print_rate_limit_info:
551 print(message)
552 logger.info(message)
553 else:
554 message = f"Updated request limit: {old_limit} -> {self.state.request_limit} req/min"
555 if self._print_rate_limit_info:
556 print(message)
557 logger.info(message)
559 # Update token limits
560 if (
561 rate_limit_info.tokens_limit
562 and rate_limit_info.tokens_limit != self.state.token_limit
563 ):
564 old_limit = self.state.token_limit
565 self.state.token_limit = rate_limit_info.tokens_limit
566 updated = True
567 if was_uninitialized:
568 message = f"Discovered token limit: {self.state.token_limit} tokens/min"
569 if self._print_rate_limit_info:
570 print(message)
571 logger.info(message)
572 else:
573 message = f"Updated token limit: {old_limit} -> {self.state.token_limit} tokens/min"
574 if self._print_rate_limit_info:
575 print(message)
576 logger.info(message)
578 if updated:
579 # Reinitialize rate limiters with new limits
580 self._init_rate_limiters()
582 # Update limits_discovered flag if both limits are now available
583 if (
584 self.state.request_limit is not None
585 and self.state.token_limit is not None
586 ):
587 self._limits_discovered = True
589 if was_uninitialized:
590 message = "Rate limiters initialized after discovery"
591 if self._print_rate_limit_info:
592 print(message)
593 # Print updated rate limit info after discovery
594 self._print_rate_limit_info_details()
595 logger.info(message)
597 # Store the rate limit info
598 self.state.last_rate_limit_info = rate_limit_info
599 self.state.last_limit_update = time.time()
601 def _estimate_tokens(self, request_data: dict[str, Any]) -> int:
602 """Estimate token usage from request data."""
603 if not self.enable_token_estimation:
604 return 0
606 # Simple token estimation
607 # This is a placeholder - real implementation would use tiktoken or similar
608 if "messages" in request_data:
609 text = ""
610 for message in request_data["messages"]:
611 if isinstance(message, dict) and "content" in message:
612 text += str(message["content"])
614 # Rough estimation: 1 token ≈ 4 characters
615 return len(text) // 4
617 return 0
619 @asynccontextmanager
620 async def _acquire_rate_limits(
621 self, estimated_tokens: int = 0
622 ) -> AsyncIterator[None]:
623 """Acquire rate limits before making a request."""
624 # Check if rate limiters are initialized
625 if self.request_limiter is None or self.token_limiter is None:
626 # Limits not yet discovered - this request will help discover them
627 logger.info(
628 "Rate limits not yet discovered, proceeding without rate limiting for discovery"
629 )
630 else:
631 # Wait for request rate limit
632 await asyncio.to_thread(self.request_limiter.try_acquire, "request")
634 # Wait for token rate limit if we have token estimation and limiters are initialized
635 if (
636 estimated_tokens > 0
637 and self.token_limiter is not None
638 and self._effective_token_limit is not None
639 ):
640 # Check if request is too large for bucket capacity
641 if estimated_tokens > self._effective_token_limit:
642 # Log warning for large requests
643 logger.warning(
644 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity "
645 f"of {self._effective_token_limit} tokens. This may cause delays."
646 )
647 # For very large requests, we'll split the acquisition
648 # Acquire tokens in chunks to avoid bucket overflow
649 remaining_tokens = estimated_tokens
650 while remaining_tokens > 0:
651 chunk_size = min(
652 remaining_tokens, self._effective_token_limit // 2
653 )
654 await asyncio.to_thread(
655 self.token_limiter.try_acquire, "token", chunk_size
656 )
657 remaining_tokens -= chunk_size
658 if remaining_tokens > 0:
659 # Brief pause to let bucket refill
660 await asyncio.sleep(0.1)
661 else:
662 # Normal acquisition for smaller requests
663 await asyncio.to_thread(
664 self.token_limiter.try_acquire, "token", estimated_tokens
665 )
667 try:
668 yield
669 finally:
670 # Update usage tracking
671 self.state.requests_used += 1
672 self.state.tokens_used += estimated_tokens
673 self.state.last_request_time = time.time()
675 @contextmanager
676 def _acquire_rate_limits_sync(self, estimated_tokens: int = 0) -> Iterator[None]:
677 """Sync version of rate limit acquisition."""
678 # Check if rate limiters are initialized
679 if self.request_limiter is None or self.token_limiter is None:
680 # Limits not yet discovered - this request will help discover them
681 logger.info(
682 "Rate limits not yet discovered, proceeding without rate limiting for discovery"
683 )
684 else:
685 # Wait for request rate limit
686 self.request_limiter.try_acquire("request")
688 # Wait for token rate limit if we have token estimation and limiters are initialized
689 if (
690 estimated_tokens > 0
691 and self.token_limiter is not None
692 and self._effective_token_limit is not None
693 ):
694 # Check if request is too large for bucket capacity
695 if estimated_tokens > self._effective_token_limit:
696 # Log warning for large requests
697 logger.warning(
698 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity "
699 f"of {self._effective_token_limit} tokens. This may cause delays."
700 )
701 # For very large requests, we'll split the acquisition
702 # Acquire tokens in chunks to avoid bucket overflow
703 remaining_tokens = estimated_tokens
704 while remaining_tokens > 0:
705 chunk_size = min(
706 remaining_tokens, self._effective_token_limit // 2
707 )
708 self.token_limiter.try_acquire("token", chunk_size)
709 remaining_tokens -= chunk_size
710 if remaining_tokens > 0:
711 # Brief pause to let bucket refill
712 time.sleep(0.1)
713 else:
714 # Normal acquisition for smaller requests
715 self.token_limiter.try_acquire("token", estimated_tokens)
717 try:
718 yield
719 finally:
720 # Update usage tracking
721 self.state.requests_used += 1
722 self.state.tokens_used += estimated_tokens
723 self.state.last_request_time = time.time()
725 def _get_retry_decorator(self) -> Any:
726 """Get retry decorator with user-configured parameters."""
727 return retry(
728 stop=stop_after_attempt(self._user_max_retries),
729 wait=wait_exponential(multiplier=self._user_base_backoff, min=1, max=60),
730 retry=retry_if_exception_type(
731 (
732 httpx.HTTPStatusError,
733 httpx.RequestError,
734 httpx.ReadTimeout,
735 httpx.ConnectTimeout,
736 )
737 ),
738 )
740 def get_current_limits(self) -> dict[str, Any]:
741 """Get current rate limit information."""
742 return {
743 "provider": self.provider.value,
744 "request_limit": self.state.request_limit,
745 "token_limit": self.state.token_limit,
746 "requests_used": self.state.requests_used,
747 "tokens_used": self.state.tokens_used,
748 "last_request_time": self.state.last_request_time,
749 "last_limit_update": self.state.last_limit_update,
750 "consecutive_rate_limit_errors": self.state.consecutive_rate_limit_errors,
751 }
753 def reset_usage_tracking(self) -> None:
754 """Reset usage tracking counters."""
755 self.state.requests_used = 0
756 self.state.tokens_used = 0
757 self.state.consecutive_rate_limit_errors = 0
759 # High-level chat completion methods
761 async def chat_completion(
762 self,
763 model: str,
764 messages: list[Message],
765 max_tokens: int | None = None,
766 temperature: float | None = None,
767 top_p: float | None = None,
768 stop: str | list[str] | None = None,
769 stream: bool = False,
770 **kwargs: Any,
771 ) -> ChatCompletionResponse:
772 """
773 Make a high-level chat completion request.
775 Args:
776 model: The model to use for completion
777 messages: List of messages in the conversation
778 max_tokens: Maximum tokens to generate
779 temperature: Sampling temperature
780 top_p: Top-p sampling parameter
781 stop: Stop sequences
782 stream: Whether to stream the response
783 **kwargs: Additional provider-specific parameters
785 Returns:
786 ChatCompletionResponse with the completion result
788 Raises:
789 ValueError: If provider cannot be determined from model
790 httpx.HTTPStatusError: For HTTP error responses
791 httpx.RequestError: For request errors
792 """
793 if not self._async_context_active:
794 raise RuntimeError("ChatLimiter must be used as an async context manager")
796 # Create request object
797 request = ChatCompletionRequest(
798 model=model,
799 messages=messages,
800 max_tokens=max_tokens,
801 temperature=temperature,
802 top_p=top_p,
803 stop=stop,
804 stream=stream,
805 **kwargs,
806 )
808 # Get the appropriate adapter
809 adapter = get_adapter(self.provider)
811 # Format the request for the provider
812 formatted_request = adapter.format_request(request)
814 # Make the HTTP request with rate limiting
815 try:
816 # Print request initiation if enabled
817 if self._print_request_initiation:
818 print(f"Sending request for model {model} (attempt 1)")
820 # Estimate tokens
821 estimated_tokens = self._estimate_tokens(formatted_request)
823 # Acquire rate limits
824 async with self._acquire_rate_limits(estimated_tokens):
825 # Make the request
826 response = await self.async_client.request(
827 "POST", adapter.get_endpoint(), json=formatted_request
828 )
830 # Extract rate limit info
831 from .providers import extract_rate_limit_info
832 rate_limit_info = extract_rate_limit_info(
833 dict(response.headers), self.config
834 )
836 # Update our rate limits
837 if self.enable_adaptive_limits:
838 self._update_rate_limits(rate_limit_info)
840 # Handle rate limit errors
841 if response.status_code == 429:
842 self.state.consecutive_rate_limit_errors += 1
843 if rate_limit_info.retry_after:
844 import asyncio
845 await asyncio.sleep(rate_limit_info.retry_after)
846 else:
847 # Exponential backoff
848 import asyncio
849 backoff = self.config.base_backoff * (
850 2**self.state.consecutive_rate_limit_errors
851 )
852 await asyncio.sleep(min(backoff, self.config.max_backoff))
854 response.raise_for_status()
855 else:
856 # Reset consecutive errors on success
857 self.state.consecutive_rate_limit_errors = 0
859 # Parse the response
860 response_data = response.json()
861 return adapter.parse_response(response_data, request)
863 except Exception as e:
864 # Handle errors and return error response
865 error_response = ChatCompletionResponse(
866 id="error",
867 model=request.model,
868 success=False,
869 error_message=str(e),
870 choices=[],
871 usage=None,
872 created=None,
873 )
874 return error_response
876 def chat_completion_sync(
877 self,
878 model: str,
879 messages: list[Message],
880 max_tokens: int | None = None,
881 temperature: float | None = None,
882 top_p: float | None = None,
883 stop: str | list[str] | None = None,
884 stream: bool = False,
885 **kwargs: Any,
886 ) -> ChatCompletionResponse:
887 """
888 Make a synchronous high-level chat completion request.
890 Args:
891 model: The model to use for completion
892 messages: List of messages in the conversation
893 max_tokens: Maximum tokens to generate
894 temperature: Sampling temperature
895 top_p: Top-p sampling parameter
896 stop: Stop sequences
897 stream: Whether to stream the response
898 **kwargs: Additional provider-specific parameters
900 Returns:
901 ChatCompletionResponse with the completion result
903 Raises:
904 ValueError: If provider cannot be determined from model
905 httpx.HTTPStatusError: For HTTP error responses
906 httpx.RequestError: For request errors
907 """
908 if not self._sync_context_active:
909 raise RuntimeError("ChatLimiter must be used as a sync context manager")
911 # Create request object
912 request = ChatCompletionRequest(
913 model=model,
914 messages=messages,
915 max_tokens=max_tokens,
916 temperature=temperature,
917 top_p=top_p,
918 stop=stop,
919 stream=stream,
920 **kwargs,
921 )
923 # Get the appropriate adapter
924 adapter = get_adapter(self.provider)
926 # Format the request for the provider
927 formatted_request = adapter.format_request(request)
929 # Make the HTTP request with rate limiting
930 try:
931 # Print request initiation if enabled
932 if self._print_request_initiation:
933 print(f"Sending request for model {model} (attempt 1)")
935 # Estimate tokens
936 estimated_tokens = self._estimate_tokens(formatted_request)
938 # Acquire rate limits
939 with self._acquire_rate_limits_sync(estimated_tokens):
940 # Make the request
941 response = self.sync_client.request(
942 "POST", adapter.get_endpoint(), json=formatted_request
943 )
945 # Extract rate limit info
946 from .providers import extract_rate_limit_info
947 rate_limit_info = extract_rate_limit_info(
948 dict(response.headers), self.config
949 )
951 # Update our rate limits
952 if self.enable_adaptive_limits:
953 self._update_rate_limits(rate_limit_info)
955 # Handle rate limit errors
956 if response.status_code == 429:
957 self.state.consecutive_rate_limit_errors += 1
958 if rate_limit_info.retry_after:
959 import time
960 time.sleep(rate_limit_info.retry_after)
961 else:
962 # Exponential backoff
963 import time
964 backoff = self.config.base_backoff * (
965 2**self.state.consecutive_rate_limit_errors
966 )
967 time.sleep(min(backoff, self.config.max_backoff))
969 response.raise_for_status()
970 else:
971 # Reset consecutive errors on success
972 self.state.consecutive_rate_limit_errors = 0
974 # Parse the response
975 response_data = response.json()
976 return adapter.parse_response(response_data, request)
978 except Exception as e:
979 # Handle errors and return error response
980 error_response = ChatCompletionResponse(
981 id="error",
982 model=request.model,
983 success=False,
984 error_message=str(e),
985 choices=[],
986 usage=None,
987 created=None,
988 )
989 return error_response
991 # Convenience methods for different message types
993 async def simple_chat(
994 self,
995 model: str,
996 prompt: str,
997 max_tokens: int | None = None,
998 temperature: float | None = None,
999 **kwargs: Any,
1000 ) -> str:
1001 """
1002 Simple chat completion that returns just the text response.
1004 Args:
1005 model: The model to use
1006 prompt: The user prompt
1007 max_tokens: Maximum tokens to generate
1008 temperature: Sampling temperature
1009 **kwargs: Additional parameters
1011 Returns:
1012 The text response from the model
1013 """
1014 messages = [Message(role=MessageRole.USER, content=prompt)]
1015 response = await self.chat_completion(
1016 model=model,
1017 messages=messages,
1018 max_tokens=max_tokens,
1019 temperature=temperature,
1020 **kwargs,
1021 )
1023 if response.choices:
1024 return response.choices[0].message.content
1025 return ""
1027 def simple_chat_sync(
1028 self,
1029 model: str,
1030 prompt: str,
1031 max_tokens: int | None = None,
1032 temperature: float | None = None,
1033 **kwargs: Any,
1034 ) -> str:
1035 """
1036 Simple synchronous chat completion that returns just the text response.
1038 Args:
1039 model: The model to use
1040 prompt: The user prompt
1041 max_tokens: Maximum tokens to generate
1042 temperature: Sampling temperature
1043 **kwargs: Additional parameters
1045 Returns:
1046 The text response from the model
1047 """
1048 messages = [Message(role=MessageRole.USER, content=prompt)]
1049 response = self.chat_completion_sync(
1050 model=model,
1051 messages=messages,
1052 max_tokens=max_tokens,
1053 temperature=temperature,
1054 **kwargs,
1055 )
1057 if response.choices:
1058 return response.choices[0].message.content
1059 return ""
1061 def set_print_rate_limit_info(self, enabled: bool) -> None:
1062 """Set whether to print rate limit information."""
1063 self._print_rate_limit_info = enabled
1065 def set_print_request_initiation(self, enabled: bool) -> None:
1066 """Set whether to print request initiation messages."""
1067 self._print_request_initiation = enabled
1069 def _print_rate_limit_info_details(self) -> None:
1070 """Print current rate limit configuration."""
1071 print(f"\n=== Rate Limit Configuration for {self.provider.value.title()} ===")
1072 print(f"Provider: {self.provider.value}")
1073 print(f"Base URL: {self.config.base_url}")
1075 # Handle None values for limits
1076 if self.state.request_limit is not None:
1077 effective_req = self._effective_request_limit or "not calculated"
1078 print(
1079 f"Request Limit: {self.state.request_limit}/minute (effective: {effective_req}/minute)"
1080 )
1081 else:
1082 print("Request Limit: Not yet discovered (will be fetched from API)")
1084 if self.state.token_limit is not None:
1085 effective_tok = self._effective_token_limit or "not calculated"
1086 print(
1087 f"Token Limit: {self.state.token_limit}/minute (effective: {effective_tok}/minute)"
1088 )
1089 else:
1090 print("Token Limit: Not yet discovered (will be fetched from API)")
1092 print(f"Request Buffer Ratio: {self.config.request_buffer_ratio}")
1093 print(f"Token Buffer Ratio: {self.config.token_buffer_ratio}")
1094 print(f"Adaptive Limits: {self.enable_adaptive_limits}")
1095 print(f"Token Estimation: {self.enable_token_estimation}")
1096 print(f"Dynamic Discovery: {self.config.supports_dynamic_limits}")
1097 print(f"Limits Discovered: {self._limits_discovered}")
1098 print("=" * 50)