Coverage for src/chat_limiter/limiter.py: 82%
377 statements
« prev ^ index » next coverage.py v7.9.2, created at 2025-07-09 16:53 +0100
« prev ^ index » next coverage.py v7.9.2, created at 2025-07-09 16:53 +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 extract_rate_limit_info,
29 get_provider_config,
30)
31from .types import (
32 ChatCompletionRequest,
33 ChatCompletionResponse,
34 Message,
35 MessageRole,
36 detect_provider_from_model,
37)
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 = timeout or 120.0 # Default to 120 seconds for better reliability
159 # Determine initial limits (user override, config default, or None for discovery)
160 initial_request_limit = (
161 request_limit or
162 self.config.default_request_limit or
163 None
164 )
165 initial_token_limit = (
166 token_limit or
167 self.config.default_token_limit or
168 None
169 )
171 # Initialize state - will be None if no defaults and no discovery yet
172 self.state = LimiterState(
173 request_limit=initial_request_limit,
174 token_limit=initial_token_limit,
175 )
177 # Flag to track if we need to discover limits
178 self._limits_discovered = initial_request_limit is not None and initial_token_limit is not None
180 # Initialize HTTP clients
181 self._init_http_clients(http_client, sync_http_client, **kwargs)
183 # Initialize rate limiters
184 self._init_rate_limiters()
186 # Context manager state
187 self._async_context_active = False
188 self._sync_context_active = False
190 # Verbose mode (can be set by batch processor)
191 self._verbose_mode = False
193 @classmethod
194 def for_model(
195 cls,
196 model: str,
197 api_key: str | None = None,
198 provider: str | Provider | None = None,
199 use_dynamic_discovery: bool = True,
200 request_limit: int | None = None,
201 token_limit: int | None = None,
202 max_retries: int | None = None,
203 base_backoff: float | None = None,
204 timeout: float | None = None,
205 **kwargs: Any,
206 ) -> "ChatLimiter":
207 """
208 Create a ChatLimiter instance automatically detecting the provider from the model name.
210 Args:
211 model: The model name (e.g., "gpt-4o", "claude-3-sonnet-20240229")
212 api_key: API key for the provider. If None, will be read from environment variables
213 (OPENAI_API_KEY, ANTHROPIC_API_KEY, OPENROUTER_API_KEY)
214 provider: Override provider detection. Can be "openai", "anthropic", "openrouter",
215 or Provider enum. If None, will be auto-detected from model name
216 use_dynamic_discovery: Whether to query live APIs for model availability (default: True).
217 Requires appropriate API keys to be available. Falls back to
218 hardcoded model lists when disabled or when API calls fail.
219 **kwargs: Additional arguments passed to ChatLimiter
221 Returns:
222 Configured ChatLimiter instance
224 Raises:
225 ValueError: If provider cannot be determined from model name or API key not found
227 Example:
228 # Auto-detect provider with dynamic discovery (default behavior)
229 async with ChatLimiter.for_model("gpt-4o") as limiter:
230 response = await limiter.simple_chat("gpt-4o", "Hello!")
232 # Override provider detection
233 async with ChatLimiter.for_model("custom-model", provider="openai") as limiter:
234 response = await limiter.simple_chat("custom-model", "Hello!")
236 # Disable dynamic discovery to use only hardcoded model lists
237 async with ChatLimiter.for_model("gpt-4o", use_dynamic_discovery=False) as limiter:
238 response = await limiter.simple_chat("gpt-4o", "Hello!")
239 """
240 import os
242 # Determine provider
243 if provider is not None:
244 # Use provided provider
245 if isinstance(provider, str):
246 provider_enum = Provider(provider)
247 else:
248 provider_enum = provider
249 provider_name = provider_enum.value
250 else:
251 # Auto-detect from model name
252 # If dynamic discovery is requested, we need to collect API keys first
253 api_keys_for_discovery = {}
254 if use_dynamic_discovery:
255 # Collect available API keys from environment
256 env_var_map = {
257 "openai": "OPENAI_API_KEY",
258 "anthropic": "ANTHROPIC_API_KEY",
259 "openrouter": "OPENROUTER_API_KEY"
260 }
262 for provider_key, env_var in env_var_map.items():
263 key_value = os.getenv(env_var)
264 if key_value:
265 api_keys_for_discovery[provider_key] = key_value
267 detected_provider = detect_provider_from_model(model, use_dynamic_discovery, api_keys_for_discovery)
268 if not detected_provider:
269 discovery_msg = " with dynamic API discovery" if use_dynamic_discovery else ""
270 raise ValueError(
271 f"Could not determine provider from model '{model}'{discovery_msg}. "
272 "Please specify the provider explicitly using the 'provider' parameter."
273 )
274 assert detected_provider is not None # Help MyPy understand type narrowing
275 provider_name = detected_provider
276 provider_enum = Provider(provider_name)
278 # Determine API key
279 if api_key is None:
280 # Try to get from environment variables
281 env_var_map = {
282 "openai": "OPENAI_API_KEY",
283 "anthropic": "ANTHROPIC_API_KEY",
284 "openrouter": "OPENROUTER_API_KEY"
285 }
287 env_var_name: str | None = env_var_map.get(provider_name)
288 if env_var_name:
289 api_key = os.getenv(env_var_name)
290 if not api_key:
291 raise ValueError(
292 f"API key not provided and {env_var_name} environment variable not set. "
293 f"Please provide api_key parameter or set {env_var_name} environment variable."
294 )
295 else:
296 raise ValueError(
297 f"Unknown provider '{provider_name}'. Cannot determine environment variable for API key."
298 )
300 return cls(
301 provider=provider_enum,
302 api_key=api_key,
303 request_limit=request_limit,
304 token_limit=token_limit,
305 max_retries=max_retries,
306 base_backoff=base_backoff,
307 timeout=timeout,
308 **kwargs
309 )
311 def _init_http_clients(
312 self,
313 http_client: httpx.AsyncClient | None,
314 sync_http_client: httpx.Client | None,
315 **kwargs: Any,
316 ) -> None:
317 """Initialize HTTP clients with proper headers."""
318 # Prepare headers
319 headers = {
320 "User-Agent": f"chat-limiter/0.1.0 ({self.provider.value})",
321 }
323 # Add provider-specific headers
324 if self.api_key:
325 if self.provider == Provider.OPENAI:
326 headers["Authorization"] = f"Bearer {self.api_key}"
327 elif self.provider == Provider.ANTHROPIC:
328 headers["x-api-key"] = self.api_key
329 headers["anthropic-version"] = "2023-06-01"
330 elif self.provider == Provider.OPENROUTER:
331 headers["Authorization"] = f"Bearer {self.api_key}"
332 headers["HTTP-Referer"] = "https://github.com/your-repo/chat-limiter"
334 # Merge with user-provided headers
335 if "headers" in kwargs:
336 headers.update(kwargs["headers"])
337 kwargs["headers"] = headers
339 # Initialize clients
340 if http_client:
341 self.async_client = http_client
342 else:
343 self.async_client = httpx.AsyncClient(
344 base_url=self.config.base_url,
345 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout
346 **kwargs,
347 )
349 if sync_http_client:
350 self.sync_client = sync_http_client
351 else:
352 self.sync_client = httpx.Client(
353 base_url=self.config.base_url,
354 timeout=httpx.Timeout(self._user_timeout), # Configurable timeout
355 **kwargs,
356 )
358 def _init_rate_limiters(self) -> None:
359 """Initialize PyrateLimiter instances."""
360 # Only initialize if we have limits
361 if self.state.request_limit is None or self.state.token_limit is None:
362 # Cannot initialize rate limiters without limits
363 # This will be called again after limits are discovered
364 self.request_limiter = None
365 self.token_limiter = None
366 self._effective_request_limit = None
367 self._effective_token_limit = None
368 return
370 # Calculate effective limits with buffer
371 effective_request_limit = int(self.state.request_limit * self.config.request_buffer_ratio)
372 effective_token_limit = int(self.state.token_limit * self.config.token_buffer_ratio)
374 # Request rate limiter
375 self.request_limiter = Limiter(
376 Rate(
377 effective_request_limit,
378 Duration.MINUTE,
379 )
380 )
382 # Token rate limiter
383 self.token_limiter = Limiter(
384 Rate(
385 effective_token_limit,
386 Duration.MINUTE,
387 )
388 )
390 # Store effective limits for logging
391 self._effective_request_limit = effective_request_limit
392 self._effective_token_limit = effective_token_limit
394 async def __aenter__(self) -> "ChatLimiter":
395 """Async context manager entry."""
396 if self._async_context_active:
397 raise RuntimeError(
398 "ChatLimiter is already active as an async context manager"
399 )
401 self._async_context_active = True
403 # Discover rate limits if supported
404 if self.config.supports_dynamic_limits:
405 await self._discover_rate_limits()
407 # Print rate limit information if verbose mode is enabled
408 if self._verbose_mode:
409 self._print_rate_limit_info()
411 return self
413 async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: object) -> None:
414 """Async context manager exit."""
415 self._async_context_active = False
416 await self.async_client.aclose()
418 def __enter__(self) -> "ChatLimiter":
419 """Sync context manager entry."""
420 if self._sync_context_active:
421 raise RuntimeError(
422 "ChatLimiter is already active as a sync context manager"
423 )
425 self._sync_context_active = True
427 # Discover rate limits if supported
428 if self.config.supports_dynamic_limits:
429 self._discover_rate_limits_sync()
431 # Print rate limit information if verbose mode is enabled
432 if self._verbose_mode:
433 self._print_rate_limit_info()
435 return self
437 def __exit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: object) -> None:
438 """Sync context manager exit."""
439 self._sync_context_active = False
440 self.sync_client.close()
442 async def _discover_rate_limits(self) -> None:
443 """Discover current rate limits from the API."""
444 try:
445 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint:
446 # OpenRouter uses a special auth endpoint
447 response = await self.async_client.get(self.config.auth_endpoint)
448 response.raise_for_status()
450 data = response.json()
451 # Update limits based on response
452 # This is a simplified version - actual implementation would parse the response
453 logger.info(f"Discovered OpenRouter limits: {data}")
455 else:
456 # For other providers, we'll discover limits on first request
457 if self._verbose_mode:
458 print(f"Rate limit discovery will happen on first request for {self.provider.value}")
459 logger.info(
460 f"Rate limit discovery will happen on first request for {self.provider.value}"
461 )
463 except Exception as e:
464 logger.warning(f"Failed to discover rate limits: {e}")
466 def _discover_rate_limits_sync(self) -> None:
467 """Sync version of rate limit discovery."""
468 try:
469 if self.provider == Provider.OPENROUTER and self.config.auth_endpoint:
470 response = self.sync_client.get(self.config.auth_endpoint)
471 response.raise_for_status()
473 data = response.json()
474 logger.info(f"Discovered OpenRouter limits: {data}")
475 else:
476 logger.info(
477 f"Rate limit discovery will happen on first request for {self.provider.value}"
478 )
480 except Exception as e:
481 logger.warning(f"Failed to discover rate limits: {e}")
483 def _update_rate_limits(self, rate_limit_info: RateLimitInfo) -> None:
484 """Update rate limits based on response headers."""
485 updated = False
486 was_uninitialized = self.state.request_limit is None or self.state.token_limit is None
488 # Update request limits
489 if (
490 rate_limit_info.requests_limit
491 and rate_limit_info.requests_limit != self.state.request_limit
492 ):
493 old_limit = self.state.request_limit
494 self.state.request_limit = rate_limit_info.requests_limit
495 updated = True
496 if was_uninitialized:
497 message = f"Discovered request limit: {self.state.request_limit} req/min"
498 if self._verbose_mode:
499 print(message)
500 logger.info(message)
501 else:
502 message = f"Updated request limit: {old_limit} -> {self.state.request_limit} req/min"
503 if self._verbose_mode:
504 print(message)
505 logger.info(message)
507 # Update token limits
508 if (
509 rate_limit_info.tokens_limit
510 and rate_limit_info.tokens_limit != self.state.token_limit
511 ):
512 old_limit = self.state.token_limit
513 self.state.token_limit = rate_limit_info.tokens_limit
514 updated = True
515 if was_uninitialized:
516 message = f"Discovered token limit: {self.state.token_limit} tokens/min"
517 if self._verbose_mode:
518 print(message)
519 logger.info(message)
520 else:
521 message = f"Updated token limit: {old_limit} -> {self.state.token_limit} tokens/min"
522 if self._verbose_mode:
523 print(message)
524 logger.info(message)
526 if updated:
527 # Reinitialize rate limiters with new limits
528 self._init_rate_limiters()
530 # Update limits_discovered flag if both limits are now available
531 if self.state.request_limit is not None and self.state.token_limit is not None:
532 self._limits_discovered = True
534 if was_uninitialized:
535 message = "Rate limiters initialized after discovery"
536 if self._verbose_mode:
537 print(message)
538 # Print updated rate limit info after discovery
539 self._print_rate_limit_info()
540 logger.info(message)
542 # Store the rate limit info
543 self.state.last_rate_limit_info = rate_limit_info
544 self.state.last_limit_update = time.time()
546 def _estimate_tokens(self, request_data: dict[str, Any]) -> int:
547 """Estimate token usage from request data."""
548 if not self.enable_token_estimation:
549 return 0
551 # Simple token estimation
552 # This is a placeholder - real implementation would use tiktoken or similar
553 if "messages" in request_data:
554 text = ""
555 for message in request_data["messages"]:
556 if isinstance(message, dict) and "content" in message:
557 text += str(message["content"])
559 # Rough estimation: 1 token ≈ 4 characters
560 return len(text) // 4
562 return 0
564 @asynccontextmanager
565 async def _acquire_rate_limits(
566 self, estimated_tokens: int = 0
567 ) -> AsyncIterator[None]:
568 """Acquire rate limits before making a request."""
569 # Check if rate limiters are initialized
570 if self.request_limiter is None or self.token_limiter is None:
571 # Limits not yet discovered - this request will help discover them
572 logger.info("Rate limits not yet discovered, proceeding without rate limiting for discovery")
573 else:
574 # Wait for request rate limit
575 await asyncio.to_thread(self.request_limiter.try_acquire, "request")
577 # Wait for token rate limit if we have token estimation and limiters are initialized
578 if estimated_tokens > 0 and self.token_limiter is not None and self._effective_token_limit is not None:
579 # Check if request is too large for bucket capacity
580 if estimated_tokens > self._effective_token_limit:
581 # Log warning for large requests
582 logger.warning(
583 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity "
584 f"of {self._effective_token_limit} tokens. This may cause delays."
585 )
586 # For very large requests, we'll split the acquisition
587 # Acquire tokens in chunks to avoid bucket overflow
588 remaining_tokens = estimated_tokens
589 while remaining_tokens > 0:
590 chunk_size = min(remaining_tokens, self._effective_token_limit // 2)
591 await asyncio.to_thread(self.token_limiter.try_acquire, "token", chunk_size)
592 remaining_tokens -= chunk_size
593 if remaining_tokens > 0:
594 # Brief pause to let bucket refill
595 await asyncio.sleep(0.1)
596 else:
597 # Normal acquisition for smaller requests
598 await asyncio.to_thread(self.token_limiter.try_acquire, "token", estimated_tokens)
600 try:
601 yield
602 finally:
603 # Update usage tracking
604 self.state.requests_used += 1
605 self.state.tokens_used += estimated_tokens
606 self.state.last_request_time = time.time()
608 @contextmanager
609 def _acquire_rate_limits_sync(self, estimated_tokens: int = 0) -> Iterator[None]:
610 """Sync version of rate limit acquisition."""
611 # Check if rate limiters are initialized
612 if self.request_limiter is None or self.token_limiter is None:
613 # Limits not yet discovered - this request will help discover them
614 logger.info("Rate limits not yet discovered, proceeding without rate limiting for discovery")
615 else:
616 # Wait for request rate limit
617 self.request_limiter.try_acquire("request")
619 # Wait for token rate limit if we have token estimation and limiters are initialized
620 if estimated_tokens > 0 and self.token_limiter is not None and self._effective_token_limit is not None:
621 # Check if request is too large for bucket capacity
622 if estimated_tokens > self._effective_token_limit:
623 # Log warning for large requests
624 logger.warning(
625 f"Request estimated at {estimated_tokens} tokens exceeds bucket capacity "
626 f"of {self._effective_token_limit} tokens. This may cause delays."
627 )
628 # For very large requests, we'll split the acquisition
629 # Acquire tokens in chunks to avoid bucket overflow
630 remaining_tokens = estimated_tokens
631 while remaining_tokens > 0:
632 chunk_size = min(remaining_tokens, self._effective_token_limit // 2)
633 self.token_limiter.try_acquire("token", chunk_size)
634 remaining_tokens -= chunk_size
635 if remaining_tokens > 0:
636 # Brief pause to let bucket refill
637 time.sleep(0.1)
638 else:
639 # Normal acquisition for smaller requests
640 self.token_limiter.try_acquire("token", estimated_tokens)
642 try:
643 yield
644 finally:
645 # Update usage tracking
646 self.state.requests_used += 1
647 self.state.tokens_used += estimated_tokens
648 self.state.last_request_time = time.time()
650 def _get_retry_decorator(self):
651 """Get retry decorator with user-configured parameters."""
652 return retry(
653 stop=stop_after_attempt(self._user_max_retries),
654 wait=wait_exponential(multiplier=self._user_base_backoff, min=1, max=60),
655 retry=retry_if_exception_type((httpx.HTTPStatusError, httpx.RequestError, httpx.ReadTimeout, httpx.ConnectTimeout)),
656 )
658 async def request(
659 self,
660 method: str,
661 url: str,
662 *,
663 json: dict[str, Any] | None = None,
664 **kwargs: Any,
665 ) -> httpx.Response:
666 """Wrapper that applies retry decorator dynamically."""
667 try:
668 return await self._get_retry_decorator()(self._request_impl)(method, url, json=json, **kwargs)
669 except Exception as e:
670 # Check if this is a retry error wrapping a timeout
671 if hasattr(e, 'last_attempt') and e.last_attempt and e.last_attempt.exception():
672 original_exception = e.last_attempt.exception()
673 if isinstance(original_exception, (httpx.ReadTimeout, httpx.ConnectTimeout)):
674 # Enhance timeout error with helpful information
675 timeout_info = (
676 f"\n💡 Timeout Error Help:\n"
677 f" Current timeout: {self._user_timeout}s\n"
678 f" To increase timeout, use: ChatLimiter.for_model('{self.provider.value}', timeout={int(self._user_timeout + 60)})\n"
679 f" Or reduce batch concurrency if processing multiple requests\n"
680 f" Retries attempted: {self._user_max_retries}\n"
681 )
682 raise type(original_exception)(str(original_exception) + timeout_info) from e
684 # For direct timeout errors (shouldn't happen due to retry decorator but just in case)
685 if isinstance(e, (httpx.ReadTimeout, httpx.ConnectTimeout)):
686 timeout_info = (
687 f"\n💡 Timeout Error Help:\n"
688 f" Current timeout: {self._user_timeout}s\n"
689 f" To increase timeout, use: ChatLimiter.for_model('{self.provider.value}', timeout={int(self._user_timeout + 60)})\n"
690 f" Or reduce batch concurrency if processing multiple requests\n"
691 )
692 raise type(e)(str(e) + timeout_info) from e
694 # Re-raise any other exceptions unchanged
695 raise
697 async def _request_impl(
698 self,
699 method: str,
700 url: str,
701 *,
702 json: dict[str, Any] | None = None,
703 **kwargs: Any,
704 ) -> httpx.Response:
705 """
706 Make an async HTTP request with rate limiting.
708 Args:
709 method: HTTP method (GET, POST, etc.)
710 url: URL or path for the request
711 json: JSON data to send
712 **kwargs: Additional arguments passed to httpx
714 Returns:
715 HTTP response
717 Raises:
718 httpx.HTTPStatusError: For HTTP error responses
719 httpx.RequestError: For request errors
720 """
721 if not self._async_context_active:
722 raise RuntimeError("ChatLimiter must be used as an async context manager")
724 # Estimate tokens if we have JSON data
725 estimated_tokens = self._estimate_tokens(json or {})
727 # Acquire rate limits
728 async with self._acquire_rate_limits(estimated_tokens):
729 # Make the request
730 response = await self.async_client.request(method, url, json=json, **kwargs)
732 # Extract rate limit info
733 rate_limit_info = extract_rate_limit_info(
734 dict(response.headers), self.config
735 )
737 # Update our rate limits
738 if self.enable_adaptive_limits:
739 self._update_rate_limits(rate_limit_info)
741 # Handle rate limit errors
742 if response.status_code == 429:
743 self.state.consecutive_rate_limit_errors += 1
744 if rate_limit_info.retry_after:
745 await asyncio.sleep(rate_limit_info.retry_after)
746 else:
747 # Exponential backoff
748 backoff = self.config.base_backoff * (
749 2**self.state.consecutive_rate_limit_errors
750 )
751 await asyncio.sleep(min(backoff, self.config.max_backoff))
753 response.raise_for_status()
754 else:
755 # Reset consecutive errors on success
756 self.state.consecutive_rate_limit_errors = 0
758 return response
760 def request_sync(
761 self,
762 method: str,
763 url: str,
764 *,
765 json: dict[str, Any] | None = None,
766 **kwargs: Any,
767 ) -> httpx.Response:
768 """Wrapper that applies retry decorator dynamically."""
769 # For sync, we need to use the sync version of retry
770 retry_decorator = retry(
771 stop=stop_after_attempt(self._user_max_retries),
772 wait=wait_exponential(multiplier=self._user_base_backoff, min=1, max=60),
773 retry=retry_if_exception_type((httpx.HTTPStatusError, httpx.RequestError, httpx.ReadTimeout, httpx.ConnectTimeout)),
774 )
775 try:
776 return retry_decorator(self._request_sync_impl)(method, url, json=json, **kwargs)
777 except (httpx.ReadTimeout, httpx.ConnectTimeout) as e:
778 # Enhance timeout error with helpful information
779 timeout_info = (
780 f"\n💡 Timeout Error Help:\n"
781 f" Current timeout: {self._user_timeout}s\n"
782 f" To increase timeout, use: ChatLimiter.for_model('{self.provider.value}', timeout={int(self._user_timeout + 60)})\n"
783 f" Or reduce batch concurrency if processing multiple requests\n"
784 )
785 raise type(e)(str(e) + timeout_info) from e
787 def _request_sync_impl(
788 self,
789 method: str,
790 url: str,
791 *,
792 json: dict[str, Any] | None = None,
793 **kwargs: Any,
794 ) -> httpx.Response:
795 """
796 Make a sync HTTP request with rate limiting.
798 Args:
799 method: HTTP method (GET, POST, etc.)
800 url: URL or path for the request
801 json: JSON data to send
802 **kwargs: Additional arguments passed to httpx
804 Returns:
805 HTTP response
807 Raises:
808 httpx.HTTPStatusError: For HTTP error responses
809 httpx.RequestError: For request errors
810 """
811 if not self._sync_context_active:
812 raise RuntimeError("ChatLimiter must be used as a sync context manager")
814 # Estimate tokens if we have JSON data
815 estimated_tokens = self._estimate_tokens(json or {})
817 # Acquire rate limits
818 with self._acquire_rate_limits_sync(estimated_tokens):
819 # Make the request
820 response = self.sync_client.request(method, url, json=json, **kwargs)
822 # Extract rate limit info
823 rate_limit_info = extract_rate_limit_info(
824 dict(response.headers), self.config
825 )
827 # Update our rate limits
828 if self.enable_adaptive_limits:
829 self._update_rate_limits(rate_limit_info)
831 # Handle rate limit errors
832 if response.status_code == 429:
833 self.state.consecutive_rate_limit_errors += 1
834 if rate_limit_info.retry_after:
835 time.sleep(rate_limit_info.retry_after)
836 else:
837 # Exponential backoff
838 backoff = self.config.base_backoff * (
839 2**self.state.consecutive_rate_limit_errors
840 )
841 time.sleep(min(backoff, self.config.max_backoff))
843 response.raise_for_status()
844 else:
845 # Reset consecutive errors on success
846 self.state.consecutive_rate_limit_errors = 0
848 return response
850 def get_current_limits(self) -> dict[str, Any]:
851 """Get current rate limit information."""
852 return {
853 "provider": self.provider.value,
854 "request_limit": self.state.request_limit,
855 "token_limit": self.state.token_limit,
856 "requests_used": self.state.requests_used,
857 "tokens_used": self.state.tokens_used,
858 "last_request_time": self.state.last_request_time,
859 "last_limit_update": self.state.last_limit_update,
860 "consecutive_rate_limit_errors": self.state.consecutive_rate_limit_errors,
861 }
863 def reset_usage_tracking(self) -> None:
864 """Reset usage tracking counters."""
865 self.state.requests_used = 0
866 self.state.tokens_used = 0
867 self.state.consecutive_rate_limit_errors = 0
869 # High-level chat completion methods
871 async def chat_completion(
872 self,
873 model: str,
874 messages: list[Message],
875 max_tokens: int | None = None,
876 temperature: float | None = None,
877 top_p: float | None = None,
878 stop: str | list[str] | None = None,
879 stream: bool = False,
880 **kwargs: Any,
881 ) -> ChatCompletionResponse:
882 """
883 Make a high-level chat completion request.
885 Args:
886 model: The model to use for completion
887 messages: List of messages in the conversation
888 max_tokens: Maximum tokens to generate
889 temperature: Sampling temperature
890 top_p: Top-p sampling parameter
891 stop: Stop sequences
892 stream: Whether to stream the response
893 **kwargs: Additional provider-specific parameters
895 Returns:
896 ChatCompletionResponse with the completion result
898 Raises:
899 ValueError: If provider cannot be determined from model
900 httpx.HTTPStatusError: For HTTP error responses
901 httpx.RequestError: For request errors
902 """
903 if not self._async_context_active:
904 raise RuntimeError("ChatLimiter must be used as an async context manager")
906 # Create request object
907 request = ChatCompletionRequest(
908 model=model,
909 messages=messages,
910 max_tokens=max_tokens,
911 temperature=temperature,
912 top_p=top_p,
913 stop=stop,
914 stream=stream,
915 **kwargs
916 )
918 # Get the appropriate adapter
919 adapter = get_adapter(self.provider)
921 # Format the request for the provider
922 formatted_request = adapter.format_request(request)
924 # Make the HTTP request
925 response = await self.request(
926 "POST",
927 adapter.get_endpoint(),
928 json=formatted_request
929 )
931 # Parse the response
932 response_data = response.json()
933 return adapter.parse_response(response_data, request)
935 def chat_completion_sync(
936 self,
937 model: str,
938 messages: list[Message],
939 max_tokens: int | None = None,
940 temperature: float | None = None,
941 top_p: float | None = None,
942 stop: str | list[str] | None = None,
943 stream: bool = False,
944 **kwargs: Any,
945 ) -> ChatCompletionResponse:
946 """
947 Make a synchronous high-level chat completion request.
949 Args:
950 model: The model to use for completion
951 messages: List of messages in the conversation
952 max_tokens: Maximum tokens to generate
953 temperature: Sampling temperature
954 top_p: Top-p sampling parameter
955 stop: Stop sequences
956 stream: Whether to stream the response
957 **kwargs: Additional provider-specific parameters
959 Returns:
960 ChatCompletionResponse with the completion result
962 Raises:
963 ValueError: If provider cannot be determined from model
964 httpx.HTTPStatusError: For HTTP error responses
965 httpx.RequestError: For request errors
966 """
967 if not self._sync_context_active:
968 raise RuntimeError("ChatLimiter must be used as a sync context manager")
970 # Create request object
971 request = ChatCompletionRequest(
972 model=model,
973 messages=messages,
974 max_tokens=max_tokens,
975 temperature=temperature,
976 top_p=top_p,
977 stop=stop,
978 stream=stream,
979 **kwargs
980 )
982 # Get the appropriate adapter
983 adapter = get_adapter(self.provider)
985 # Format the request for the provider
986 formatted_request = adapter.format_request(request)
988 # Make the HTTP request
989 response = self.request_sync(
990 "POST",
991 adapter.get_endpoint(),
992 json=formatted_request
993 )
995 # Parse the response
996 response_data = response.json()
997 return adapter.parse_response(response_data, request)
999 # Convenience methods for different message types
1001 async def simple_chat(
1002 self,
1003 model: str,
1004 prompt: str,
1005 max_tokens: int | None = None,
1006 temperature: float | None = None,
1007 **kwargs: Any,
1008 ) -> str:
1009 """
1010 Simple chat completion that returns just the text response.
1012 Args:
1013 model: The model to use
1014 prompt: The user prompt
1015 max_tokens: Maximum tokens to generate
1016 temperature: Sampling temperature
1017 **kwargs: Additional parameters
1019 Returns:
1020 The text response from the model
1021 """
1022 messages = [Message(role=MessageRole.USER, content=prompt)]
1023 response = await self.chat_completion(
1024 model=model,
1025 messages=messages,
1026 max_tokens=max_tokens,
1027 temperature=temperature,
1028 **kwargs
1029 )
1031 if response.choices:
1032 return response.choices[0].message.content
1033 return ""
1035 def simple_chat_sync(
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 synchronous 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 = self.chat_completion_sync(
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 set_verbose_mode(self, verbose: bool) -> None:
1070 """Set verbose mode for detailed logging."""
1071 self._verbose_mode = verbose
1073 def _print_rate_limit_info(self) -> None:
1074 """Print current rate limit configuration."""
1075 print(f"\n=== Rate Limit Configuration for {self.provider.value.title()} ===")
1076 print(f"Provider: {self.provider.value}")
1077 print(f"Base URL: {self.config.base_url}")
1079 # Handle None values for limits
1080 if self.state.request_limit is not None:
1081 effective_req = self._effective_request_limit or "not calculated"
1082 print(f"Request Limit: {self.state.request_limit}/minute (effective: {effective_req}/minute)")
1083 else:
1084 print("Request Limit: Not yet discovered (will be fetched from API)")
1086 if self.state.token_limit is not None:
1087 effective_tok = self._effective_token_limit or "not calculated"
1088 print(f"Token Limit: {self.state.token_limit}/minute (effective: {effective_tok}/minute)")
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)