fix: refund max_parallel_requests on disconnect from outer streaming generators

The cancellation refund previously lived in async_post_call_streaming_iterator_hook,
but that hook is nested inside the outer streaming generators and a nested async
generator only receives GeneratorExit on garbage collection (non-deterministic).
With only the v3 limiter enabled, /chat/completions also bypasses the hook entirely
(needs_iterator_wrap() is false). Move the release into async_data_generator and
async_streaming_data_generator, the generators Starlette closes on client disconnect,
so the refund fires deterministically on every streaming route. Warn when no event
loop is running, and document the window TTL refresh on the decrement
This commit is contained in:
Armaan Sandhu 2026-06-09 20:41:53 +05:30
parent b2cf030291
commit 1af00f49b8
5 changed files with 208 additions and 83 deletions

View file

@ -2061,6 +2061,17 @@ class ProxyBaseLLMRequestProcessing:
)
)
yield serialize_chunk(chunk)
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit
# are BaseException and bypass the success/failure logging
# callbacks that release the pre-call max_parallel_requests +1;
# release it here. This is the outermost generator Starlette closes
# on disconnect, so the nested iterator hook (which only sees
# GeneratorExit on GC) cannot own the refund.
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(

View file

@ -2963,6 +2963,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
rate_limit_type="max_parallel_requests",
),
increment_value=-1,
# Refresh the window TTL on the decrement, matching the
# failure path. max_parallel_requests is a concurrency
# gauge, not a rolling-window count, so the key must
# outlive in-flight requests rather than expire mid-stream.
ttl=self.window_size,
)
],

View file

