mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(rate-limiting): concurrency reservations never released on real HTTP requests
Live-proxy verification of the disconnect-leak fix surfaced a far more severe, pre-existing bug: tag_rate_limiter's pending-concurrency-key handoff used a contextvars.ContextVar to pass reservations from admission (async_filter_deployments) to release (async_log_success_event / async_log_failure_event). The real proxy request pipeline forks the streaming response through several distinct asyncio Tasks (create_response's disconnect race, the streaming generator's own task, ...); a ContextVar only propagates into tasks forked after a value is set, so release ran in a task that never saw admission's write. Confirmed via task-id tracing on a live proxy that even a normal, fully-completed streaming request never released its concurrency slot -- not just the disconnect case. Replaces the ContextVar with a field directly on the request's own Logging.model_call_details dict, which is explicitly passed by object reference through both admission's request_kwargs and release's kwargs (confirmed identical object identity on a live request), so it survives task boundaries by construction. Deliberately not keyed by litellm_call_id instead: that field is caller-controlled via the x-litellm-call-id header, and an earlier design already tried and rejected that approach for exactly this reason (letting unrelated concurrent requests merge reservations). async_release_disconnect_state_hook now takes request_data so it can reach the same model_call_details. Rewrote the concurrency-release tests to wire a shared model_call_details across admission/release (mirroring production) instead of relying on ambient task context, and re-verified live against a real proxy: both normal completion and disconnect-before-first-chunk now correctly free the slot.
This commit is contained in:
parent
9b3891ee9a
commit
f70e48872f
5 changed files with 233 additions and 188 deletions
|
|
@ -765,7 +765,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
"""
|
||||
return
|
||||
|
||||
async def async_release_disconnect_state_hook(self) -> None:
|
||||
async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None:
|
||||
"""
|
||||
Release per-request state reserved outside of `async_log_success_event`
|
||||
/ `async_log_failure_event` for a request whose streaming response is
|
||||
|
|
@ -773,7 +773,14 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
are `BaseException`, so they bypass both of those callbacks entirely.
|
||||
|
||||
Called from the proxy's shielded streaming cleanup only when no
|
||||
disconnect-time success event fired for this request. Must be
|
||||
disconnect-time success event fired for this request. `request_data`
|
||||
is the same proxy request-data dict threaded through the rest of that
|
||||
cleanup (carries ``litellm_logging_obj`` and other request-scoped
|
||||
state); implementations needing per-request correlation should key
|
||||
off an object reachable from it (e.g. ``litellm_logging_obj``'s own
|
||||
identity or its ``model_call_details``), never off a caller-supplied
|
||||
value like ``litellm_call_id`` (settable via the ``x-litellm-call-id``
|
||||
header), which would let two unrelated requests collide. Must be
|
||||
idempotent and never raise -- a callback that never reserved such
|
||||
state has nothing to do here.
|
||||
|
||||
|
|
|
|||
|
|
@ -398,7 +398,7 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
|
|||
return True
|
||||
|
||||
|
||||
async def _release_disconnect_state_on_all_callbacks() -> None:
|
||||
async def _release_disconnect_state_on_all_callbacks(request_data: Mapping[str, object]) -> None:
|
||||
"""
|
||||
A client disconnect throws GeneratorExit/CancelledError into the streaming
|
||||
generator, so neither the success nor failure logging callback runs for it
|
||||
|
|
@ -417,7 +417,7 @@ async def _release_disconnect_state_on_all_callbacks() -> None:
|
|||
if not isinstance(callback, CustomLogger):
|
||||
continue
|
||||
try:
|
||||
await callback.async_release_disconnect_state_hook()
|
||||
await callback.async_release_disconnect_state_hook(request_data)
|
||||
except Exception as e: # noqa: BLE001 # one callback's cleanup must never block another's or the response teardown
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to run async_release_disconnect_state_hook for %s: %s", type(callback).__name__, e
|
||||
|
|
@ -3392,7 +3392,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
):
|
||||
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(user_api_key_dict)
|
||||
if not success_event_owns_slot_release:
|
||||
await _release_disconnect_state_on_all_callbacks()
|
||||
await _release_disconnect_state_on_all_callbacks(request_data)
|
||||
|
||||
if hasattr(response, "aclose"):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
"""Tag-scoped token, request, dollar, and concurrency rate limits."""
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import hashlib
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
|
|
@ -468,34 +467,34 @@ _CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600
|
|||
# TagRateLimitEntry.max_in_memory_cache_size) each reservation was
|
||||
# incremented under: releasing a reservation must decrement the exact same
|
||||
# cache partition it was incremented on, or the release silently no-ops on
|
||||
# the wrong (default) partition and the reservation leaks forever. Held via
|
||||
# a ContextVar bound to a mutable holder object (not an immutable tuple
|
||||
# rebound with `.set()`) because `asyncio.create_task` only copies which
|
||||
# *object* a ContextVar is bound to, not a snapshot of that object's
|
||||
# contents: a `.set()` performed inside a task forked off this context
|
||||
# mutates only that task's own binding, invisible to the parent task that
|
||||
# continues on to a fallback hop. Mutating a shared holder in place is
|
||||
# visible from every task forked after the holder was first created,
|
||||
# regardless of which task performs the mutation.
|
||||
class _PendingConcurrencyKeys:
|
||||
__slots__ = ("keys",)
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.keys: list[tuple[str, _PartitionKey]] = [] # mutable-ok: shared across forks by design; see docstring
|
||||
|
||||
|
||||
_pending_concurrency_keys: Final[contextvars.ContextVar[_PendingConcurrencyKeys | None]] = contextvars.ContextVar(
|
||||
"tag_rate_limiter_pending_concurrency_keys", default=None
|
||||
)
|
||||
|
||||
|
||||
def _pending_concurrency_holder() -> _PendingConcurrencyKeys:
|
||||
existing: Final = _pending_concurrency_keys.get()
|
||||
if existing is not None:
|
||||
return existing
|
||||
holder: Final = _PendingConcurrencyKeys()
|
||||
_pending_concurrency_keys.set(holder)
|
||||
return holder
|
||||
# the wrong (default) partition and the reservation leaks forever.
|
||||
#
|
||||
# Stashed directly on `Logging.model_call_details` under this field, not a
|
||||
# `contextvars.ContextVar`: the real proxy request pipeline forks the
|
||||
# streaming response through several distinct asyncio Tasks (the disconnect
|
||||
# race in `create_response`, the streaming generator's own task, ...), and a
|
||||
# ContextVar only propagates forward into tasks forked *after* a value was
|
||||
# `.set()` -- a task that isn't a descendant of admission's task never sees
|
||||
# it, so release silently finds nothing and every reservation leaks until
|
||||
# `_CONCURRENCY_MIN_SAFETY_TTL_SECONDS`, disconnect or not (confirmed live:
|
||||
# even a fully-completed, non-disconnected streaming request never released
|
||||
# its slot). `model_call_details` is a single dict, explicitly passed by
|
||||
# object reference through both admission's `request_kwargs` (as
|
||||
# `request_kwargs["litellm_logging_obj"].model_call_details`) and release's
|
||||
# `kwargs` (`async_log_success_event`/`async_log_failure_event`'s `kwargs`
|
||||
# argument *is* `model_call_details` -- see their own callers), so it
|
||||
# survives task boundaries by construction, not by ambient context.
|
||||
#
|
||||
# Deliberately not keyed by `litellm_call_id` instead: that field is
|
||||
# caller-controlled via the `x-litellm-call-id` request header, so two
|
||||
# unrelated concurrent requests sharing a caller-chosen id would merge their
|
||||
# reservations under a shared identifier -- letting one request's release
|
||||
# free a different request's still-live slot. `model_call_details` is a
|
||||
# plain Python object with no caller-visible identifier, created fresh
|
||||
# server-side per logical request (and shared across that request's own
|
||||
# fallback hops, matching the original chain-wide release semantics), so it
|
||||
# can't be forged or guessed.
|
||||
_PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_tag_rate_limiter_pending_concurrency_keys"
|
||||
|
||||
|
||||
class _TagRateLimitIndex:
|
||||
|
|
@ -633,7 +632,7 @@ def _increment_operation_for_limit(
|
|||
now: float,
|
||||
) -> RedisPipelineIncrementOperation | None:
|
||||
if configured_limit.unit == "concurrency":
|
||||
return None # released above, from _pending_concurrency_keys
|
||||
return None # released above, via _pop_pending_concurrency_keys
|
||||
if configured_limit.deployment_scope is not None and deployment_id not in configured_limit.deployment_scope:
|
||||
return None
|
||||
tag_value: Final = _extract_identity(tags, configured_limit.entry.tag_id)
|
||||
|
|
@ -704,6 +703,26 @@ def _partition_key(entry: TagRateLimitEntry) -> _PartitionKey:
|
|||
)
|
||||
|
||||
|
||||
def _queue_pending_concurrency_reservations(
|
||||
request_kwargs: Mapping[str, object], reservations: Sequence[tuple[str, _PartitionKey]]
|
||||
) -> None:
|
||||
"""Stash reservations on the request's own `model_call_details` -- see
|
||||
`_PENDING_CONCURRENCY_KEYS_FIELD`'s docstring for why this, not a
|
||||
ContextVar or `litellm_call_id`. Silently a no-op without a real logging
|
||||
object (defensive only; every real request has one): the reservation
|
||||
still self-heals via `_CONCURRENCY_MIN_SAFETY_TTL_SECONDS`, just later.
|
||||
"""
|
||||
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
|
||||
pending = model_call_details.get(_PENDING_CONCURRENCY_KEYS_FIELD)
|
||||
if pending is None:
|
||||
pending = [] # mutable-ok: shared, request-scoped accumulator; see field's own docstring
|
||||
model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD] = pending
|
||||
pending.extend(reservations) # mutable-ok: see comment above
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CachePartition:
|
||||
internal_usage_cache: InternalUsageCache
|
||||
|
|
@ -984,7 +1003,7 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
|
|||
if configured_limit.unit == "concurrency"
|
||||
)
|
||||
if concurrency_reservations:
|
||||
_pending_concurrency_holder().keys.extend(concurrency_reservations)
|
||||
_queue_pending_concurrency_reservations(resolved_request_kwargs, concurrency_reservations)
|
||||
|
||||
return healthy_deployments
|
||||
|
||||
|
|
@ -1099,7 +1118,7 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
|
|||
|
||||
Each reservation is released against the exact cache partition
|
||||
(`_partition_for(partition_key)`) its increment used -- see
|
||||
`_PendingConcurrencyKeys`'s docstring for why this must match.
|
||||
`_PENDING_CONCURRENCY_KEYS_FIELD`'s docstring for why this must match.
|
||||
"""
|
||||
for key, partition_key in reservations:
|
||||
try:
|
||||
|
|
@ -1109,24 +1128,24 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
|
|||
verbose_proxy_logger.warning("tag_rate_limiter: failed to release concurrency slot %s: %s", key, e)
|
||||
|
||||
@staticmethod
|
||||
def _pop_pending_concurrency_keys() -> tuple[tuple[str, _PartitionKey], ...]:
|
||||
def _pop_pending_concurrency_keys(kwargs: Mapping[str, object]) -> tuple[tuple[str, _PartitionKey], ...]:
|
||||
# Snapshot then remove only those exact keys, never a blanket clear:
|
||||
# a sibling hop can still be live and appending to the same shared
|
||||
# holder concurrently (see the holder's own comment above), so
|
||||
# wiping the whole list here would silently strand that hop's
|
||||
# reservation instead of releasing it later.
|
||||
holder: Final = _pending_concurrency_keys.get()
|
||||
if holder is None or not holder.keys:
|
||||
# a sibling hop sharing this same request's model_call_details can
|
||||
# still be live and appending concurrently (see the field's own
|
||||
# docstring), so wiping the whole list here would silently strand
|
||||
# that hop's reservation instead of releasing it later.
|
||||
pending: Final = kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD)
|
||||
if not isinstance(pending, list) or not pending:
|
||||
return ()
|
||||
keys: Final = tuple(holder.keys)
|
||||
keys: Final = tuple(pending)
|
||||
for key in keys:
|
||||
try:
|
||||
holder.keys.remove(key)
|
||||
pending.remove(key)
|
||||
except ValueError:
|
||||
pass
|
||||
return keys
|
||||
|
||||
async def async_release_disconnect_state_hook(self) -> None:
|
||||
async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None:
|
||||
"""
|
||||
A client disconnecting before the first streamed chunk raises
|
||||
CancelledError/GeneratorExit, which bypasses both async_log_success_event
|
||||
|
|
@ -1136,7 +1155,11 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
|
|||
expires, letting a caller who repeatedly opens and immediately drops
|
||||
streaming requests exhaust their own tag's concurrency limit for free.
|
||||
"""
|
||||
release_keys: Final = self._pop_pending_concurrency_keys()
|
||||
logging_obj: Final = request_data.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)
|
||||
|
||||
|
|
@ -1148,12 +1171,12 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
|
|||
if detail.get("error") == "tag_rate_limit_exceeded":
|
||||
return
|
||||
|
||||
release_keys: Final = self._pop_pending_concurrency_keys()
|
||||
release_keys: Final = self._pop_pending_concurrency_keys(kwargs)
|
||||
if release_keys:
|
||||
await self._release_keys(release_keys)
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
|
||||
release_keys: Final = self._pop_pending_concurrency_keys()
|
||||
release_keys: Final = self._pop_pending_concurrency_keys(kwargs)
|
||||
if release_keys:
|
||||
asyncio.create_task(self._release_keys(release_keys))
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Unit tests for tag-scoped token/request/dollar rate limiting.
|
|||
import asyncio
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -25,8 +26,9 @@ from litellm.proxy.hooks.tag_rate_limiter import (
|
|||
_fixed_length_identity,
|
||||
_inflight_key,
|
||||
_partition_key,
|
||||
_pending_concurrency_holder,
|
||||
_PENDING_CONCURRENCY_KEYS_FIELD,
|
||||
_PROXY_TagRateLimiter,
|
||||
_queue_pending_concurrency_reservations,
|
||||
)
|
||||
from litellm.types.router import TagRateLimitEntry
|
||||
|
||||
|
|
@ -54,6 +56,28 @@ def _make_limiter(time_controller: TimeController) -> _PROXY_TagRateLimiter:
|
|||
)
|
||||
|
||||
|
||||
def _call_context(tags: list[str]) -> tuple[dict, dict]:
|
||||
"""
|
||||
A (request_kwargs, kwargs) pair sharing one `model_call_details` dict,
|
||||
mirroring production: admission reads `request_kwargs["litellm_logging_obj"]
|
||||
.model_call_details`, and the `kwargs` passed to async_log_success_event /
|
||||
async_log_failure_event / async_release_disconnect_state_hook's
|
||||
request_data *is* that same model_call_details dict (or carries the same
|
||||
logging_obj) -- see _PENDING_CONCURRENCY_KEYS_FIELD's docstring. A plain
|
||||
SimpleNamespace stands in for the real Logging object; only its
|
||||
model_call_details attribute is used.
|
||||
"""
|
||||
model_call_details: dict = {}
|
||||
logging_obj = SimpleNamespace(model_call_details=model_call_details)
|
||||
request_kwargs = {"metadata": {"tags": tags}, "litellm_logging_obj": logging_obj}
|
||||
# kwargs must be the *same* dict object model_call_details is, so that
|
||||
# admission's writes onto model_call_details are visible when this kwargs
|
||||
# is later passed to a release hook -- see the docstring above.
|
||||
model_call_details["litellm_logging_obj"] = logging_obj
|
||||
model_call_details["metadata"] = {"tags": tags}
|
||||
return request_kwargs, model_call_details
|
||||
|
||||
|
||||
def _deployment(model_name: str, deployment_id: str, tag_rate_limits: dict) -> dict:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
|
|
@ -984,9 +1008,9 @@ async def test_concurrency_slot_released_on_success_frees_capacity(time_controll
|
|||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
kwargs = {"metadata": {"tags": ["end_user_id:u1"]}}
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
|
||||
)
|
||||
|
||||
# At capacity: a second concurrent request is rejected.
|
||||
|
|
@ -1031,9 +1055,9 @@ async def test_concurrency_slot_released_on_disconnect_frees_capacity(time_contr
|
|||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
kwargs = {"metadata": {"tags": ["end_user_id:u1"]}}
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
|
||||
)
|
||||
|
||||
# At capacity: a second concurrent request is rejected.
|
||||
|
|
@ -1047,7 +1071,7 @@ async def test_concurrency_slot_released_on_disconnect_frees_capacity(time_contr
|
|||
|
||||
# The first request's client disconnects -- neither logging callback fires --
|
||||
# but the disconnect hook still releases its slot, freeing capacity again.
|
||||
await limiter.async_release_disconnect_state_hook()
|
||||
await limiter.async_release_disconnect_state_hook(request_kwargs)
|
||||
|
||||
result = await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
|
|
@ -1065,15 +1089,17 @@ async def test_concurrency_slot_released_on_failure_frees_capacity(time_controll
|
|||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
kwargs["standard_logging_object"] = {"model_group": "grp"}
|
||||
await limiter.async_log_failure_event(
|
||||
kwargs={"standard_logging_object": {"model_group": "grp"}, "metadata": {"tags": ["end_user_id:u1"]}},
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
|
|
@ -1106,17 +1132,16 @@ async def test_concurrency_slot_released_on_fallback_recovered_hop_failure(time_
|
|||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
kwargs["standard_logging_object"] = {"model_group": "grp", "model_id": "dep-1"}
|
||||
await limiter.async_log_failure_event(
|
||||
kwargs={
|
||||
"standard_logging_object": {"model_group": "grp", "model_id": "dep-1"},
|
||||
"metadata": {"tags": ["end_user_id:u1"]},
|
||||
},
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
|
|
@ -1132,82 +1157,69 @@ async def test_concurrency_slot_released_on_fallback_recovered_hop_failure(time_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pending_concurrency_context_does_not_leak_across_concurrent_tasks(time_controller):
|
||||
async def test_pending_concurrency_reservations_do_not_leak_across_unrelated_requests(time_controller):
|
||||
"""
|
||||
Security regression test, current design: `_pending_concurrency_keys` is
|
||||
a `contextvars.ContextVar`, isolated per asyncio task/context rather
|
||||
than a plain shared dict or list -- which matters because two genuinely
|
||||
concurrent, unrelated requests each get their own task in production (a
|
||||
hard ASGI guarantee, not something litellm or this hook controls), so
|
||||
they can never share a context regardless of what identifiers (tags,
|
||||
keys, litellm_call_id) they happen to reuse. Prove this directly: if
|
||||
this were a shared collection instead of a real `ContextVar`, one task's
|
||||
own release would incorrectly drain the other task's still-pending
|
||||
reservation too, since nothing would distinguish which task accumulated
|
||||
which key. An earlier design correlated reservations using
|
||||
litellm_call_id specifically -- caller-controlled via the
|
||||
x-litellm-call-id header -- as a registry key; that's what let two
|
||||
unrelated concurrent requests merge reservations in the first place, and
|
||||
is why this test isolates via real tasks rather than a shared id at all.
|
||||
Security regression test, current design: pending concurrency keys are
|
||||
stashed on the admitting request's own `model_call_details` dict (see
|
||||
`_PENDING_CONCURRENCY_KEYS_FIELD`'s docstring), never in a registry keyed
|
||||
by anything caller-visible or by ambient asyncio context. Two unrelated
|
||||
concurrent requests each get their own `model_call_details` in
|
||||
production, so one request's release can never see or drain a different
|
||||
request's still-pending reservation, regardless of which asyncio task
|
||||
each happens to run in and even when both share the identical tag value
|
||||
(an earlier design keyed reservations by `litellm_call_id` -- settable by
|
||||
the caller via the `x-litellm-call-id` header -- which let two unrelated
|
||||
requests merge reservations simply by choosing the same id).
|
||||
"""
|
||||
limiter = _make_limiter(time_controller)
|
||||
router = _concurrency_router(limit=2)
|
||||
router = _concurrency_router(limit=1)
|
||||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
async def _admit(tag_value):
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": [f"end_user_id:{tag_value}"]}},
|
||||
)
|
||||
request_a, kwargs_a = _call_context(["end_user_id:shared"])
|
||||
request_b, kwargs_b = _call_context(["end_user_id:shared"])
|
||||
|
||||
async def _release(tag_value):
|
||||
await limiter.async_log_success_event(
|
||||
kwargs={
|
||||
"standard_logging_object": {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
},
|
||||
"metadata": {"tags": [f"end_user_id:{tag_value}"]},
|
||||
},
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
)
|
||||
|
||||
# Two separate, genuinely concurrent tasks admit -- reaching capacity.
|
||||
task_a = asyncio.create_task(_admit("a"))
|
||||
task_b = asyncio.create_task(_admit("b"))
|
||||
await task_a
|
||||
await task_b
|
||||
|
||||
# Task A releases its own reservation, in its own task -- this must not
|
||||
# also release task B's still-pending one.
|
||||
await asyncio.create_task(_release("a"))
|
||||
await asyncio.sleep(0)
|
||||
|
||||
# Exactly one slot was freed: a fresh request is admitted (back to 2 in flight)...
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:a"]}},
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_a
|
||||
)
|
||||
# ...but a second one does not, since B's reservation is genuinely still
|
||||
# held. If task isolation were broken, task A's release would have
|
||||
# drained B's reservation too, and this would wrongly admit.
|
||||
# B shares A's tag value but is a genuinely separate request/object: at
|
||||
# capacity (limit=1), B is rejected and never reserves anything.
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_b
|
||||
)
|
||||
|
||||
# B's own failure event releases via its own (empty) model_call_details --
|
||||
# this must not accidentally drain A's still-live reservation.
|
||||
kwargs_b["standard_logging_object"] = {"model_group": "grp", "model_id": "dep-1"}
|
||||
await limiter.async_log_failure_event(kwargs=kwargs_b, response_obj=None, start_time=0, end_time=0)
|
||||
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:a"]}},
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:shared"]}},
|
||||
)
|
||||
|
||||
# A's own success event correctly releases its own reservation.
|
||||
kwargs_a["standard_logging_object"] = {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
}
|
||||
await limiter.async_log_success_event(kwargs=kwargs_a, response_obj=None, start_time=0, end_time=0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
result = await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:shared"]}},
|
||||
)
|
||||
assert result == healthy
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(time_controller):
|
||||
|
|
@ -1217,18 +1229,18 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
|
|||
`Logging.has_run_logging`'s `has_logged_async_failure` guard); a later
|
||||
failed hop (a retry or a further fallback) never gets its own failure
|
||||
event at all. Reservations still accumulate at admission for every hop
|
||||
regardless (onto `_pending_concurrency_keys`), so whichever event fires
|
||||
next must release everything accumulated since the last release, not
|
||||
just its own hop's key.
|
||||
regardless (onto the request's own `model_call_details`, shared across
|
||||
every hop of one logical request -- see `_PENDING_CONCURRENCY_KEYS_FIELD`'s
|
||||
docstring), so whichever event fires next must release everything
|
||||
accumulated since the last release, not just its own hop's key.
|
||||
|
||||
Hop 3's eventual success is fired as a child task of the same admission
|
||||
chain -- exactly like litellm's real dispatch, where `wrapper_async`
|
||||
create_task's the success path and `LoggingWorker.enqueue` explicitly
|
||||
propagates the calling context -- to prove the fix survives the actual
|
||||
task boundary a real success event crosses in production, not just a
|
||||
same-coroutine call that would pass regardless of whether
|
||||
`_pending_concurrency_keys` were a real `ContextVar` or an ordinary
|
||||
variable.
|
||||
same-coroutine call that would pass regardless of whether the pending
|
||||
keys lived on a real shared object or an ordinary per-task variable.
|
||||
"""
|
||||
limiter = _make_limiter(time_controller)
|
||||
router = _concurrency_router(limit=2)
|
||||
|
|
@ -1236,6 +1248,11 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
|
|||
healthy = router.model_list
|
||||
|
||||
async def _one_logical_request():
|
||||
# All three hops of this one logical request share the same
|
||||
# model_call_details, exactly as real fallback hops share one
|
||||
# Logging object -- only litellm_call_id differs per hop.
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
|
||||
# Hop 1 admits and fails; its failure event is the one that fires
|
||||
# (dedup allows exactly the first failure through), releasing its
|
||||
# own key immediately.
|
||||
|
|
@ -1243,10 +1260,11 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
|
|||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
kwargs["standard_logging_object"] = {"model_group": "grp"}
|
||||
await limiter.async_log_failure_event(
|
||||
kwargs={"standard_logging_object": {"model_group": "grp"}, "metadata": {"tags": ["end_user_id:u1"]}},
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
|
|
@ -1258,7 +1276,7 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
|
|||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
# Hop 3 admits and succeeds. Its success event, dispatched as a
|
||||
|
|
@ -1268,19 +1286,18 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti
|
|||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
async def _hop_3_success_event():
|
||||
kwargs["standard_logging_object"] = {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
}
|
||||
await limiter.async_log_success_event(
|
||||
kwargs={
|
||||
"standard_logging_object": {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
}
|
||||
},
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
|
|
@ -1930,53 +1947,53 @@ def test_concurrency_ttl_floor_does_not_shorten_a_longer_period_seconds():
|
|||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# pending-concurrency-key holder must survive a detached asyncio.create_task
|
||||
# fork (e.g. litellm's own failure-logging dispatch) without a rebind in that
|
||||
# forked task hiding the release from the parent, and a release must never
|
||||
# sweep up a key a still-live sibling hop appended in the meantime
|
||||
# pending-concurrency-key field on model_call_details must survive a detached
|
||||
# asyncio.create_task fork (e.g. litellm's own failure-logging dispatch),
|
||||
# and a release must never sweep up a key a still-live sibling hop appended
|
||||
# in the meantime. This dict-on-a-shared-object design is what replaced a
|
||||
# contextvars.ContextVar-based holder that silently failed to release
|
||||
# anything once release ran in a task that wasn't a descendant of admission's
|
||||
# own task -- exactly what happens in the real proxy request pipeline (see
|
||||
# _PENDING_CONCURRENCY_KEYS_FIELD's docstring).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_in_a_forked_task_is_visible_to_the_parent_context():
|
||||
_pending_concurrency_holder().keys.clear()
|
||||
_pending_concurrency_holder().keys.append("key1")
|
||||
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
|
||||
|
||||
async def detached_release():
|
||||
return _PROXY_TagRateLimiter._pop_pending_concurrency_keys()
|
||||
return _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details)
|
||||
|
||||
released = await asyncio.create_task(detached_release())
|
||||
assert released == ("key1",)
|
||||
|
||||
# The parent's own binding must see the same, now-empty holder --
|
||||
# not a stale copy still holding "key1".
|
||||
assert _pending_concurrency_holder().keys == []
|
||||
# The parent's own view of the same dict must see the release too.
|
||||
assert model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot():
|
||||
_pending_concurrency_holder().keys.clear()
|
||||
_pending_concurrency_holder().keys.append("key1")
|
||||
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
|
||||
|
||||
async def detached_release_then_sibling_admits():
|
||||
released = _PROXY_TagRateLimiter._pop_pending_concurrency_keys()
|
||||
# A sibling hop's admission, appending to the same shared holder,
|
||||
released = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details)
|
||||
# A sibling hop's admission, appending to the same shared dict,
|
||||
# interleaved right after this release's snapshot was taken.
|
||||
_pending_concurrency_holder().keys.append("key2")
|
||||
model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD].append("key2")
|
||||
return released
|
||||
|
||||
released = await asyncio.create_task(detached_release_then_sibling_admits())
|
||||
assert released == ("key1",)
|
||||
# key2 must still be pending for its own hop's eventual release.
|
||||
assert _pending_concurrency_holder().keys == ["key2"]
|
||||
assert model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD] == ["key2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_is_not_repeated_for_the_same_snapshot():
|
||||
_pending_concurrency_holder().keys.clear()
|
||||
_pending_concurrency_holder().keys.append("key1")
|
||||
first = _PROXY_TagRateLimiter._pop_pending_concurrency_keys()
|
||||
second = _PROXY_TagRateLimiter._pop_pending_concurrency_keys()
|
||||
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
|
||||
first = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details)
|
||||
second = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details)
|
||||
assert first == ("key1",)
|
||||
assert second == ()
|
||||
|
||||
|
|
@ -2420,56 +2437,54 @@ async def test_concurrency_scope_by_key_hash_gives_independent_reservations_per_
|
|||
block keyB's admission, and releasing keyA's reservation (via the
|
||||
standard_logging_object.metadata.user_api_key_hash channel) must free
|
||||
keyA's capacity, not keyB's. Each key is modeled as its own logical
|
||||
request: one task does admission and then spawns its own release as a
|
||||
child task, exactly like litellm's real dispatch (`wrapper_async`
|
||||
create_task's the success path, itself a descendant of the same
|
||||
admission-time task/context chain) -- release must never be spawned as
|
||||
an unrelated sibling task from the test's own top level, which would
|
||||
start from a fresh context that never saw the admission's `ContextVar`
|
||||
write at all, an artifact of this test's own construction rather than a
|
||||
real bug.
|
||||
request with its own model_call_details, and keyA's release is spawned
|
||||
as a genuinely separate child task (mirroring litellm's real dispatch)
|
||||
to prove release survives that task boundary via the shared
|
||||
model_call_details object, not via which task happens to run it.
|
||||
"""
|
||||
limiter = _make_limiter(time_controller)
|
||||
router = _concurrency_router_scoped_by_key(limit=1)
|
||||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
async def _admit(key: str):
|
||||
async def _admit(key: str, request_kwargs: dict):
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": ["end_user_id:u1"], "user_api_key": key}},
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
async def _release(key: str):
|
||||
async def _release(key: str, kwargs: dict):
|
||||
kwargs["standard_logging_object"] = {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
"metadata": {"user_api_key_hash": key},
|
||||
}
|
||||
await limiter.async_log_success_event(
|
||||
kwargs={
|
||||
"standard_logging_object": {
|
||||
"model_group": "grp",
|
||||
"model_id": "dep-1",
|
||||
"total_tokens": 0,
|
||||
"response_cost": 0,
|
||||
"metadata": {"user_api_key_hash": key},
|
||||
},
|
||||
"metadata": {"tags": ["end_user_id:u1"]},
|
||||
},
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=0,
|
||||
end_time=0,
|
||||
)
|
||||
|
||||
ready_to_release = asyncio.Event()
|
||||
key_a_request, key_a_kwargs = _call_context(["end_user_id:u1"])
|
||||
key_a_request["metadata"]["user_api_key"] = "keyA"
|
||||
key_b_request, _key_b_kwargs = _call_context(["end_user_id:u1"])
|
||||
key_b_request["metadata"]["user_api_key"] = "keyB"
|
||||
|
||||
async def _key_a_admits_then_waits_then_releases_from_the_same_context_chain():
|
||||
await _admit("keyA")
|
||||
await _admit("keyA", key_a_request)
|
||||
await ready_to_release.wait()
|
||||
await asyncio.create_task(_release("keyA"))
|
||||
await asyncio.create_task(_release("keyA", key_a_kwargs))
|
||||
|
||||
# keyA occupies its own single slot; keyB, same tag value, different
|
||||
# key, still admits since it has its own bucket.
|
||||
key_a_task = asyncio.create_task(_key_a_admits_then_waits_then_releases_from_the_same_context_chain())
|
||||
key_b_task = asyncio.create_task(_admit("keyB"))
|
||||
key_b_task = asyncio.create_task(_admit("keyB", key_b_request))
|
||||
await key_b_task
|
||||
# Let key_a_task's admission run up to (but not past) `ready_to_release.wait()`.
|
||||
await asyncio.sleep(0)
|
||||
|
|
@ -2913,9 +2928,9 @@ async def test_concurrency_slot_with_a_cache_size_override_is_released_against_t
|
|||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
|
||||
kwargs = {"metadata": {"tags": ["end_user_id:u1"]}}
|
||||
request_kwargs, kwargs = _call_context(["end_user_id:u1"])
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs
|
||||
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
|
||||
)
|
||||
|
||||
# At capacity: a second concurrent reservation for the same tag is rejected.
|
||||
|
|
|
|||
|
|
@ -5601,7 +5601,7 @@ class _RecordingDisconnectHookLogger(CustomLogger):
|
|||
super().__init__()
|
||||
self.disconnect_hook_calls = 0
|
||||
|
||||
async def async_release_disconnect_state_hook(self) -> None:
|
||||
async def async_release_disconnect_state_hook(self, request_data: dict) -> None:
|
||||
self.disconnect_hook_calls += 1
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue