mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge 6db862d7ae into bb72815e70
This commit is contained in:
commit
c81e9f1b07
2 changed files with 266 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue