diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 50ff8d8475f..911e21ad3d1 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -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. diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 79d3cbf29a5..2597ece0552 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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: diff --git a/litellm/proxy/hooks/tag_rate_limiter.py b/litellm/proxy/hooks/tag_rate_limiter.py index 02cbfc1db02..b38f981e819 100644 --- a/litellm/proxy/hooks/tag_rate_limiter.py +++ b/litellm/proxy/hooks/tag_rate_limiter.py @@ -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)) 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 dec549bcb6e..eb8fd892519 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py @@ -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. diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 7e82dffd0ec..a408978f84b 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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