mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
3aaa60d064
commit
8674da5c01
3 changed files with 143 additions and 30 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
]
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue