fix(rate-limiting): release the final fallback hop's slot via a cache mirror

The earlier async_post_call_failure_hook attempt (293aed9b2c) read
model_call_details off request_data["litellm_logging_obj"], but
proxy/utils.py's post_call_failure_hook deliberately pops that key before
invoking any callback ("Remove before callbacks iterate -- not
serialisable"), confirmed live: the hook always no-op'd.

A ContextVar-based mirror was tried next and also confirmed broken live:
its value never reached the task that calls post_call_failure_hook, since
Router's own per-hop execution does not keep that task a descendant of the
one that ran the final hop's admission. Mutating request_kwargs directly is
unsafe too, since litellm forwards unrecognized kwargs to the actual
provider call as extra_body.

litellm_call_id is the one identifier stable across every one of those
objects, so this mirrors the latest hop's own reservation in the same
external cache (Redis or in-memory) the reservations already live in, keyed
by call_id, and reads it back from there instead. Every normal release path
now also clears this same cache entry, so a hop already released the normal
way is never found "stale" and double-released. Verified live for both a
plain and a `stream: true` request whose entire fallback chain fails before
any token is produced: both now reach the real provider on a follow-up call
instead of being rejected by this hook's own stuck reservation.
This commit is contained in:
Deepanshu 2026-08-25 13:41:00 -04:00
parent 311bf593f4
commit c3593e63f2
2 changed files with 166 additions and 25 deletions

View file

@ -2,6 +2,7 @@
import asyncio
import hashlib
import json
from collections.abc import Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass, replace
from datetime import datetime
@ -645,6 +646,78 @@ _CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600
# can't be forged or guessed.
_PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_model_based_tag_rate_limits_pending_concurrency_keys"
# Mirrors the latest hop's own queued reservation in the same external cache
# the reservations themselves live in, keyed by litellm_call_id, for the one
# release path that cannot reach model_call_details at all:
# proxy/utils.py's post_call_failure_hook deliberately pops litellm_logging_obj
# off request_data before invoking any callback's async_post_call_failure_hook
# ("Remove before callbacks iterate — not serialisable"), so a fallback
# chain's own final, chain-exhausting failure -- which only this hook fires
# for, since litellm's has_logged_async_failure dedup blocks
# async_log_failure_event for every hop after the first -- has no
# model_call_details to pop a reservation from.
#
# Neither a ContextVar nor the flat request_kwargs dict works here (both
# confirmed live, not just reasoned about): a ContextVar's value only
# propagates into descendant tasks, and Router's own per-hop/per-attempt
# execution does not keep the task that later calls post_call_failure_hook
# a descendant of the task that ran the final hop's own admission, so a
# value set there is invisible by the time this fires. request_kwargs is a
# distinct object every hop (confirmed via id()), is a third, unrelated
# object again by the time post_call_failure_hook runs, and mutating it
# directly leaks the mutated key into the actual provider call as an
# `extra_body` param, since litellm forwards unrecognized kwargs verbatim.
# litellm_call_id is the one identifier that is stable across every one of
# those objects, so an external cache keyed by it -- the same Redis/
# in-memory store the reservations themselves already live in -- is the
# only channel that survives all three failure modes at once.
_PENDING_RESERVATIONS_CACHE_KEY_PREFIX: Final = "model_based_tag_rate_limits:pending_reservations:"
def _pending_reservations_cache_key(call_id: str) -> str:
return f"{_PENDING_RESERVATIONS_CACHE_KEY_PREFIX}{call_id}"
def _encode_reservations(reservations: Sequence[tuple[str, "_PartitionKey"]]) -> str:
return json.dumps(
tuple(
(key, partition_key if partition_key is None else tuple(partition_key))
for key, partition_key in reservations
)
)
def _as_decoded_list(raw: object) -> Sequence[object] | None:
# InMemoryCache.get_cache always attempts json.loads on read regardless
# of what was stored (see its own implementation), so a value written as
# our own already-JSON-encoded string comes back pre-decoded into a list
# when served from the in-memory layer; only a real Redis round trip
# hands back the raw string that still needs decoding here.
if isinstance(raw, list):
return raw
if not isinstance(raw, str):
return None
try:
decoded: Final = json.loads(raw)
except (TypeError, ValueError):
return None
return decoded if isinstance(decoded, list) else None
def _decode_reservations(raw: object) -> tuple[tuple[str, "_PartitionKey"], ...]:
decoded: Final = _as_decoded_list(raw)
if decoded is None:
return ()
entries: Final = [] # mutable-ok: accumulator over an externally-decoded, untrusted-shape list; immediately frozen below
for item in decoded:
if not (isinstance(item, list) and len(item) == 2 and isinstance(item[0], str)):
continue
partition_key_raw = item[1]
partition_key: _PartitionKey = tuple(partition_key_raw) if isinstance(partition_key_raw, list) else None # pyright: ignore[reportGeneralTypeIssues] # decoded from our own _encode_reservations output; shape validated above
entries.append((item[0], partition_key)) # mutable-ok: see comment above
return tuple(entries)
# The admission-time timestamp a hop's token/dollar checks classified their
# bucket against, stashed on the same model_call_details object so success
# accounting recomputes the identical bucket_id (int(now) // period_seconds)
@ -1248,9 +1321,29 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
)
if concurrency_reservations:
_queue_pending_concurrency_reservations(resolved_request_kwargs, concurrency_reservations)
await self._mirror_pending_reservations(
resolved_request_kwargs.get("litellm_call_id"), concurrency_reservations
)
return healthy_deployments
async def _mirror_pending_reservations(
self, call_id: object, reservations: Sequence[tuple[str, "_PartitionKey"]]
) -> None:
if not isinstance(call_id, str):
return
try:
await self.internal_usage_cache.async_set_cache(
key=_pending_reservations_cache_key(call_id),
value=_encode_reservations(reservations),
ttl=_CONCURRENCY_MIN_SAFETY_TTL_SECONDS,
litellm_parent_otel_span=None,
)
except Exception as e: # noqa: BLE001 - a failed mirror write must never block admission; the reservation still self-heals via its own TTL
verbose_proxy_logger.warning(
"model_based_tag_rate_limits_hook: failed to mirror pending reservations for call_id=%s: %s", call_id, e
)
@staticmethod
def _ttl_for(configured_limit: _ConfiguredLimit) -> int:
if configured_limit.unit == "concurrency":
@ -1415,19 +1508,35 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
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.
whose own failure exhausts the retry chain -- async_post_call_failure_hook
closes that residual case instead, via the cache mirror
`_PENDING_RESERVATIONS_CACHE_KEY_PREFIX` documents.
"""
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)
release_keys: Final = await 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], ...]:
async def _pop_pending_concurrency_keys(
self, kwargs: Mapping[str, object]
) -> tuple[tuple[str, _PartitionKey], ...]:
# Every caller of this method is itself a normal release path, so
# also clear the async_post_call_failure_hook cache mirror for the
# same call_id right here: whatever this pop is about to release
# must never be found there later and double-released.
call_id: Final = kwargs.get("litellm_call_id")
if isinstance(call_id, str):
try:
await self.internal_usage_cache.dual_cache.async_delete_cache(_pending_reservations_cache_key(call_id))
except Exception as e: # noqa: BLE001 - a failed mirror clear must never block the real release below
verbose_proxy_logger.warning(
"model_based_tag_rate_limits_hook: failed to clear mirrored reservations for call_id=%s: %s",
call_id,
e,
)
# Snapshot then remove only those exact keys, never a blanket clear:
# a sibling hop sharing this same request's model_call_details can
# still be live and appending concurrently (see the field's own
@ -1458,7 +1567,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
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)
release_keys: Final = await self._pop_pending_concurrency_keys(model_call_details)
if release_keys:
await self._release_keys(release_keys)
@ -1480,16 +1589,37 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
after the last one. This hook fires exactly once per proxy request,
at the point the proxy gives up and returns an error to the caller,
regardless of how many hops ran or whether the completion-level
callback was suppressed for this one -- the one reliable place left
to release whatever reservation is still pending.
callback was suppressed for this one.
Reads the cache mirror written by `_mirror_pending_reservations`, not
`model_call_details`: proxy/utils.py's post_call_failure_hook pops
`litellm_logging_obj` off `request_data` before invoking any callback
here ("Remove before callbacks iterate — not serialisable"), and
neither a ContextVar nor `request_data` itself survives to this
point either (see `_PENDING_RESERVATIONS_CACHE_KEY_PREFIX`'s own
docstring for why, confirmed live for each).
"""
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):
call_id: Final = request_data.get("litellm_call_id")
if not isinstance(call_id, str):
return
release_keys: Final = self._pop_pending_concurrency_keys(model_call_details)
if release_keys:
await self._release_keys(release_keys)
cache_key: Final = _pending_reservations_cache_key(call_id)
try:
raw: Final = await self.internal_usage_cache.async_get_cache(key=cache_key, litellm_parent_otel_span=None)
except Exception as e: # noqa: BLE001 - a failed mirror read must never raise into the caller's request path
verbose_proxy_logger.warning(
"model_based_tag_rate_limits_hook: failed to read mirrored reservations for call_id=%s: %s", call_id, e
)
return
release_keys: Final = _decode_reservations(raw)
if not release_keys:
return
try:
await self.internal_usage_cache.dual_cache.async_delete_cache(cache_key)
except Exception as e: # noqa: BLE001 - a failed mirror clear must never block the real release below
verbose_proxy_logger.warning(
"model_based_tag_rate_limits_hook: failed to clear mirrored reservations for call_id=%s: %s", call_id, e
)
await self._release_keys(release_keys)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
# No special-case skip for this hook's own tag_rate_limit_exceeded
@ -1501,12 +1631,12 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
# global_tag_rate_limits_hook raises the identical marker -- that
# rejection can land after this hook already reserved a slot for the
# same request, and that slot must still be released.
release_keys: Final = self._pop_pending_concurrency_keys(kwargs)
release_keys: Final = await 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(kwargs)
release_keys: Final = await self._pop_pending_concurrency_keys(kwargs)
if release_keys:
release_task: Final = asyncio.create_task(self._release_keys(release_keys))
_BACKGROUND_TASKS.add(release_task) # mutable-ok: see _BACKGROUND_TASKS's own docstring

