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
« prev ^ index » next coverage.py v7.9.2, created at 2025-10-02 12:23 +0100
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 # 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 }
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 )
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 )
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 }
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"
359 # Merge with user-provided headers
360 if "headers" in kwargs:
361 headers.update(kwargs["headers"])
362 kwargs["headers"] = headers
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 )
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 )
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
395 # Dispose existing limiters to prevent background leaker thread accumulation
396 self._dispose_rate_limiters()
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 )
406 # Request rate limiter
407 self.request_limiter = Limiter(
408 Rate(
409 effective_request_limit,
410 Duration.MINUTE,
411 )
412 )
414 # Token rate limiter
415 self.token_limiter = Limiter(
416 Rate(
417 effective_token_limit,
418 Duration.MINUTE,
419 )
420 )
422 # Store effective limits for logging
423 self._effective_request_limit = effective_request_limit
424 self._effective_token_limit = effective_token_limit
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 )
433 self._async_context_active = True
435 # Discover rate limits if supported
436 if self.config.supports_dynamic_limits:
437 await self._discover_rate_limits()
439 # Print rate limit information if enabled
440 if self._print_rate_limit_info:
441 self._print_rate_limit_info_details()
443 return self
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()
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 )
463 self._sync_context_active = True
465 # Discover rate limits if supported
466 if self.config.supports_dynamic_limits:
467 self._discover_rate_limits_sync()
469 # Print rate limit information if enabled
470 if self._print_rate_limit_info:
471 self._print_rate_limit_info_details()
473 return self
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()
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
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
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()
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}")
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 )
525 except Exception as e:
526 logger.warning(f"Failed to discover rate limits: {e}")
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()
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 )
542 except Exception as e:
543 logger.warning(f"Failed to discover rate limits: {e}")
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 )
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)
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)
592 if updated:
593 # Reinitialize rate limiters with new limits
594 self._init_rate_limiters()
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
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)
611 # Store the rate limit info
612 self.state.last_rate_limit_info = rate_limit_info
613 self.state.last_limit_update = time.time()
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
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"])
628 # Rough estimation: 1 token ≈ 4 characters
629 return len(text) // 4
631 return 0
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")
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 )
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()
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")
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)
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()
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 )
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 }
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
773 # High-level chat completion methods
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.
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
799 Returns:
800 ChatCompletionResponse with the completion result
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 )
819 # Get the appropriate adapter
820 adapter = get_adapter(self.provider)
822 # Format the request for the provider
823 formatted_request = adapter.format_request(request)
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)")
831 # Estimate tokens
832 estimated_tokens = self._estimate_tokens(formatted_request)
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
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 )
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 )
861 # Update our rate limits
862 if self.enable_adaptive_limits:
863 self._update_rate_limits(rate_limit_info)
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))
879 response.raise_for_status()
880 else:
881 # Reset consecutive errors on success
882 self.state.consecutive_rate_limit_errors = 0
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()
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
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.
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
928 Returns:
929 ChatCompletionResponse with the completion result
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 )
948 # Get the appropriate adapter
949 adapter = get_adapter(self.provider)
951 # Format the request for the provider
952 formatted_request = adapter.format_request(request)
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)")
960 # Estimate tokens
961 estimated_tokens = self._estimate_tokens(formatted_request)
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
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 )
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 )
990 # Update our rate limits
991 if self.enable_adaptive_limits:
992 self._update_rate_limits(rate_limit_info)
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))
1008 response.raise_for_status()
1009 else:
1010 # Reset consecutive errors on success
1011 self.state.consecutive_rate_limit_errors = 0
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()
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
1033 # Convenience methods for different message types
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.
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
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 )
1065 if response.choices:
1066 return response.choices[0].message.content
1067 return ""
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.
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
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 )
1099 if response.choices:
1100 return response.choices[0].message.content
1101 return ""
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
1107 def set_print_request_initiation(self, enabled: bool) -> None:
1108 """Set whether to print request initiation messages."""
1109 self._print_request_initiation = enabled
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}")
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)")
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)")
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)