diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index b9a9f3cebb7..c257073b4fa 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 6b70cea65a3..49744080317 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 ): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cc2f4298474..b813b152bf7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ebd5b5d90cd..4aa555164b0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 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 d10311b9f41..6d6485d219d 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 @@ -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)