diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 62751fb68a4..7516e43e764 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 ): diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5ad42b5e1be..9d7344e5196 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 @@ -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 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 676f623a5dd..cc8b04a875c 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 @@ -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