fix(proxy): release max_parallel_requests slot when a stream is cancelled mid-flight (#27955) (#30020)

* fix(proxy): release max_parallel_requests slot when a stream is cancelled mid-flight (#27955)

* 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-11 16:07:31 +05:30 • committed by GitHub
parent adf7c97714
commit 35dc441092
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 343 additions and 0 deletions

View file

@ -2075,6 +2075,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

@ -2944,6 +2944,46 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
f"Error in rate limit failure event: {str(e)}"
)
async def async_release_max_parallel_requests_on_disconnect(
self, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""
Release the api-key ``max_parallel_requests`` slot that
``async_pre_call_hook`` reserved, for a request that ended without
either logging callback firing.
The +1 is normally undone by ``async_log_success_event`` (natural
stream completion) or ``async_log_failure_event`` (LLM error). When a
client cancels a stream mid-flight, the cancellation surfaces as
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
runs, so without this the counter leaks one slot per cancelled stream
until the key wedges at its limit.
"""
if (
not user_api_key_dict.api_key
or user_api_key_dict.max_parallel_requests is None
):
return
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=[
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
key="api_key",
value=user_api_key_dict.api_key,
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,
)
],
litellm_parent_otel_span=None,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response
):

View file

@ -7113,6 +7113,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

@ -137,6 +137,9 @@ from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.repositories.budget_repository import BudgetRepository
@ -2688,6 +2691,42 @@ class ProxyLogging:
logging_obj._deferred_stream_complete_args = None
asyncio.create_task(_deferred_cb(*_args))
def _release_max_parallel_requests_on_disconnect(
self, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""
Release the api-key max_parallel_requests slot when a streaming
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.
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")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
return
try:
asyncio.create_task(
limiter.async_release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
)
except RuntimeError:
# 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):
"""
Initialize the response taking too long task if user is using slack alerting

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
@ -20,6 +21,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
EmbeddingResponse,
ModelResponse,
@ -3187,3 +3189,243 @@ def test_get_key_mcp_rpm_limit_precedence():
none_set = UserAPIKeyAuth(api_key=hash_token("sk-mcp-key"))
assert get_key_mcp_rpm_limit(none_set) is None
assert get_team_mcp_rpm_limit(none_set) is None
async def _seed_max_parallel_requests_counter(
dual_cache: DualCache, counter_key: str, window_size: int
) -> None:
await dual_cache.async_increment_cache_pipeline(
increment_list=[
RedisPipelineIncrementOperation(
key=counter_key, increment_value=1, ttl=window_size
)
]
)
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():
"""
Regression for issue #27955: a stream cancelled mid-flight must release the
pre-call +1 reservation. The success/failure logging callbacks never fire
on cancellation, so without an explicit release the api-key counter climbs
by one per cancelled request until the key wedges at its limit. The release
must decrement the api-key max_parallel_requests counter by exactly one.
"""
_api_key = hash_token("sk-12345")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_counter(
local_cache, counter_key, handler.window_size
)
assert await local_cache.async_get_cache(key=counter_key) == 1
await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict)
assert await local_cache.async_get_cache(key=counter_key) == 0
@pytest.mark.asyncio
async def test_release_max_parallel_requests_on_disconnect_noop_v3():
"""
The release must be a no-op when the key never reserved a parallel slot
(no api_key, or max_parallel_requests unset). Otherwise a cancelled
no-limit request would drive an unrelated counter negative.
"""
_api_key = hash_token("sk-12345")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await handler.async_release_max_parallel_requests_on_disconnect(
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None)
)
assert await local_cache.async_get_cache(key=counter_key) is None
await handler.async_release_max_parallel_requests_on_disconnect(
UserAPIKeyAuth(api_key=None, max_parallel_requests=5)
)
assert await local_cache.async_get_cache(key=counter_key) is None
@pytest.mark.parametrize("disconnect", ["cancel", "aclose"])
@pytest.mark.asyncio
async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3(
disconnect,
):
"""
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.
"""
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
async def upstream():
yield ModelResponse()
if disconnect == "cancel":
raise asyncio.CancelledError()
while True:
yield ModelResponse()
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()
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)