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)