diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 1e65da5b867..d1fb03fd7e6 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -12,6 +12,7 @@ from collections.abc import Callable, Mapping, Sequence, Set from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, @@ -2020,7 +2021,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def reserve_tpm_tokens( self, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], estimated_tokens: int, parent_otel_span: Span | None = None, ) -> RateLimitResponse: @@ -2806,6 +2807,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return descriptors + @staticmethod + def _deduplicate_descriptors( + descriptors: Sequence[RateLimitDescriptor], + ) -> tuple[RateLimitDescriptor, ...]: + """ + Collapse descriptors repeating a (key, value) pair into one entry. + + A descriptor identifies one counter, and every consumer below charges + the request once per descriptor it is given. A repeat therefore + increments the sliding window twice and reserves TPM tokens twice + against that same counter, halving the limit the operator configured. + """ + by_identity: Final = MappingProxyType({(d["key"], d["value"]): d for d in descriptors}) + return tuple(by_identity.values()) + async def _check_model_has_recent_failures( self, model: str, @@ -3000,7 +3016,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _handle_rate_limit_error( self, response: RateLimitResponse, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], requested_model: str | None = None, ) -> None: """Handle rate limit exceeded by raising :class:`ProxyRateLimitError` (a 429).""" @@ -3427,7 +3443,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Create rate limit descriptors - descriptors: Final = self._create_rate_limit_descriptors( + assembled_descriptors: Final = self._create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=data, rpm_limit_type=rpm_limit_type, @@ -3440,23 +3456,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self._add_team_model_rate_limit_descriptor_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model, - descriptors=descriptors, + descriptors=assembled_descriptors, ) # Project Level Rate Limits self._add_project_model_rate_limit_descriptor_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model, - descriptors=descriptors, + descriptors=assembled_descriptors, ) self.add_project_io_token_rate_limit_descriptors_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model, - descriptors=descriptors, + descriptors=assembled_descriptors, ) # Org Level Rate Limits - descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) + assembled_descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) + + descriptors: Final = self._deduplicate_descriptors(assembled_descriptors) # Only check rate limits if we have descriptors with actual limits if descriptors: diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index fc0088b28d7..6e8d680826a 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -20,6 +20,7 @@ from litellm.caching.caching import DualCache from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + CHECK_AND_INCREMENT_BY_N_SCRIPT, PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, RequestRateLimiterStash, @@ -6004,3 +6005,243 @@ async def test_success_hook_leaves_stash_untouched_for_non_batch_responses(): data={}, user_api_key_dict=user, response=ModelResponse(usage=Usage(total_tokens=5)) ) assert get_request_stash().batch_enqueued_reservation == reservation + + +def _team_model_limit_auth(api_key: str, team_id: str, **limits) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key=hash_token(api_key), + team_id=team_id, + team_metadata=dict(limits), + ) + + +async def _counter_value(handler, counter_key: str) -> int: + raw = await handler.internal_usage_cache.async_get_cache( + key=counter_key, litellm_parent_otel_span=None, local_only=True + ) + return int(raw or 0) + + +class _RecordingUsageCache(InternalUsageCache): + """ + Injected in place of the handler's usage cache so a test can see which counter + keys one request charges, without replacing methods on the handler itself. + """ + + def __init__(self, dual_cache: DualCache): + super().__init__(dual_cache) + self.batched_key_reads: List[List[str]] = [] + + async def async_batch_get_cache( + self, + keys, + parent_otel_span=None, + local_only: bool = False, + ): + self.batched_key_reads.append([k for k in keys if k is not None]) + return await super().async_batch_get_cache( + keys=keys, parent_otel_span=parent_otel_span, local_only=local_only + ) + + +def test_deduplicate_descriptors_collapses_repeats_keeping_limits_v3(): + """ + The dedup helper keeps one entry per (key, value), preserves its rate_limit + and leaves distinct descriptors and their order alone. + """ + limits = {"requests_per_unit": 10, "tokens_per_unit": 100000, "window_size": 60} + team = {"key": "model_per_team", "value": "team-dup:gpt-4", "rate_limit": limits} + key_scoped = {"key": "model_per_key", "value": "sk-hash:gpt-4", "rate_limit": limits} + + deduped = _PROXY_MaxParallelRequestsHandler._deduplicate_descriptors( + [team, key_scoped, dict(team)] + ) + + assert [(d["key"], d["value"]) for d in deduped] == [ + ("model_per_team", "team-dup:gpt-4"), + ("model_per_key", "sk-hash:gpt-4"), + ] + assert deduped[0]["rate_limit"]["requests_per_unit"] == 10 + assert deduped[0]["rate_limit"]["tokens_per_unit"] == 100000 + + +@pytest.mark.asyncio +async def test_model_per_team_counter_charged_once_from_team_metadata_v3(): + """ + Regression test: team_metadata model limits must charge their counter once. + + They were appended twice, once inside _create_rate_limit_descriptors and once + by _add_team_model_rate_limit_descriptor_from_metadata, so the same counter key + landed in one batch read twice and every request incremented it by 2. + + The usage cache is injected rather than patched onto the handler, so the real + rate-limit path runs and the test observes it through a real collaborator. + """ + usage_cache = _RecordingUsageCache(DualCache()) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=usage_cache) + + user_api_key_dict = _team_model_limit_auth( + "sk-team-dup", "team-dup", model_rpm_limit={"gpt-4": 10} + ) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=usage_cache.dual_cache, + data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]}, + call_type="", + ) + + counter_key = handler.create_rate_limit_keys( + "model_per_team", "team-dup:gpt-4", "requests" + ) + charged = [read.count(counter_key) for read in usage_cache.batched_key_reads] + assert any( + count > 0 for count in charged + ), f"model_per_team was never charged; reads={usage_cache.batched_key_reads}" + assert all( + count <= 1 for count in charged + ), f"model_per_team charged more than once in a single read: {charged}" + + +@pytest.mark.asyncio +async def test_model_per_team_rpm_counter_increments_once_per_request_v3(): + """ + Regression test: one request must add exactly 1 to the team per-model RPM + counter. The duplicate descriptor put the same counter key into + keys_to_fetch twice, so the sliding window incremented it by 2 and the + configured RPM was enforced at half its value. + """ + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + assert ( + handler.batch_rate_limiter_script is None + ), "Test premise: no Redis, so the in-memory sliding window does the increment" + + user_api_key_dict = _team_model_limit_auth( + "sk-team-count", "team-count", model_rpm_limit={"gpt-4": 100} + ) + counter_key = handler.create_rate_limit_keys( + "model_per_team", "team-count:gpt-4", "requests" + ) + + for expected_count in (1, 2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]}, + call_type="", + ) + assert await _counter_value(handler, counter_key) == expected_count + + +class _ScriptRegisteringRedisCache: + """ + Stand-in for RedisCache carrying only what the handler's constructor uses: + it registers Lua scripts. Injecting it through DualCache makes the handler + wire up its own Redis reservation path, so the test drives that path without + assigning anything onto the handler. + """ + + def __init__(self, store: Dict[str, int]): + self._store = store + + def async_register_script(self, script: str): + if script is CHECK_AND_INCREMENT_BY_N_SCRIPT: + return _fake_check_and_increment_script(self._store) + return None + + +def _fake_check_and_increment_script(store: Dict[str, int]): + """ + Stand-in for CHECK_AND_INCREMENT_BY_N_LUA with the Redis INCRBY semantics + the real script has. The in-memory fallback snapshots every counter before + writing any, so a repeated descriptor there silently overwrites instead of + accumulating; only this path shows the double reservation. + """ + + async def script(keys: List[str], args: List[int]) -> List[int]: + results: List[int] = [0] + for i in range(0, len(keys), 2): + window_key = keys[i] + counter_key = keys[i + 1] + increment = int(args[(i // 2) * 4 + 1]) + if window_key in store: + store[counter_key] = store.get(counter_key, 0) + increment + else: + store[window_key] = 0 + store[counter_key] = increment + results.extend([store[counter_key], store[window_key]]) + return results + + return script + + +@pytest.mark.asyncio +async def test_model_per_team_tpm_reserved_once_per_request_v3(): + """ + Regression test: the atomic TPM reservation must charge the team per-model + token counter once. reserve_tpm_tokens emits one increment per descriptor + and the Redis path INCRBYs each in turn, so the duplicate descriptor + reserved the estimate twice and halved the effective TPM. + """ + redis_store: Dict[str, int] = {} + local_cache = DualCache(redis_cache=_ScriptRegisteringRedisCache(redis_store)) + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + assert handler.tpm_reservation_enabled, "Test premise: reservation path is on" + + user_api_key_dict = _team_model_limit_auth( + "sk-team-tpm", "team-tpm", model_tpm_limit={"gpt-4": 100000} + ) + tokens_key = handler.create_rate_limit_keys( + "model_per_team", "team-tpm:gpt-4", "tokens" + ) + + pre_call_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 5, + } + expected_reservation = handler._estimate_tokens_for_request(data=pre_call_data) + assert expected_reservation > 0, "Test premise: something must be reserved" + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=dict(pre_call_data), + call_type="", + ) + + assert redis_store[tokens_key] == expected_reservation + + +@pytest.mark.asyncio +async def test_model_per_key_rpm_counter_increments_once_per_request_v3(): + """ + Guard that the model_per_key path, which was never duplicated, still + charges its counter exactly once after the model_per_team dedup. + """ + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + api_key = hash_token("sk-key-count") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"model_rpm_limit": {"gpt-4": 100}}, + ) + counter_key = handler.create_rate_limit_keys( + "model_per_key", f"{api_key}:gpt-4", "requests" + ) + + for expected_count in (1, 2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]}, + call_type="", + ) + assert await _counter_value(handler, counter_key) == expected_count