diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 315fbcba310..09d52b78da8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, ) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 64318778bc2..f854d8f94e5 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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,