test(proxy): inject the clock into the v1 parallel request limiter and pin it in its tests (#45522)

Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 00:16:03 -07:00 • committed by GitHub
parent 488594e03f
commit ee47da334d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 247 additions and 75 deletions

View file

@ -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,

View file

@ -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"

View file

@ -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,

View file

@ -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(