diff --git a/litellm/proxy/hooks/tag_rate_limiter.py b/litellm/proxy/hooks/tag_rate_limiter.py index ba27421673e..927c7bd0981 100644 --- a/litellm/proxy/hooks/tag_rate_limiter.py +++ b/litellm/proxy/hooks/tag_rate_limiter.py @@ -833,10 +833,16 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer A later key's own admission raising (a transient Redis error, or this coroutine being cancelled mid-call, e.g. the caller disconnecting) is treated the same as a normal rejection for refund - purposes: every earlier admission in this batch is refunded via the - `finally` block below before the exception propagates, so a - mid-batch infra failure can't leave a permanently-charged counter or - a leaked concurrency reservation behind for the rest of that key's TTL. + purposes, with one difference: a clean rejection is guaranteed by + TAG_RL_CHECK_AND_INCR_SCRIPT to never have incremented that key (it + returns before calling INCRBY), so only the earlier admissions need + refunding. A raise gives no such guarantee -- Redis can commit the + INCRBY and still have the call raise if the response back to us is + lost (a timeout, a dropped connection) -- so that key's own possibly + -committed increment is refunded too. Refunding a key that in fact + never committed is harmless (floors at 0); skipping one that did + commit would leak a permanently-charged counter or concurrency + reservation for the rest of that key's TTL. Returns (failing_index, values). On success, failing_index is None and values holds each key's new post-increment value, same order as @@ -854,16 +860,19 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer 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): admitted = False + completed = False try: admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl) + completed = True finally: - # Runs on a normal rejection (admitted stays False) and on - # any exception/cancellation from the awaited call above - # (admitted never gets assigned, so it's still the False set - # just before the try) -- either way, everything admitted so - # far in this batch must be refunded before this key's own - # outcome is used. - if not admitted: + # completed=False means the awaited call itself raised or + # was cancelled -- refund through this index inclusive, per + # the docstring above. completed=True and admitted=False is + # a clean rejection -- refund only the earlier ones, since + # this key's own increment never happened. + if not completed: + await self._refund_admitted(checks, up_to_index=index + 1) + elif not admitted: await self._refund_admitted(checks, up_to_index=index) if admitted: admitted_values.append(value) # mutable-ok: see accumulator comment above diff --git a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py index 0194a21cf3b..bef93da4752 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py @@ -2083,6 +2083,38 @@ async def test_exception_mid_batch_refunds_every_earlier_admission_before_propag assert (float(admitted_value) if admitted_value is not None else 0.0) == 0.0 +@pytest.mark.asyncio +async def test_a_committed_increment_whose_response_is_lost_is_also_refunded(time_controller): + """ + Regression test: a key can commit its own increment (e.g. Redis runs + the INCRBY) and still have the call raise if the response back to us is + lost (a timeout, a dropped connection) -- the caller can't tell a lost + response apart from a call that never reached Redis at all. The earlier + fix only refunded indices *before* the one that raised, leaving this + key's own possibly-committed increment permanently charged. It must be + refunded too, not just the earlier ones in the same batch. + """ + raising_key = "{tag_rl:test:lost-response-refund:a}:requests" + + class _FlakyLimiter(_PROXY_TagRateLimiter): + async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int): + 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) + # ...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) + + 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, raising_key, 10.0, 1.0, 60)]) + + raising_key_value = await flaky.internal_usage_cache.async_get_cache(key=raising_key, litellm_parent_otel_span=None) + assert (float(raising_key_value) if raising_key_value is not None else 0.0) == 0.0 + + # --------------------------------------------------------------------------- # scope_by_key_hash -- opt-in per-calling-key bucket separation # ---------------------------------------------------------------------------