diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 0f884387673..7c02ac308f3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -1,5 +1,6 @@ import asyncio import sys +from collections.abc import Callable from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn @@ -36,6 +37,10 @@ else: InternalUsageCache = Any +def _precise_minute(now: datetime) -> str: + return now.strftime("%Y-%m-%d-%H-%M") + + def _response_total_tokens(response_obj: object) -> int: if not isinstance(response_obj, (ModelResponse, EmbeddingResponse, TextCompletionResponse)): return 0 @@ -54,8 +59,9 @@ class CacheObject(TypedDict): class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Class variables or attributes - def __init__(self, internal_usage_cache: InternalUsageCache): + def __init__(self, internal_usage_cache: InternalUsageCache, *, clock: Callable[[], datetime] = datetime.now): self.internal_usage_cache = internal_usage_cache + self._clock: Final = clock def print_verbose(self, print_statement) -> None: try: @@ -149,7 +155,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): def time_to_next_minute(self) -> float: # Get the current time - now: Final = datetime.now() + now: Final = self._clock() # Calculate the next minute next_minute: Final = (now + timedelta(minutes=1)).replace(second=0, microsecond=0) @@ -306,10 +312,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) _model = data.get("model", None) - current_date: Final = datetime.now().strftime("%Y-%m-%d") - current_hour: Final = datetime.now().strftime("%H") - current_minute: Final = datetime.now().strftime("%M") - precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}" + precise_minute: Final = _precise_minute(self._clock()) cache_objects: Final[CacheObject] = await self.get_all_cache_objects( current_global_requests=( @@ -538,10 +541,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=litellm_parent_otel_span, ) - current_date: Final = datetime.now().strftime("%Y-%m-%d") - current_hour: Final = datetime.now().strftime("%H") - current_minute: Final = datetime.now().strftime("%M") - precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}" + precise_minute: Final = _precise_minute(self._clock()) total_tokens: int = _response_total_tokens(response_obj) @@ -737,10 +737,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=litellm_parent_otel_span, ) - current_date: Final = datetime.now().strftime("%Y-%m-%d") - current_hour: Final = datetime.now().strftime("%H") - current_minute: Final = datetime.now().strftime("%M") - precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}" + precise_minute: Final = _precise_minute(self._clock()) request_count_api_key: Final = f"{user_api_key}::{precise_minute}::request_count" @@ -813,10 +810,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): Retrieve the key's remaining rate limits. """ api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) - current_date: Final = datetime.now().strftime("%Y-%m-%d") - current_hour: Final = datetime.now().strftime("%H") - current_minute: Final = datetime.now().strftime("%M") - precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}" + precise_minute: Final = _precise_minute(self._clock()) request_count_api_key: Final = f"{api_key}::{precise_minute}::request_count" current: Final[CurrentItemRateLimit | None] = await self.internal_usage_cache.async_get_cache( key=request_count_api_key, diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index 1772f41a9d3..a2e43b3bc9c 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -2,22 +2,49 @@ Unit Tests for the max parallel request limiter v1 for the proxy """ +import itertools +from collections.abc import Callable, Iterator from datetime import datetime +from typing import Final import pytest from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.parallel_request_limiter import ( PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, hash_token -from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +from litellm.types.utils import EmbeddingResponse, ModelResponse, TextCompletionResponse, Usage + +FROZEN_INSTANT: Final = datetime(2026, 1, 31, 23, 59, 30) +LAST_MICROSECOND_OF_JANUARY: Final = datetime(2026, 1, 31, 23, 59, 59, 999999) +FIRST_MICROSECOND_OF_FEBRUARY: Final = datetime(2026, 2, 1, 0, 0, 0, 1) +LAST_MINUTE_OF_JANUARY: Final = "2026-01-31-23-59" +FIRST_MINUTE_OF_FEBRUARY: Final = "2026-02-01-00-00" +TORN_MINUTE_OF_JANUARY: Final = "2026-01-31-00-00" + + +def _frozen_clock() -> datetime: + return FROZEN_INSTANT + + +def _clock_reading(instants: Iterator[datetime]) -> Callable[[], datetime]: + return lambda: next(instants) + + +def _clock_rolling_over_after_first_read() -> Callable[[], datetime]: + return _clock_reading( + itertools.chain([LAST_MICROSECOND_OF_JANUARY], itertools.repeat(FIRST_MICROSECOND_OF_FEBRUARY)) + ) @pytest.mark.asyncio async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token(): - handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()), clock=_frozen_clock + ) session = UserAPIKeyAuth( api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", user_id="alice", @@ -30,7 +57,7 @@ async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_t user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" ) - precise_minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + precise_minute = FROZEN_INSTANT.strftime("%Y-%m-%d-%H-%M") counted = await handler.internal_usage_cache.async_get_cache( key=f"cli-session-alice::{precise_minute}::request_count", litellm_parent_otel_span=None ) @@ -63,13 +90,10 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ end_user_id = "customer-1" parallel_request_handler = PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) + internal_usage_cache=InternalUsageCache(DualCache()), clock=_frozen_clock ) - current_date = datetime.now().strftime("%Y-%m-%d") - current_hour = datetime.now().strftime("%H") - current_minute = datetime.now().strftime("%M") - precise_minute = f"{current_date}-{current_hour}-{current_minute}" + precise_minute = FROZEN_INSTANT.strftime("%Y-%m-%d-%H-%M") scope_ids = [_api_key, user_id, team_id, end_user_id] for scope_id in scope_ids: @@ -94,8 +118,8 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ await parallel_request_handler.async_log_success_event( kwargs=kwargs, response_obj=response_obj, - start_time=datetime.now(), - end_time=datetime.now(), + start_time=FROZEN_INSTANT, + end_time=FROZEN_INSTANT, ) for scope_id in scope_ids: @@ -107,3 +131,162 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ f"expected 50 tokens counted for {scope_id}, " f"got {current['current_tpm']}" ) + + +@pytest.mark.asyncio +async def test_a_pre_call_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read() + ) + session: Final = UserAPIKeyAuth(api_key="sk-torn-pre", max_parallel_requests=5) + api_key: Final = session.api_key + + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None + ) == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1} + for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY): + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None + ) is None + + +@pytest.mark.asyncio +async def test_a_success_event_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read() + ) + api_key: Final = hash_token("sk-torn-success") + await internal_usage_cache.async_set_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", + value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1}, + litellm_parent_otel_span=None, + ) + + await handler.async_log_success_event( + kwargs={ + "litellm_params": { + "metadata": {"user_api_key": api_key, "user_api_key_model_max_budget": {}} + } + }, + response_obj=ModelResponse(usage=Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7)), + start_time=LAST_MICROSECOND_OF_JANUARY, + end_time=FIRST_MICROSECOND_OF_FEBRUARY, + ) + + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None + ) == {"current_requests": 0, "current_tpm": 7, "current_rpm": 1} + for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY): + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None + ) is None + + +@pytest.mark.asyncio +async def test_a_failure_event_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read() + ) + api_key: Final = hash_token("sk-torn-failure") + await internal_usage_cache.async_set_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", + value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1}, + litellm_parent_otel_span=None, + ) + + await handler.async_log_failure_event( + kwargs={ + "litellm_params": {"metadata": {"user_api_key": api_key}}, + "exception": Exception("upstream boom"), + }, + response_obj=None, + start_time=LAST_MICROSECOND_OF_JANUARY, + end_time=FIRST_MICROSECOND_OF_FEBRUARY, + ) + + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None + ) == {"current_requests": 0, "current_tpm": 0, "current_rpm": 1} + for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY): + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None + ) is None + + +@pytest.mark.asyncio +async def test_a_post_call_headers_read_across_a_minute_rollover_uses_the_bucket_of_its_first_clock_read(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read() + ) + user_api_key_dict: Final = UserAPIKeyAuth(api_key="sk-torn-post", rpm_limit=5, tpm_limit=100) + api_key: Final = user_api_key_dict.api_key + await internal_usage_cache.async_set_cache( + key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", + value={"current_requests": 1, "current_tpm": 10, "current_rpm": 1}, + litellm_parent_otel_span=None, + ) + response: Final = ModelResponse() + response._hidden_params = {} + + await handler.async_post_call_success_hook( + data={"model": "gpt-4o-mini"}, + user_api_key_dict=user_api_key_dict, + response=response, + ) + + assert response._hidden_params["additional_headers"] == { + "x-ratelimit-remaining-requests": 4, + "x-ratelimit-limit-requests": 5, + "x-ratelimit-remaining-tokens": 90, + "x-ratelimit-limit-tokens": 100, + } + + +@pytest.mark.asyncio +async def test_a_request_in_one_minute_is_not_counted_by_a_pre_call_in_the_next_minute(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, + clock=_clock_reading(iter([datetime(2026, 3, 10, 12, 0, 0), datetime(2026, 3, 10, 12, 1, 0)])), + ) + session: Final = UserAPIKeyAuth(api_key="sk-minute-reset", rpm_limit=1) + api_key: Final = session.api_key + + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + + assert await internal_usage_cache.async_get_cache( + key=f"{api_key}::2026-03-10-12-01::request_count", litellm_parent_otel_span=None + ) == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1} + + +@pytest.mark.asyncio +async def test_retry_after_is_the_seconds_until_the_next_minute_of_the_injected_clock(): + internal_usage_cache: Final = InternalUsageCache(DualCache()) + handler: Final = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, + clock=_clock_reading(itertools.repeat(datetime(2026, 3, 10, 12, 0, 45, 500000))), + ) + session: Final = UserAPIKeyAuth(api_key="sk-retry-after", rpm_limit=1) + + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + + assert exc_info.value.headers["retry-after"] == "14.5" diff --git a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py index 16b5406bc21..187100c24d8 100644 --- a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -33,6 +33,7 @@ fallback path (unknown model, missing model) for every limiter. """ import sys +from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -305,7 +306,9 @@ async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit( Trip the per-key RPM cap and assert the raised exception carries ``model`` / ``llm_provider`` resolved from ``data["model"]``. """ - handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0) + ) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-test", max_parallel_requests=10, @@ -404,7 +407,9 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): When ``data["model"]`` is unparseable, the resolver falls back to ``litellm_proxy`` — and crucially does *not* leak a secondary exception. """ - handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0) + ) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-unknown", max_parallel_requests=10, @@ -438,7 +443,9 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): @pytest.mark.asyncio async def test_parallel_request_limiter_v1_missing_model_falls_back(): - handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0) + ) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-no-model", max_parallel_requests=10, diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 1aed60ee5e2..cc7cce937d7 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -6951,19 +6951,13 @@ class TestPreCallWithFallbacksOnLocalRateLimit: primary_model = "gpt-4" fallback_model = "gpt-3.5-turbo" - # Freeze the limiter's clock so the per-minute counter key is stable and - # the pre-seeded counter is guaranteed to be the one it reads. - class _FrozenClock(datetime.datetime): - @classmethod - def now(cls, tz=None): - return cls(2026, 1, 1, 12, 30, 0) - precise_minute = "2026-01-01-12-30" # Real per-key per-model TPM limiter + a key carrying the customer's # `model_tpm_limit` metadata (only the primary is capped). limiter = PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) + internal_usage_cache=InternalUsageCache(DualCache()), + clock=lambda: datetime.datetime(2026, 1, 1, 12, 30, 0), ) user_api_key_dict = UserAPIKeyAuth( api_key="sk-lit3890", @@ -7006,30 +7000,27 @@ class TestPreCallWithFallbacksOnLocalRateLimit: mock_router = MagicMock() mock_router.fallbacks = [{primary_model: [fallback_model]}] - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock + with patch.object( + processor, + "common_processing_pre_call_logic", + side_effect=real_limiter_pre_call, ): - with patch.object( - processor, - "common_processing_pre_call_logic", - side_effect=real_limiter_pre_call, - ): - data, logging_obj = await processor._pre_call_with_fallbacks( - request=MagicMock(), - general_settings={}, - proxy_logging_obj=MagicMock(), - user_api_key_dict=user_api_key_dict, - version=None, - proxy_config=MagicMock(), - user_model=None, - user_temperature=None, - user_request_timeout=None, - user_max_tokens=None, - user_api_base=None, - model=primary_model, - route_type="acompletion", - llm_router=mock_router, - ) + data, logging_obj = await processor._pre_call_with_fallbacks( + request=MagicMock(), + general_settings={}, + proxy_logging_obj=MagicMock(), + user_api_key_dict=user_api_key_dict, + version=None, + proxy_config=MagicMock(), + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + model=primary_model, + route_type="acompletion", + llm_router=mock_router, + ) # The capped primary tripped the real limiter, and the fallback (which # has no per-model cap) served the request — no 429 to the client. @@ -7038,19 +7029,16 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Sanity-check the premise: the limiter genuinely raises a # ProxyRateLimitError for the capped primary under the frozen clock. - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock - ): - with pytest.raises(ProxyRateLimitError): - await limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=DualCache(), - data={ - "model": primary_model, - "messages": [{"role": "user", "content": "hi"}], - }, - call_type="acompletion", - ) + with pytest.raises(ProxyRateLimitError): + await limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={ + "model": primary_model, + "messages": [{"role": "user", "content": "hi"}], + }, + call_type="acompletion", + ) @staticmethod def _v3_limiter_rig(