@ -7097,6 +7097,17 @@ async def async_data_generator( # noqa: PLR0915
if not request_data.get("_litellm_skip_openai_stream_done"):
done_message = "[DONE]"
yield f"data: {done_message}\n\n"
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit are
# BaseException, so they bypass the success/failure logging callbacks
# that normally release the pre-call max_parallel_requests +1; release
# it here. This is the outermost generator Starlette closes on
# disconnect, so it fires reliably regardless of needs_iterator_wrap
# (a nested iterator hook would only see GeneratorExit on GC).
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(

View file

@ -2616,12 +2616,8 @@ class ProxyLogging:
# through each of them adds N pass-through trampolines per chunk for
# zero behavior change. Skip the chain entirely and stream through.
if not caps.iterator_overrides:
try:
async for chunk in response:
yield chunk
except (asyncio.CancelledError, GeneratorExit):
self._release_max_parallel_requests_on_disconnect(user_api_key_dict)
raise
async for chunk in response:
yield chunk
ProxyLogging._fire_deferred_stream_logging(request_data)
return
@ -2665,12 +2661,8 @@ class ProxyLogging:
)
# Actually iterate through the chained async generator and yield chunks
try:
async for chunk in current_response:
yield chunk
except (asyncio.CancelledError, GeneratorExit):
self._release_max_parallel_requests_on_disconnect(user_api_key_dict)
raise
async for chunk in current_response:
yield chunk
# Fire deferred logging AFTER all guardrail end-of-stream blocks
# completed. unified_guardrail writes guardrail_information during
@ -2706,7 +2698,14 @@ class ProxyLogging:
response is cancelled mid-flight (client disconnect). Neither the
success nor failure logging callback fires on the resulting
CancelledError / GeneratorExit, so the pre-call +1 would otherwise
leak. Scheduled fire-and-forget (no await) because awaiting is not
leak.
Must be called from the outermost streaming generator (the one
Starlette drives and closes on disconnect). A nested iterator-hook
generator only receives GeneratorExit when it is garbage collected,
which is non-deterministic, so the refund cannot live there.
Scheduled fire-and-forget (no await) because awaiting is not
permitted while unwinding a GeneratorExit.
"""
limiter = self.get_proxy_hook("parallel_request_limiter")
@ -2719,7 +2718,13 @@ class ProxyLogging:
)
)
except RuntimeError:
pass
# No running event loop (e.g. interpreter/loop shutdown); the
# counter's window TTL will reclaim the slot.
verbose_proxy_logger.warning(
"parallel_request_limiter_v3: could not schedule "
"max_parallel_requests release on disconnect; no running "
"event loop. Slot will be reclaimed when its window TTL expires"
)
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
"""

View file

@ -6,6 +6,7 @@ import asyncio
import os
import sys
import time
from contextlib import contextmanager
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional
@ -3135,6 +3136,36 @@ async def _seed_max_parallel_requests_counter(
)
async def _build_seeded_limiter():
"""Build a v3 limiter whose api-key counter already holds the pre-call +1."""
api_key = hash_token("sk-disconnect")
cache = DualCache()
limiter = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(cache)
)
counter_key = f"{{api_key:{api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(cache, counter_key, limiter.window_size)
user_api_key_dict = UserAPIKeyAuth(api_key=api_key, max_parallel_requests=2)
return limiter, cache, counter_key, user_api_key_dict
@contextmanager
def _override_litellm_callbacks(new_callbacks):
"""Swap litellm.callbacks so _callback_capabilities recomputes deterministically."""
saved = litellm.callbacks
litellm.callbacks = new_callbacks
try:
yield
finally:
litellm.callbacks = saved
async def _drain_release_task():
# The disconnect release is scheduled fire-and-forget via create_task.
for _ in range(5):
await asyncio.sleep(0)
@pytest.mark.asyncio
async def test_release_max_parallel_requests_on_disconnect_v3():
"""
@ -3187,84 +3218,147 @@ async def test_release_max_parallel_requests_on_disconnect_noop_v3():
assert await local_cache.async_get_cache(key=counter_key) is None
@pytest.mark.parametrize("disconnect", ["cancel", "aclose"])
@pytest.mark.asyncio
async def test_streaming_iterator_hook_releases_counter_on_cancel_v3():
async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3(
disconnect,
):
"""
End-to-end regression for issue #27955. When the client cancels a stream
mid-flight, async_post_call_streaming_iterator_hook must release the
api-key max_parallel_requests reservation and re-raise the CancelledError.
Before the fix the counter stayed at +1 because the cancellation slipped
past the hook's post-loop cleanup.
Regression for issue #27955 on the outer SSE generator (used by /v1/messages
and other event-stream routes). A client that disconnects mid-stream raises
GeneratorExit (aclose) or CancelledError into async_streaming_data_generator;
both are BaseException and bypass the success/failure logging callbacks, so
the generator itself must refund the pre-call max_parallel_requests +1.
Releasing inside the nested iterator hook does not work because that
generator is only closed on garbage collection, which is non-deterministic.
"""
_api_key = hash_token("sk-12345")
limiter_cache = DualCache()
limiter = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(limiter_cache)
)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(
limiter_cache, counter_key, limiter.window_size
)
assert await limiter_cache.async_get_cache(key=counter_key) == 1
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
assert await cache.async_get_cache(key=counter_key) == 1
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
async def cancelling_stream():
async def upstream():
yield ModelResponse()
raise asyncio.CancelledError()
with pytest.raises(asyncio.CancelledError):
async for _ in proxy_logging_obj.async_post_call_streaming_iterator_hook(
response=cancelling_stream(),
user_api_key_dict=user_api_key_dict,
request_data={},
):
pass
# The release is scheduled fire-and-forget via create_task; let it run.
for _ in range(5):
await asyncio.sleep(0)
assert await limiter_cache.async_get_cache(key=counter_key) == 0
@pytest.mark.asyncio
async def test_streaming_iterator_hook_releases_counter_on_aclose_v3():
"""
Companion to the cancel test for issue #27955. When the client disconnects,
the nested streaming generators are closed, which throws GeneratorExit (not
CancelledError) into the iterator hook. That path must also release the
api-key max_parallel_requests reservation.
"""
_api_key = hash_token("sk-12345")
limiter_cache = DualCache()
limiter = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(limiter_cache)
)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(
limiter_cache, counter_key, limiter.window_size
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
async def open_stream():
if disconnect == "cancel":
raise asyncio.CancelledError()
while True:
yield ModelResponse()
hook_iter = proxy_logging_obj.async_post_call_streaming_iterator_hook(
response=open_stream(),
user_api_key_dict=user_api_key_dict,
request_data={},
)
await hook_iter.__anext__()
await hook_iter.aclose()
with _override_litellm_callbacks([]):
gen = ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "claude-test"},
proxy_logging_obj=proxy_logging_obj,
)
await gen.__anext__()
if disconnect == "cancel":
with pytest.raises(asyncio.CancelledError):
await gen.__anext__()
else:
await gen.aclose()
await _drain_release_task()
for _ in range(5):
await asyncio.sleep(0)
assert await limiter_cache.async_get_cache(key=counter_key) == 0
assert await cache.async_get_cache(key=counter_key) == 0
@pytest.mark.parametrize("disconnect", ["cancel", "aclose"])
@pytest.mark.asyncio
async def test_async_data_generator_releases_counter_on_disconnect_v3(disconnect):
"""
Regression for issue #27955 on the chat-completions outer generator
(proxy_server.async_data_generator). With only the v3 parallel limiter
enabled, needs_iterator_wrap() is False, so this generator iterates the
upstream response directly and the iterator hook is bypassed entirely -- the
gap that let a disconnect leak the slot in the default limiter-only config.
A mid-stream disconnect must still refund the pre-call +1.
"""
import litellm.proxy.proxy_server as proxy_server
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
proxy_logging_obj = proxy_server.proxy_logging_obj
saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter")
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
async def upstream():
yield ModelResponse()
if disconnect == "cancel":
raise asyncio.CancelledError()
while True:
yield ModelResponse()
try:
with _override_litellm_callbacks([]):
assert proxy_logging_obj.needs_iterator_wrap() is False
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "gpt-test"},
)
await gen.__anext__()
if disconnect == "cancel":
with pytest.raises(asyncio.CancelledError):
await gen.__anext__()
else:
await gen.aclose()
await _drain_release_task()
assert await cache.async_get_cache(key=counter_key) == 0
finally:
if saved_hook is not None:
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = (
saved_hook
)
else:
proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None)
@pytest.mark.asyncio
async def test_async_data_generator_releases_counter_when_wrapped_v3():
"""
Companion to the no-wrap case for issue #27955. With an iterator-override
callback active, needs_iterator_wrap() is True and async_data_generator
drives the chained iterator hook. The refund must still fire exactly once
from the outer generator: the counter returns to 0 (not -1), proving the
nested hook does not also refund and there is no double decrement.
"""
from litellm.integrations.custom_logger import CustomLogger
import litellm.proxy.proxy_server as proxy_server
class _PassthroughIteratorOverride(CustomLogger):
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict, response, request_data
):
async for chunk in response:
yield chunk
limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter()
proxy_logging_obj = proxy_server.proxy_logging_obj
saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter")
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
async def upstream():
while True:
yield ModelResponse()
try:
with _override_litellm_callbacks([_PassthroughIteratorOverride()]):
assert proxy_logging_obj.needs_iterator_wrap() is True
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={"model": "gpt-test"},
)
await gen.__anext__()
await gen.aclose()
await _drain_release_task()
assert await cache.async_get_cache(key=counter_key) == 0
finally:
if saved_hook is not None:
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = (
saved_hook
)
else:
proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None)