fix(rate-limiting): release tag concurrency reservations on client disconnect

A client disconnecting before the first streamed chunk raises CancelledError/
GeneratorExit, which bypasses both async_log_success_event and
async_log_failure_event -- the only two places tag_rate_limiter releases a
concurrency reservation queued at admission. Without a release, the
reservation sits held until the 1-hour safety-net TTL, letting a caller
exhaust their own tag's concurrency limit for free by repeatedly opening and
dropping streaming requests.

Adds an optional, default-no-op async_release_disconnect_state_hook on
CustomLogger, wires it into the proxy's shielded streaming-disconnect
cleanup (only when no disconnect-time success event already fired), and
implements it in tag_rate_limiter to release pending reservations.
This commit is contained in:
Deepanshu 2026-08-19 11:19:09 -04:00
parent 69004900b1
commit 9b3891ee9a
5 changed files with 175 additions and 0 deletions

View file

@ -765,6 +765,22 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
return
async def async_release_disconnect_state_hook(self) -> None:
"""
Release per-request state reserved outside of `async_log_success_event`
/ `async_log_failure_event` for a request whose streaming response is
torn down by a client disconnect: `CancelledError` / `GeneratorExit`
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
idempotent and never raise -- a callback that never reserved such
state has nothing to do here.
Default does nothing.
"""
return
async def async_should_run_chat_completion_agentic_loop(
self,
response: Any,

View file

@ -33,6 +33,7 @@ from litellm.constants import (
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 (
@ -397,6 +398,32 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
return True
async def _release_disconnect_state_on_all_callbacks() -> 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()
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:
@ -3364,6 +3391,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()
if hasattr(response, "aclose"):
try:

View file

@ -1126,6 +1126,20 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer
pass
return keys
async def async_release_disconnect_state_hook(self) -> None:
"""
A client disconnecting before the first streamed chunk raises
CancelledError/GeneratorExit, which bypasses both async_log_success_event
and async_log_failure_event below -- the only two places a concurrency
reservation queued during admission is normally popped and released.
Without this, the reservation would sit held until _CONCURRENCY_MIN_SAFETY_TTL_SECONDS
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()
if release_keys:
await self._release_keys(release_keys)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
if isinstance(kwargs.get("exception"), ProxyRateLimitError):
detail: Final = (

View file

@ -1017,6 +1017,47 @@ async def test_concurrency_slot_released_on_success_frees_capacity(time_controll
assert result == healthy
@pytest.mark.asyncio
async def test_concurrency_slot_released_on_disconnect_frees_capacity(time_controller):
"""
A client disconnecting before the first streamed chunk raises
CancelledError/GeneratorExit, which bypasses both async_log_success_event
and async_log_failure_event entirely -- neither fires, so the reservation
would otherwise sit held until the safety-net TTL. The proxy's disconnect
cleanup calls async_release_disconnect_state_hook instead in that case.
"""
limiter = _make_limiter(time_controller)
router = _concurrency_router(limit=1)
limiter.update_variables(llm_router=router)
healthy = router.model_list
kwargs = {"metadata": {"tags": ["end_user_id:u1"]}}
await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs
)
# At capacity: a second concurrent request is rejected.
with pytest.raises(ProxyRateLimitError):
await limiter.async_filter_deployments(
model="grp",
healthy_deployments=healthy,
messages=None,
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
)
# 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()
result = await limiter.async_filter_deployments(
model="grp",
healthy_deployments=healthy,
messages=None,
request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}},
)
assert result == healthy
@pytest.mark.asyncio
async def test_concurrency_slot_released_on_failure_frees_capacity(time_controller):
limiter = _make_limiter(time_controller)

View file

@ -5596,6 +5596,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) -> None:
self.disconnect_hook_calls += 1
class TestStreamingClientDisconnectBilling:
"""
A client disconnect throws GeneratorExit into the proxy streaming
@ -5990,6 +5999,72 @@ 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):
"""
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()
original_callbacks = litellm.callbacks
litellm.callbacks = [recorder]
try:
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(),
),
)
finally:
litellm.callbacks = original_callbacks
assert recorder.disconnect_hook_calls == 1
@pytest.mark.asyncio
async def test_disconnect_billing_skips_callback_disconnect_hook(self):
"""
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()
original_callbacks = litellm.callbacks
litellm.callbacks = [recorder]
try:
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(),
),
)
finally:
litellm.callbacks = original_callbacks
assert recorder.disconnect_hook_calls == 0
def _apply_stream_usage_tracking(
data: dict,