View file

@ -2361,6 +2361,14 @@ async def test_concurrency_slot_released_by_post_call_failure_hook_on_the_final_
proxy gives up and returns an error to the caller, regardless of how
many hops ran or whether the completion-level callback was suppressed --
it must release whatever reservation is still pending at that point.
request_data here is a distinct dict object from admission's own
request_kwargs, with no litellm_logging_obj at all: proxy/utils.py's
post_call_failure_hook pops that key off request_data before invoking
any callback ("Remove before callbacks iterate — not serialisable"),
and confirmed live, request_data is a third, unrelated object from
every hop's own request_kwargs by the time this fires. litellm_call_id
is the only identifier stable across all of them.
"""
limiter = _make_limiter(time_controller)
router = _concurrency_router(limit=1)
@ -2370,13 +2378,13 @@ async def test_concurrency_slot_released_by_post_call_failure_hook_on_the_final_
# This hop's admission reserves the slot; its own failure is the chain's
# final one, so async_log_failure_event never fires for it (simulating
# litellm's has_logged_async_failure dedup blocking the callback here).
request_kwargs, _kwargs = _call_context(["end_user_id:u1"])
request_kwargs = {"metadata": {"tags": ["end_user_id:u1"]}, "litellm_call_id": "call-final"}
await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs
)
await limiter.async_post_call_failure_hook(
request_data=request_kwargs,
request_data={"litellm_call_id": "call-final"},
original_exception=Exception("all deployments failed"),
user_api_key_dict=UserAPIKeyAuth(api_key="hash"),
)
@ -3397,11 +3405,12 @@ def test_concurrency_ttl_floor_does_not_shorten_a_longer_period_seconds():
@pytest.mark.asyncio
async def test_release_in_a_forked_task_is_visible_to_the_parent_context():
async def test_release_in_a_forked_task_is_visible_to_the_parent_context(time_controller):
limiter = _make_limiter(time_controller)
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
async def detached_release():
return _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details)
return await limiter._pop_pending_concurrency_keys(model_call_details)
released = await asyncio.create_task(detached_release())
assert released == ("key1",)
@ -3411,11 +3420,12 @@ async def test_release_in_a_forked_task_is_visible_to_the_parent_context():
@pytest.mark.asyncio
async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot():
async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot(time_controller):
limiter = _make_limiter(time_controller)
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
async def detached_release_then_sibling_admits():
released = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details)
released = await limiter._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.
model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD].append("key2")
@ -3428,10 +3438,11 @@ async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot():
@pytest.mark.asyncio
async def test_release_is_not_repeated_for_the_same_snapshot():
async def test_release_is_not_repeated_for_the_same_snapshot(time_controller):
limiter = _make_limiter(time_controller)
model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]}
first = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details)
second = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details)
first = await limiter._pop_pending_concurrency_keys(model_call_details)
second = await limiter._pop_pending_concurrency_keys(model_call_details)
assert first == ("key1",)
assert second == ()