This commit is contained in:
Dan Brima 2026-08-27 20:17:30 -05:00 • committed by GitHub
commit c81e9f1b07
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 266 additions and 7 deletions

View file

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

View file

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