From 1af00f49b8be2d65c65ed2e58e2eac118d841db3 Mon Sep 17 00:00:00 2001 From: Armaan Sandhu Date: Tue, 9 Jun 2026 20:41:53 +0530 Subject: [PATCH] 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 --- litellm/proxy/common_request_processing.py | 11 + .../hooks/parallel_request_limiter_v3.py | 4 + litellm/proxy/proxy_server.py | 11 + litellm/proxy/utils.py | 33 +-- .../hooks/test_parallel_request_limiter_v3.py | 232 ++++++++++++------ 5 files changed, 208 insertions(+), 83 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6558543370d..f1a8e9706ea 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 7516e43e764..66424a8136c 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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, ) ], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 213f682b8f7..c7cf0182c30 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9d7344e5196..471bcb9c612 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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): """ diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index cc8b04a875c..dfe6aebfbf2 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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)