fix(rate-limiting): close three new leak/mismatch findings from this review round

- resolve_any now dedups routing-group candidates in sorted order, not raw
  frozenset iteration order, which depends on the process's hash seed and
  could pick a different resolved_group (and therefore Redis key) across
  worker processes for the identical candidate set.
- Success-side token/dollar accounting now derives scope_by_key_hash's key
  hash the same way admission does (straight from request metadata), instead
  of standard_logging_object's own derived field, which silently drops to
  None whenever the raw key value isn't SHA-256-shaped.
- A concurrency reservation from a failed retry/fallback hop is now released
  at the very start of the next hop's own admission call: litellm only ever
  fires async_log_failure_event once per request (the first failed hop wins,
  every later hop's failure is silently deduped), so a later hop's own
  reservation would otherwise never be released before its TTL.
- The opt-in cancel_on_disconnect path now also releases callback-held
  per-request state (e.g. a concurrency slot) before converting the
  cancellation to a 499, mirroring the existing streaming-disconnect cleanup:
  asyncio.CancelledError bypasses the normal success/failure logging
  callbacks here too, since it's a BaseException, not an Exception.
This commit is contained in:
Deepanshu 2026-08-20 15:18:26 -04:00
parent e7d5b51a39
commit be8ecc7c26
4 changed files with 234 additions and 14 deletions

View file

@ -1481,6 +1481,7 @@ async def _cancel_llm_call_on_client_disconnect(
async def _await_llm_call_cancelling_on_disconnect(
request: Request,
llm_api_call: "asyncio.Future[_LlmCallT]",
request_data: Mapping[str, object],
) -> _LlmCallT:
disconnect_event: Final = asyncio.Event()
monitor: Final = asyncio.create_task(_cancel_llm_call_on_client_disconnect(request, llm_api_call, disconnect_event))
@ -1488,6 +1489,14 @@ async def _await_llm_call_cancelling_on_disconnect(
return await llm_api_call
except asyncio.CancelledError:
if disconnect_event.is_set():
# This cancellation never reaches litellm.utils.wrapper_async's own
# except block (asyncio.CancelledError is a BaseException, not an
# Exception, since Python 3.8), so async_log_failure_event never
# fires for it -- the same gap async_release_disconnect_state_hook
# was added for on the streaming path (see
# _finalize_streaming_generator_cleanup), just reached here via a
# cancelled non-streaming call instead of a mid-stream disconnect.
await _release_disconnect_state_on_all_callbacks(request_data)
raise HTTPException(
status_code=499,
detail=_CLIENT_DISCONNECT_DETAIL,
@ -2312,7 +2321,7 @@ class ProxyBaseLLMRequestProcessing:
try:
if general_settings.get("cancel_on_disconnect", False):
responses = await _await_llm_call_cancelling_on_disconnect(request, llm_responses)
responses = await _await_llm_call_cancelling_on_disconnect(request, llm_responses, self.data)
else:
responses = await llm_responses
finally:

View file

@ -384,12 +384,18 @@ class _LimitsIndex:
are left as separate entries, same as before this dedup: resolving
that ambiguity needs knowing which deployment will be picked, which
isn't known yet at this admission-time hook.
Candidates are deduped in sorted order, not raw `frozenset` iteration
order: `frozenset` order depends on the process's hash seed, so two
workers resolving the identical candidate set could otherwise pick
different members as `resolved_group` and end up checking/accounting
against different Redis keys for what's meant to be one shared bucket.
"""
direct: Final = self.resolve(model, team_id)
if direct:
return direct
deduped: Final[dict[tuple[object, ...], _ConfiguredLimit]] = {} # mutable-ok: see docstring above
for name in frozenset(candidate_model_names):
for name in sorted(frozenset(candidate_model_names)):
for limit in self.by_model_name.get(name, ()):
key = (
limit.unit,
@ -971,6 +977,7 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
return healthy_deployments
resolved_request_kwargs: Final = request_kwargs or _EMPTY_MAPPING
await self._release_stale_hop_reservations(resolved_request_kwargs)
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(resolved_request_kwargs)
team_id: Final = _extract_team_id(resolved_request_kwargs, metadata_variable_name)
candidate_model_names: Final = tuple(
@ -1164,6 +1171,35 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
except Exception as e: # noqa: BLE001 - releasing a slot must never raise into the caller's request path
verbose_proxy_logger.warning("tag_rate_limiter: failed to release concurrency slot %s: %s", key, e)
async def _release_stale_hop_reservations(self, request_kwargs: Mapping[str, object]) -> None:
"""
A concurrency reservation still queued when a *new* hop's admission
runs can only belong to an earlier hop of this same request that
already concluded and failed: Router awaits one hop's entire attempt
(call plus its own failure handling) before starting the next, and a
hop that instead succeeded ends the request there via
async_log_success_event, which already pops everything -- so
admission is never re-entered while an earlier hop's reservation is
still legitimately in flight.
LiteLLM only invokes a request's CustomLogger.async_log_failure_event
once per request, for whichever hop fails first (its internal
has_logged_async_failure dedup silently skips every later hop's own
failure), so every hop after that one would otherwise never release
its predecessor's key until _CONCURRENCY_MIN_SAFETY_TTL_SECONDS.
Releasing here, at the one point guaranteed to re-run before every
subsequent hop, closes that gap for every hop except a final one
whose own failure exhausts the retry chain -- that residual case
still self-heals via the same TTL floor.
"""
logging_obj: Final = request_kwargs.get("litellm_logging_obj")
model_call_details: Final = getattr(logging_obj, "model_call_details", None)
if not isinstance(model_call_details, dict):
return
release_keys: Final = self._pop_pending_concurrency_keys(model_call_details)
if release_keys:
await self._release_keys(release_keys)
@staticmethod
def _pop_pending_concurrency_keys(kwargs: Mapping[str, object]) -> tuple[tuple[str, _PartitionKey], ...]:
# Snapshot then remove only those exact keys, never a blanket clear:
@ -1232,7 +1268,21 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
standard_logging_metadata: Final = standard_logging_object.get("metadata") or _EMPTY_MAPPING
team_id: Final = standard_logging_metadata.get("user_api_key_team_id")
key_hash: Final = standard_logging_metadata.get("user_api_key_hash")
# kwargs here is Logging.model_call_details, not the router's flat
# request kwargs admission sees: metadata/litellm_metadata are never
# top-level here, only nested under kwargs["litellm_params"] (see
# Logging.update_environment_variables).
litellm_params_for_metadata: Final = kwargs.get("litellm_params") or kwargs
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(litellm_params_for_metadata)
# standard_logging_object.metadata.user_api_key_hash is only ever
# populated when the raw value happens to look like a SHA-256 hash
# (see litellm_logging.py's get_standard_logging_metadata), so it
# silently drops to None for any key whose hash doesn't pass that
# shape check even though admission's own _extract_key_hash reads
# the same field unconditionally -- reading straight from kwargs
# here instead keeps this bucket identical to the one admission
# already scoped the check against.
key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name)
# model_group is the caller-visible name, which Router deliberately
# keeps distinct from the serving deployment's own model_name for a
# routing-group call (see resolve_any's docstring). Passing only the
@ -1266,16 +1316,12 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
if not configured:
return
# kwargs here is Logging.model_call_details, not the router's flat
# request kwargs admission sees: metadata/litellm_metadata are never
# top-level here, only nested under kwargs["litellm_params"] (see
# Logging.update_environment_variables). Resolving the field name
# against kwargs itself always picks the "metadata" default, so on
# LITELLM_METADATA_ROUTES (/v1/messages, /responses, ...) this read
# the caller's native, tag-less metadata instead of the real,
# server-computed litellm_metadata.tags admission already used.
litellm_params_for_metadata: Final = kwargs.get("litellm_params") or kwargs
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(litellm_params_for_metadata)
# Resolving the field name against kwargs itself always picks the
# "metadata" default, so on LITELLM_METADATA_ROUTES (/v1/messages,
# /responses, ...) this would read the caller's native, tag-less
# metadata instead of the real, server-computed litellm_metadata.tags
# admission already used -- metadata_variable_name above is already
# resolved against litellm_params_for_metadata to avoid that.
tags: Final = _get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)
if not tags:
return

View file

@ -3,6 +3,9 @@ Unit tests for tag-scoped token/request/dollar rate limiting.
"""
import asyncio
import os
import subprocess
import sys
import uuid
from datetime import datetime, timedelta
from types import SimpleNamespace
@ -559,6 +562,50 @@ def test_resolve_any_keeps_divergent_signatures_across_member_model_names_separa
assert {c.entry.limit for c in resolved} == {1, 2}
def test_resolve_any_picks_the_same_resolved_group_regardless_of_hash_seed():
"""
Two members with an identical signature dedup to whichever one
`frozenset(candidate_model_names)` iterates first. Plain `frozenset`
iteration order for strings is seeded from `PYTHONHASHSEED`, which is
randomized per process by default, so two proxy worker processes (or the
same process across a restart) resolving the identical member set could
pick different members as `resolved_group` -- fragmenting what's meant to
be one shared Redis bucket into two. This can't be observed from within
one interpreter (a single process has one fixed seed for its lifetime),
so this spawns two real subprocesses pinned to seeds empirically known to
order these three names differently under a plain, unsorted frozenset --
see the bug report this regression-tests for the exact reproduction.
"""
script = (
"from litellm.proxy.hooks.tag_rate_limiter import _build_limits_index\n"
"def _deployment(model_name, deployment_id, tag_rate_limits):\n"
" return {'model_name': model_name, 'litellm_params': {'model': 'gpt-4o'},"
" 'model_info': {'id': deployment_id, 'tag_rate_limits': tag_rate_limits}}\n"
"limits = {'concurrency_limits': {'limits': [{'name': 'inflight', 'tag_id': 'end_user_id',"
" 'limit': 1, 'period_seconds': 300}]}}\n"
"index = _build_limits_index(["
"_deployment('backend-a', 'dep-a', limits),"
"_deployment('backend-b', 'dep-b', limits),"
"_deployment('backend-c', 'dep-c', limits)])\n"
"resolved = index.resolve_any('my-group', team_id=None,"
" candidate_model_names=('backend-a', 'backend-b', 'backend-c'))\n"
"print(resolved[0].resolved_group)\n"
)
# seed=1 and seed=3 are empirically confirmed to order these three
# literal strings differently under plain (unsorted) frozenset iteration.
results = {
seed: subprocess.run(
[sys.executable, "-c", script],
env={**os.environ, "PYTHONHASHSEED": seed},
capture_output=True,
text=True,
check=True,
).stdout.strip()
for seed in ("1", "3")
}
assert results["1"] == results["3"] == "backend-a"
@pytest.mark.asyncio
async def test_filter_deployments_per_entry_fail_open_when_tag_absent(time_controller):
"""
@ -939,6 +986,59 @@ async def test_log_success_event_accounts_against_the_same_bucket_admission_chec
)
@pytest.mark.asyncio
async def test_log_success_event_accounts_against_the_key_hash_admission_checked(time_controller):
"""
Admission's `_extract_key_hash` reads `metadata.user_api_key` unconditionally
whenever scope_by_key_hash is set -- that field is already the hashed
token by the time it reaches this hook (see the function's own
docstring), regardless of its shape. `standard_logging_object.metadata`
only ever carries the derived `user_api_key_hash` field, and only when the
raw value happens to look like a SHA-256 hex digest (see
litellm_logging.py's get_standard_logging_metadata) -- a virtual key
represented any other way makes that field silently absent, so reading it
on the success side would account against key_hash=None while admission
scoped the check against the real value, letting usage silently bypass a
per-key limit whenever the key's own representation isn't SHA-256-shaped.
"""
token_limits = {
"token_limits": {
"limits": [
{"name": "daily", "tag_id": "end_user_id", "limit": 500000, "period_seconds": 86400, "scope_by_key_hash": True}
]
}
}
router = litellm.Router(model_list=[_deployment("grp", "dep-1", token_limits)])
limiter = _make_limiter(time_controller)
limiter.update_variables(llm_router=router)
# "keyA" deliberately isn't SHA-256-shaped, so standard_logging_object's
# own redaction/derivation step would never populate user_api_key_hash
# for it -- it's simply absent, matching production for a key hash that
# doesn't pass that shape check.
kwargs = {
"metadata": {"tags": ["end_user_id:u1"], "user_api_key": "keyA"},
"standard_logging_object": {
"model_group": "grp",
"model_id": "dep-1",
"total_tokens": 42,
"response_cost": 0.01,
"metadata": {},
},
}
await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0)
await asyncio.sleep(0)
now = time_controller.now().timestamp()
keyed_bucket = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash="keyA")
unkeyed_bucket = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash=None)
assert (
float(await limiter.internal_usage_cache.async_get_cache(key=keyed_bucket, litellm_parent_otel_span=None))
== 42.0
)
assert await limiter.internal_usage_cache.async_get_cache(key=unkeyed_bucket, litellm_parent_otel_span=None) is None
# ---------------------------------------------------------------------------
# concurrency limits -- reserve at admission, release on success/failure
# ---------------------------------------------------------------------------
@ -1612,6 +1712,45 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
assert result == healthy
@pytest.mark.asyncio
async def test_next_hops_admission_releases_a_prior_hops_leaked_reservation(time_controller):
"""
Regression test for a leak that a success/failure-event-only release
strategy can never close: litellm's has_logged_async_failure dedup lets
exactly one hop's async_log_failure_event fire per logical request (see
test_concurrency_released_for_every_hop_across_a_real_task_boundary), so
a hop that fails *after* that one event has already fired gets no
failure event of its own at all -- not "delayed until the next event",
genuinely never. Only the next hop's own admission call is guaranteed to
run afterward, so release must happen there, not wait for some later
success/failure event that this specific hop will never get.
Concurrency limit of 1 makes this observable directly: if hop 2's
admission doesn't release hop 1's leaked reservation before checking its
own, it raises ProxyRateLimitError against a bucket that's actually free.
"""
limiter = _make_limiter(time_controller)
router = _concurrency_router(limit=1)
limiter.update_variables(llm_router=router)
healthy = router.model_list
request_kwargs, _kwargs = _call_context(["end_user_id:u1"])
# Hop 1 admits (the only slot) and then fails with no failure event ever
# following it -- simulating every hop after litellm's one dedup-allowed
# failure event has already fired for an earlier hop of this request.
result = await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
)
assert result == healthy
# Hop 2's own admission call must release hop 1's stale reservation
# before checking its own -- if it didn't, this raises ProxyRateLimitError.
result = await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
)
assert result == healthy
@pytest.mark.asyncio
async def test_own_rejection_does_not_release_a_live_reservation(time_controller):
"""

View file

@ -4063,7 +4063,33 @@ class TestCancelOnDisconnect:
llm_call.cancel()
with pytest.raises(asyncio.CancelledError):
await _await_llm_call_cancelling_on_disconnect(request, llm_call)
await _await_llm_call_cancelling_on_disconnect(request, llm_call, {})
async def test_disconnect_releases_callback_state_before_499(self, monkeypatch):
"""
asyncio.CancelledError is a BaseException, not an Exception, so it
never reaches litellm.utils.wrapper_async's own except block -- the
cancelled call's async_log_failure_event never fires, and the 499
this raises is later handled by post_call_failure_hook, a different
hook a CustomLogger like tag_rate_limiter doesn't implement. Without
an explicit release here, a callback that reserved per-request state
at admission (a concurrency slot) leaks it until that state's own
safety TTL. This mirrors the streaming disconnect case
(_finalize_streaming_generator_cleanup), just for a non-streaming
call cancelled via the opt-in cancel_on_disconnect flag.
"""
recorder = _RecordingDisconnectHookLogger()
monkeypatch.setattr(litellm, "callbacks", [recorder])
request = self._request([{"type": "http.disconnect"}])
llm_call = asyncio.get_running_loop().create_future()
with pytest.raises(HTTPException) as exc_info:
await _await_llm_call_cancelling_on_disconnect(
request, llm_call, {"litellm_logging_obj": MagicMock()}
)
assert exc_info.value.status_code == 499
assert recorder.disconnect_hook_calls == 1
async def _drive_base_process_llm_request(
self, monkeypatch, general_settings: dict, llm_call, request: Request