mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): refresh a concurrency key's Redis TTL on every admission, not just its first
TAG_RL_CHECK_AND_INCR_SCRIPT only called 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 purely as a crash-safety net for a reservation whose explicit release never runs -- so a still-active bucket under sustained traffic would expire mid-flight, silently admitting past the cap and letting a later release decrement an unrelated, newer cohort's counter. Adds a refresh_ttl script argument, true only for the concurrency caller, and verified against a real Redis instance since the in-memory fallback (which already refreshes unconditionally) can't reproduce this. bugbot caught this on review.
This commit is contained in:
parent
54a7f6f7b3
commit
513c273750
3 changed files with 88 additions and 23 deletions
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
]
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue