mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
e7d5b51a39
commit
be8ecc7c26
4 changed files with 234 additions and 14 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue