mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(rate-limiting): refresh the global hook's own concurrency key ttl on every admission
veria-ai finding on the previous commit: TAG_RL_CHECK_AND_INCR_SCRIPT is shared with model_based_tag_rate_limits_hook, but only that hook's own Redis and in-memory call sites were updated to carry refresh_ttl through. global_tag_rate_limits_hook's own _check_and_increment_one still called the script with three args and never passed refresh_ttl to InMemoryCache, so a global concurrency bucket's ttl stayed fixed from its first admission and could expire mid-flight under sustained traffic, admitting past the cap.
This commit is contained in:
parent
8674da5c01
commit
fee2c41359
2 changed files with 122 additions and 8 deletions
|
|
@ -346,10 +346,12 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
return built
|
||||
|
||||
async def _check_and_increment_one(
|
||||
self, cache: InternalUsageCache, key: str, limit: float, increment: float, ttl: int
|
||||
self, cache: InternalUsageCache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool
|
||||
) -> tuple[bool, float]:
|
||||
if self._check_and_incr_script is not None:
|
||||
raw: Final = await self._check_and_incr_script(keys=(key,), args=(limit, increment, ttl))
|
||||
raw: Final = await self._check_and_incr_script(
|
||||
keys=(key,), args=(limit, increment, ttl, 1 if refresh_ttl else 0)
|
||||
)
|
||||
return bool(raw[0]), float(raw[1])
|
||||
async with self._lock:
|
||||
current_value: Final = await cache.async_get_cache(key=key, litellm_parent_otel_span=None)
|
||||
|
|
@ -357,7 +359,9 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
if current + increment > limit:
|
||||
return False, current
|
||||
new_value: Final = current + increment
|
||||
await cache.async_set_cache(key=key, value=new_value, ttl=ttl, litellm_parent_otel_span=None)
|
||||
await cache.async_set_cache(
|
||||
key=key, value=new_value, ttl=ttl, refresh_ttl=refresh_ttl, litellm_parent_otel_span=None
|
||||
)
|
||||
return True, new_value
|
||||
|
||||
async def _decrement_floor_zero(self, cache: InternalUsageCache, key: str, delta: float) -> None:
|
||||
|
|
@ -371,7 +375,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
|
||||
async def _atomic_check_and_increment(
|
||||
self,
|
||||
checks: Sequence[tuple[InternalUsageCache, str, float, float, int]],
|
||||
checks: Sequence[tuple[InternalUsageCache, str, float, float, int, bool]],
|
||||
) -> tuple[int | None, tuple[float, ...]]:
|
||||
"""All-or-nothing atomic admission across `checks` -- see
|
||||
`model_based_tag_rate_limits_hook._PROXY_ModelBasedTagRateLimitsHook._atomic_check_and_increment`'s
|
||||
|
|
@ -382,10 +386,10 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
if not checks:
|
||||
return None, ()
|
||||
admitted_values: Final = [] # mutable-ok: sequential async accumulator, discardable on early rejection
|
||||
for index, (cache, key, limit, increment, ttl) in enumerate(checks):
|
||||
for index, (cache, key, limit, increment, ttl, refresh_ttl) in enumerate(checks):
|
||||
admitted = False
|
||||
try:
|
||||
admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl)
|
||||
admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl)
|
||||
finally:
|
||||
if not admitted:
|
||||
await self._refund_admitted(checks, up_to_index=index)
|
||||
|
|
@ -396,10 +400,10 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
return None, tuple(admitted_values)
|
||||
|
||||
async def _refund_admitted(
|
||||
self, checks: Sequence[tuple[InternalUsageCache, str, float, float, int]], up_to_index: int
|
||||
self, checks: Sequence[tuple[InternalUsageCache, str, float, float, int, bool]], up_to_index: int
|
||||
) -> None:
|
||||
for refund_index in range(up_to_index):
|
||||
refund_cache, refund_key, _limit, refund_increment, _ttl = checks[refund_index]
|
||||
refund_cache, refund_key, _limit, refund_increment, _ttl, _refresh_ttl = checks[refund_index]
|
||||
try:
|
||||
await self._decrement_floor_zero(refund_cache, refund_key, -refund_increment)
|
||||
except Exception as e: # noqa: BLE001 - one failed refund must not block refunding the rest
|
||||
|
|
@ -631,6 +635,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
|
|||
)
|
||||
else 1.0,
|
||||
self._ttl_for(check.unit, check.entry),
|
||||
check.unit == "concurrency",
|
||||
)
|
||||
for partition, check in zip(atomic_partitions, atomic_checks)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,9 @@ Unit tests for the global-scope, model-independent tag rate limiter.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
|
@ -48,6 +51,18 @@ def _data(tags: list[str], call_id: str = "call-1") -> dict:
|
|||
return {"metadata": {"tags": tags}, "litellm_call_id": call_id}
|
||||
|
||||
|
||||
def _redis_hook(time_controller: TimeController):
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
redis_host = os.getenv("REDIS_HOST")
|
||||
redis_port = os.getenv("REDIS_PORT")
|
||||
if not redis_host or not redis_port:
|
||||
pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set")
|
||||
redis_cache = RedisCache(host=redis_host, port=int(redis_port), password=os.getenv("REDIS_PASSWORD"))
|
||||
dual_cache = DualCache(redis_cache=redis_cache)
|
||||
return _PROXY_GlobalTagRateLimitsHook(internal_usage_cache=dual_cache, time_provider=time_controller.now), redis_cache
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# No-op when unconfigured
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -1359,3 +1374,97 @@ async def test_rejection_detail_does_not_disclose_the_resolved_tag_value(time_co
|
|||
)
|
||||
|
||||
assert "tag_value" not in exc_info.value.detail
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Concurrency ttl refresh -- veria-ai finding: this hook's own atomic checks
|
||||
# never carried refresh_ttl through to TAG_RL_CHECK_AND_INCR_SCRIPT or
|
||||
# InMemoryCache.set_cache, even though it shares that script with
|
||||
# model_based_tag_rate_limits_hook (whose own call sites got fixed first)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_backed_concurrency_ttl_refreshes_on_every_admission(time_controller):
|
||||
"""
|
||||
A concurrency bucket isn't epoch-windowed like requests/tokens/dollars --
|
||||
its ttl exists only as a crash-safety net for a reservation whose
|
||||
explicit release never runs -- so a still-active bucket receiving
|
||||
continuous admissions must keep extending that ttl, or it expires
|
||||
mid-flight under sustained traffic, silently admitting past the cap.
|
||||
"""
|
||||
hook, redis_cache = _redis_hook(time_controller)
|
||||
try:
|
||||
await redis_cache.ping()
|
||||
except Exception as e:
|
||||
pytest.skip(f"Redis connection failed: {e!s}")
|
||||
|
||||
key = f"{{tag_rl:test:global-ttl-refresh:{uuid.uuid4().hex}}}:inflight"
|
||||
cache = hook.internal_usage_cache
|
||||
try:
|
||||
admitted, _ = await hook._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
|
||||
assert admitted
|
||||
ttl_after_first_admission = await redis_cache.init_async_client().ttl(key)
|
||||
assert ttl_after_first_admission > 0
|
||||
|
||||
await asyncio.sleep(2)
|
||||
|
||||
admitted, _ = await hook._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
|
||||
assert admitted
|
||||
ttl_after_second_admission = await redis_cache.init_async_client().ttl(key)
|
||||
assert ttl_after_second_admission >= 2
|
||||
finally:
|
||||
await redis_cache.async_delete_cache(key=key)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_in_memory_concurrency_ttl_refreshes_on_every_admission(time_controller):
|
||||
"""
|
||||
In-memory mirror of the Redis test above: InMemoryCache.allow_ttl_override
|
||||
leaves a still-live ttl untouched, so without refresh_ttl reaching
|
||||
set_cache a concurrency counter's expiry stayed fixed from its first
|
||||
admission even under sustained traffic.
|
||||
"""
|
||||
hook = _make_hook(time_controller)
|
||||
cache = hook.internal_usage_cache
|
||||
in_memory_cache = cache.dual_cache.in_memory_cache
|
||||
key = f"tag_rl:test:global-in-memory-ttl-refresh:{uuid.uuid4().hex}"
|
||||
|
||||
admitted, _ = await hook._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
|
||||
assert admitted
|
||||
ttl_after_first_admission = in_memory_cache.ttl_dict[key]
|
||||
|
||||
admitted, _ = await hook._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
|
||||
assert admitted
|
||||
ttl_after_second_admission = in_memory_cache.ttl_dict[key]
|
||||
|
||||
assert ttl_after_second_admission > ttl_after_first_admission
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrency_limit_admission_refreshes_ttl_end_to_end(time_controller, monkeypatch):
|
||||
"""
|
||||
End-to-end regression through async_pre_call_hook itself (not just the
|
||||
low-level _check_and_increment_one helper above): a concurrency entry's
|
||||
bucket key must carry a live ttl after admission, proving refresh_ttl is
|
||||
actually wired from the classified check through to the atomic batch,
|
||||
not just present on the helper's own signature.
|
||||
"""
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"global_tag_rate_limits",
|
||||
{
|
||||
"concurrency_limits": {
|
||||
"limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}]
|
||||
}
|
||||
},
|
||||
)
|
||||
hook = _make_hook(time_controller)
|
||||
await hook.async_pre_call_hook(
|
||||
user_api_key_dict=_key(), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion"
|
||||
)
|
||||
|
||||
in_memory_cache = hook.internal_usage_cache.dual_cache.in_memory_cache
|
||||
inflight_keys = [key for key in in_memory_cache.ttl_dict if key.endswith(":inflight")]
|
||||
assert len(inflight_keys) == 1
|
||||
assert in_memory_cache.ttl_dict[inflight_keys[0]] > time.time()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue