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

This commit is contained in:
Armaan Sandhu 2026-06-09 16:58:32 +05:30
parent 51ba6e39cd
commit b2cf030291
3 changed files with 222 additions and 4 deletions

View file

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

@ -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
@ -2613,8 +2616,12 @@ 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:
async for chunk in response:
yield chunk
try:
async for chunk in response:
yield chunk
except (asyncio.CancelledError, GeneratorExit):
self._release_max_parallel_requests_on_disconnect(user_api_key_dict)
raise
ProxyLogging._fire_deferred_stream_logging(request_data)
return
@ -2658,8 +2665,12 @@ class ProxyLogging:
)
# Actually iterate through the chained async generator and yield chunks
async for chunk in current_response:
yield chunk
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
# Fire deferred logging AFTER all guardrail end-of-stream blocks
# completed. unified_guardrail writes guardrail_information during
@ -2687,6 +2698,29 @@ 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. 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:
pass
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

@ -20,6 +20,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 ModelResponse, Usage
@ -3120,3 +3121,150 @@ 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
)
]
)
@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.asyncio
async def test_streaming_iterator_hook_releases_counter_on_cancel_v3():
"""
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.
"""
_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
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():
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():
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()
for _ in range(5):
await asyncio.sleep(0)
assert await limiter_cache.async_get_cache(key=counter_key) == 0