fix(rate-limiting): refresh a concurrency key's ttl on every admission, not just its first (ported from #38292)

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 mirrors the same bypass onto
InMemoryCache.set_cache for the no-Redis fallback path (allow_ttl_override
otherwise leaves a still-live ttl untouched).

Also stamps the identical team_scope onto a team-owned deployment's
by_model_name entry whenever any deployment in that group has a team
alias: Router.should_include_deployment lets a same-team caller reach the
deployment by its own internal model_name, not only its
team_public_model_name alias, and both paths must resolve to the same
bucket or a caller could split its usage across two independent counters
by alternating which name it calls with.

Verified against a real local Redis instance (redis-server v8.8.0 on a
scratch port) since the in-memory fallback can't reproduce the Redis TTL
behavior on its own.
This commit is contained in:
Deepanshu 2026-08-27 07:35:43 -04:00
parent 3aaa60d064
commit 8674da5c01
3 changed files with 143 additions and 30 deletions

View file

@ -163,7 +163,12 @@ class InMemoryCache(BaseCache):
return
self.cache_dict[key] = value
if self.allow_ttl_override(key): # if ttl is not set, set it to default ttl
# refresh_ttl bypasses allow_ttl_override's "leave a still-live ttl
# alone" guard -- a caller only sets it for a counter whose ttl must
# keep extending on every write (e.g. a concurrency reservation's
# crash-safety-net ttl), never for one that must stay fixed to its
# original epoch window (e.g. a fixed-period rate-limit bucket).
if kwargs.get("refresh_ttl") or self.allow_ttl_override(key): # if ttl is not set, set it to default ttl
if "ttl" in kwargs and kwargs["ttl"] is not None:
self.ttl_dict[key] = time.time() + float(kwargs["ttl"])
heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key))

View file

@ -123,19 +123,35 @@ _BACKGROUND_TASKS: Final[set["asyncio.Task[None]"]] = set() # mutable-ok: see c
# `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 }
"""
@ -611,7 +627,20 @@ def _build_limits_index(model_list: Sequence[Mapping[str, object]]) -> _LimitsIn
sorted_by_model_name: Final = sorted(model_list, key=lambda deployment: deployment["model_name"])
by_model_name: Final[Mapping[str, tuple[_ConfiguredLimit, ...]]] = MappingProxyType(
{
model_name: configured
model_name: (
# `Router.should_include_deployment` lets a same-team caller
# reach a team-owned deployment by its own internal
# model_name, not only its team_public_model_name alias
# (litellm auto-generates a name unique per (team_id, uuid),
# so every deployment in this group shares one team_id when
# any does) -- stamping the identical team_scope here as the
# alias entry below gets keeps both paths resolving to the
# same bucket, so a caller can't split its usage across two
# independent counters just by alternating which name it calls.
tuple(replace(limit, team_scope=team_scope) for limit in configured)
if (team_scope := next((key[0] for dep in group if (key := _team_alias_key(dep))), None)) is not None
else configured
)
for model_name, deployment_group in groupby(
sorted_by_model_name, key=lambda deployment: deployment["model_name"]
)
@ -1184,12 +1213,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:
@ -1198,7 +1231,9 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
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:
@ -1212,7 +1247,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
@ -1269,10 +1304,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
@ -1290,10 +1325,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
@ -1400,6 +1435,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)
)

View file

@ -3634,6 +3634,75 @@ 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):
"""
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.init_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.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):
"""
The Redis path's refresh_ttl fix above was never mirrored onto the
in-memory fallback, which called async_set_cache unconditionally --
InMemoryCache.allow_ttl_override leaves a still-live ttl untouched, so
a concurrency counter's expiry stayed fixed from its first admission
even with refresh_ttl=True, the same silent-past-the-cap failure mode
the Redis fix closed.
"""
limiter = _make_limiter(time_controller)
cache = limiter.internal_usage_cache
in_memory_cache = cache.dual_cache.in_memory_cache
key = f"tag_rl:test:in-memory-ttl-refresh:{uuid.uuid4().hex}"
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 = in_memory_cache.ttl_dict[key]
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 = in_memory_cache.ttl_dict[key]
assert ttl_after_second_admission > ttl_after_first_admission
# ---------------------------------------------------------------------------
# team_public_model_name alias -- index lookup must not miss
# ---------------------------------------------------------------------------
@ -3647,6 +3716,12 @@ def test_build_limits_index_is_also_keyed_by_team_public_model_name():
model_group_alias). The index must resolve either name to the same
configured limits, or a team-aliased chain's limits are silently never
checked.
Security regression: Router.should_include_deployment also lets a
same-team (or team-unconstrained) caller reach this deployment by its
own internal model_name, not only the alias. Both paths must resolve to
the identical team_scope, or a caller could split its usage across two
independent buckets just by alternating which name it calls with.
"""
deployment = _deployment(
"real-model-name",
@ -3660,10 +3735,7 @@ def test_build_limits_index_is_also_keyed_by_team_public_model_name():
by_alias = index.resolve("team-alias-name", team_id="team-1")
assert by_name != ()
assert [c.entry for c in by_name] == [c.entry for c in by_alias]
# The alias resolution must carry the team_id into the bucket scope --
# see test_build_limits_index_keeps_different_teams_same_alias_separate
# for why (two teams can publish the identical alias string).
assert by_name[0].team_scope is None
assert by_name[0].team_scope == "team-1"
assert by_alias[0].team_scope == "team-1"
@ -4137,9 +4209,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),
]
)
@ -4164,18 +4236,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),
]
)
@ -4202,22 +4274,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),
]
)