diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 7edf4672963..284e780cc6f 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -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)) 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 66ee3fccb44..4899f8fbd2a 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -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) ) 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 0f19ede858c..d52d7ca2c18 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 @@ -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), ] )