mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
137b854c50
commit
2b99f4af33
2 changed files with 138 additions and 11 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue