diff --git a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index 2f1cf21f3af..e68dbc81208 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -934,12 +934,16 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] 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]: """Single-key atomic check-and-increment. Always one key per Lua - call -- see TAG_RL_CHECK_AND_INCR_SCRIPT's module docstring for why.""" + call -- see TAG_RL_CHECK_AND_INCR_SCRIPT's module docstring for why, + and for why `refresh_ttl` must be True for a concurrency key and + False for a requests key.""" 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: @@ -962,7 +966,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] 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 across every (cache, key, limit, increment, ttl) in @@ -1019,10 +1023,10 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # accumulated so far in favor of refunding and returning early, so # this can't be expressed as a one-shot comprehension. admitted_values: Final = [] # mutable-ok: sequential async accumulator, discardable on early rejection; see comment above - 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: # Runs on a normal rejection (admitted stays False) and on # any exception/cancellation from the awaited call above @@ -1040,10 +1044,10 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] 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 @@ -1150,6 +1154,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # with nothing to replace it. 0.0 if configured_limit.unit == "requests" and key in stale_request_keys else 1.0, self._ttl_for(configured_limit), + configured_limit.unit == "concurrency", ) for partition, (configured_limit, _tag_value, key) in zip(atomic_partitions, atomic_checks) ) diff --git a/litellm/proxy/hooks/tag_rate_limits_shared.py b/litellm/proxy/hooks/tag_rate_limits_shared.py index ad8e7bfe29d..f5a79f7de17 100644 --- a/litellm/proxy/hooks/tag_rate_limits_shared.py +++ b/litellm/proxy/hooks/tag_rate_limits_shared.py @@ -93,19 +93,36 @@ CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600 # `atomic_check_and_increment_by_n` in parallel_request_limiter_v3.py, applied # per-key instead of per-descriptor since each key already is one hash-tag # group by construction. +# +# refresh_ttl (ARGV[4]) distinguishes the two callers of this script: +# "requests" is an epoch-bucketed fixed window, whose TTL must be set once +# (at first write) and never extended, or the bucket outlives the epoch it's +# meant to reset at. "concurrency" is not windowed at all -- its TTL exists +# purely as a crash-safety net for a reservation whose explicit release never +# runs -- so a still-active bucket must keep pushing that TTL out on every +# admission, or a long-lived burst of continuous traffic expires the whole +# counter mid-flight (silently admitting past the cap, and letting a release +# for a since-reset counter decrement an unrelated, newer cohort). TAG_RL_CHECK_AND_INCR_SCRIPT: Final = """ local key = KEYS[1] local limit = tonumber(ARGV[1]) local increment = tonumber(ARGV[2]) local ttl = tonumber(ARGV[3]) +local refresh_ttl = tonumber(ARGV[4]) local current = tonumber(redis.call('GET', key) or 0) if current + increment > limit then return { 0, current } end local new_value = redis.call('INCRBY', key, increment) -local current_ttl = redis.call('TTL', key) -if current_ttl == -1 and ttl > 0 then - redis.call('EXPIRE', key, ttl) +if ttl > 0 then + if refresh_ttl == 1 then + redis.call('EXPIRE', key, ttl) + else + local current_ttl = redis.call('TTL', key) + if current_ttl == -1 then + redis.call('EXPIRE', key, ttl) + end + end end return { 1, new_value } """ diff --git a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index 33edf7519a3..4a0f46e72aa 100644 --- a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -3340,6 +3340,49 @@ async def test_redis_backed_token_admission_sees_increments_the_in_memory_cache_ await redis_cache.async_delete_cache(key=token_key) +@pytest.mark.asyncio +async def test_redis_backed_concurrency_ttl_refreshes_on_every_admission(time_controller): + """ + Bugbot finding: TAG_RL_CHECK_AND_INCR_SCRIPT only ran EXPIRE when a key + had no TTL at all, so a concurrency counter's expiry was fixed from its + first admission and never pushed out by later ones. 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. + """ + limiter, redis_cache = _redis_limiter(time_controller) + try: + await redis_cache.ping() + except Exception as e: + pytest.skip(f"Redis connection failed: {e!s}") + + key = f"{{tag_rl:test:ttl-refresh:{uuid.uuid4().hex}}}:inflight" + cache = limiter.internal_usage_cache + try: + # A short, fixed ttl (bypassing _ttl_for's 3600s safety floor, which + # would make a real-time before/after comparison too slow to assert + # on deterministically) with refresh_ttl=True, matching how a + # concurrency check is actually admitted. + admitted, _ = await limiter._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.redis_async_client.ttl(key) + assert ttl_after_first_admission > 0 + + await asyncio.sleep(2) + + # A second admission on the same still-live key, most of the way + # through the first admission's ttl, must push the ttl back out to + # the full window again, not leave it counting down toward zero. + admitted, _ = await limiter._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.redis_async_client.ttl(key) + assert ttl_after_second_admission >= 2 + finally: + await redis_cache.async_delete_cache(key=key) + + # --------------------------------------------------------------------------- # team_public_model_name alias -- index lookup must not miss # --------------------------------------------------------------------------- @@ -3843,9 +3886,9 @@ async def test_refund_failure_on_one_key_does_not_block_others_or_raise(time_con failing_index, values = await flaky._atomic_check_and_increment( [ - (flaky.internal_usage_cache, failing_key, 10.0, 1.0, 60), - (flaky.internal_usage_cache, other_key, 10.0, 1.0, 60), - (flaky.internal_usage_cache, rejecting_key, 0.0, 1.0, 60), + (flaky.internal_usage_cache, failing_key, 10.0, 1.0, 60, False), + (flaky.internal_usage_cache, other_key, 10.0, 1.0, 60, False), + (flaky.internal_usage_cache, rejecting_key, 0.0, 1.0, 60, False), ] ) @@ -3870,18 +3913,18 @@ async def test_exception_mid_batch_refunds_every_earlier_admission_before_propag raising_key = "{tag_rl:test:exception-refund:b}:requests" class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook): - async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int): + async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool): if key == raising_key: raise RuntimeError("simulated transient redis failure") - return await super()._check_and_increment_one(cache, key, limit, increment, ttl) + return await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl) flaky = _FlakyLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now) with pytest.raises(RuntimeError): await flaky._atomic_check_and_increment( [ - (flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60), - (flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60), + (flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60, False), + (flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60, False), ] ) @@ -3908,22 +3951,22 @@ async def test_a_raising_keys_own_ambiguous_outcome_is_never_refunded(time_contr raising_key = "{tag_rl:test:ambiguous-no-refund:b}:requests" class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook): - async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int): + async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool): if key == raising_key: # Simulate Redis committing the increment before the # response is lost: the write actually happens... - await super()._check_and_increment_one(cache, key, limit, increment, ttl) + await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl) # ...but the caller never finds out. raise RuntimeError("simulated lost response after a committed redis write") - return await super()._check_and_increment_one(cache, key, limit, increment, ttl) + return await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl) flaky = _FlakyLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now) with pytest.raises(RuntimeError): await flaky._atomic_check_and_increment( [ - (flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60), - (flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60), + (flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60, False), + (flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60, False), ] )