fix(proxy): release rate limit hook state on client disconnect

A client disconnect throws GeneratorExit/CancelledError into the
request path, so neither the success nor failure logging callback
runs and a concurrency slot reserved at admission leaks until its own
safety TTL. Gives every registered CustomLogger a chance to release
such state via the new async_release_disconnect_state_hook, called
from both the streaming and non-streaming cancel-on-disconnect paths.
This commit is contained in:
Deepanshu 2026-08-25 22:01:57 -04:00
parent 137b854c50
commit 2b99f4af33
2 changed files with 138 additions and 11 deletions

View file

@ -29,10 +29,10 @@ from litellm.constants import (
NON_INFERENCE_CALL_TYPES,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
STREAM_SSE_DATA_PREFIX,
STREAM_SSE_KEEPALIVE_PING_BYTES,
UNSAFE_PROXY_RESPONSE_HEADERS,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.get_supported_openai_params import (
@ -203,10 +203,6 @@ _CLIENT_DISCONNECTED_ERROR_INFORMATION: Final[StandardLoggingPayloadErrorInforma
}
def _withheld_provider_output(response: object) -> bool:
return getattr(response, "has_buffered_provider_output", False) is True
def _should_return_raw_model_name(request_data: dict[str, object]) -> bool:
return any(
isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True
@ -397,6 +393,32 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
return True
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
(see the callers of this function). A callback that reserves per-request
state outside of those two callbacks (e.g. a concurrency slot admitted
before the first chunk) would otherwise leak that state until its own
safety-net TTL. Give every registered callback a chance to release such
state via the optional, default-no-op ``async_release_disconnect_state_hook``.
Only ``CustomLogger`` instances are considered, never raw string entries:
by the time a request can reach this proxy-only cleanup path, startup's
``ProxyLogging._init_litellm_callbacks`` has already replaced every string
entry in ``litellm.callbacks`` with its initialized instance in place.
"""
for callback in litellm.callbacks:
if not isinstance(callback, CustomLogger):
continue
try:
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
)
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
pending_tasks: Final = [task for task in tasks if not task.done()]
for task in pending_tasks:
@ -1454,6 +1476,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))
@ -1461,6 +1484,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,
@ -2285,7 +2316,9 @@ 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( # rebind-ok: assigned in exactly one of these two mutually exclusive branches
request, llm_responses, self.data
)
else:
responses = await llm_responses
finally:
@ -3364,6 +3397,8 @@ class ProxyBaseLLMRequestProcessing:
and user_api_key_dict is not None
):
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(request_data)
if hasattr(response, "aclose"):
try:
@ -3447,9 +3482,8 @@ class ProxyBaseLLMRequestProcessing:
# so a GeneratorExit on client disconnect is raised there and any
# statement after the yield never runs. The slow-path hook is
# awaited above, so a cancellation during it still leaves this
# False and refunds. A keepalive ping carries no provider output,
# so it must not suppress that refund.
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
# False and refunds.
delivered_chunk = True
yield serialize_chunk(chunk)
stream_completed = True
except (asyncio.CancelledError, GeneratorExit):
@ -3463,7 +3497,7 @@ class ProxyBaseLLMRequestProcessing:
# only sees GeneratorExit on GC) cannot own the refund.
if not stream_completed:
client_disconnected = True
if not delivered_chunk and not _withheld_provider_output(response):
if not delivered_chunk:
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation_on_cancel,
)

View file

@ -4063,7 +4063,33 @@ class TestCancelOnDisconnect:
llm_call.cancel()
with pytest.raises(asyncio.CancelledError):
await _await_llm_call_cancelling_on_disconnect(request, llm_call)
await _await_llm_call_cancelling_on_disconnect(request, llm_call, {})
async def test_disconnect_releases_callback_state_before_499(self, monkeypatch):
"""
asyncio.CancelledError is a BaseException, not an Exception, so it
never reaches litellm.utils.wrapper_async's own except block -- the
cancelled call's async_log_failure_event never fires, and the 499
this raises is later handled by post_call_failure_hook, a different
hook a CustomLogger like model_based_tag_rate_limits_hook 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
@ -5596,6 +5622,15 @@ class _RecordingSuccessLogger(CustomLogger):
self.success_events.append({"kwargs": kwargs, "response_obj": response_obj})
class _RecordingDisconnectHookLogger(CustomLogger):
def __init__(self):
super().__init__()
self.disconnect_hook_calls = 0
async def async_release_disconnect_state_hook(self, request_data: dict) -> None:
self.disconnect_hook_calls += 1
class TestStreamingClientDisconnectBilling:
"""
A client disconnect throws GeneratorExit into the proxy streaming
@ -5990,6 +6025,64 @@ class TestStreamingClientDisconnectBilling:
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.cached_tokens == 7
@pytest.mark.asyncio
async def test_disconnect_without_billable_chunks_releases_callback_state(self, monkeypatch):
"""
A callback that reserves per-request state outside of the success/failure
logging callbacks (e.g. a concurrency slot admitted before the first
chunk) would otherwise leak it on a disconnect with nothing to bill,
since neither logging callback ever fires for it. The disconnect
cleanup must give every registered callback a chance to release such
state via async_release_disconnect_state_hook.
"""
import types
response = await self._start_partial_stream()
empty_response = types.SimpleNamespace(chunks=[], messages=None)
recorder = _RecordingDisconnectHookLogger()
monkeypatch.setattr(litellm, "callbacks", [recorder])
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=None,
request_data={"litellm_logging_obj": response.logging_obj},
response=empty_response,
stream_completed=False,
client_disconnected=True,
user_api_key_dict=MagicMock(),
proxy_logging_obj=types.SimpleNamespace(
_arelease_max_parallel_requests_on_disconnect=AsyncMock(),
),
)
assert recorder.disconnect_hook_calls == 1
@pytest.mark.asyncio
async def test_disconnect_billing_skips_callback_disconnect_hook(self, monkeypatch):
"""
When a disconnect-time success event already fired (partial billing
dispatched it), that event's own async_log_success_event already ran
for every registered callback. The disconnect hook must not also run
in that case, so a callback with idempotent-but-not-free release logic
does not do redundant work on every disconnect.
"""
import types
recorder = _RecordingDisconnectHookLogger()
monkeypatch.setattr(litellm, "callbacks", [recorder])
response = await self._start_partial_stream()
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=None,
request_data={"litellm_logging_obj": response.logging_obj},
response=response,
stream_completed=False,
client_disconnected=True,
user_api_key_dict=MagicMock(),
proxy_logging_obj=types.SimpleNamespace(
_arelease_max_parallel_requests_on_disconnect=AsyncMock(),
),
)
assert recorder.disconnect_hook_calls == 0
def _apply_stream_usage_tracking(
data: dict,