diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 430a789d2a0..c23f27243e1 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -2846,6 +2846,7 @@ async def make_call( sync_stream=False, logging_obj=logging_obj, response_headers=response.headers, + response=response, ) # LOGGING logging_obj.post_call( @@ -2889,6 +2890,7 @@ def make_sync_call( sync_stream=True, logging_obj=logging_obj, response_headers=response.headers, + response=response, ) # LOGGING @@ -3350,12 +3352,14 @@ class ModelResponseIterator: sync_stream: bool, logging_obj: LoggingClass, response_headers: Optional[Dict[str, str]] = None, + response: Optional[httpx.Response] = None, ): from litellm.litellm_core_utils.prompt_templates.common_utils import ( check_is_function_call, ) self.streaming_response = streaming_response + self.response = response self.chunk_type: Literal["valid_json", "accumulated_json"] = "valid_json" self.accumulated_json = "" self.sent_first_chunk = False @@ -3653,3 +3657,37 @@ class ModelResponseIterator: raise StopAsyncIteration except ValueError as e: raise RuntimeError(f"Error parsing chunk: {e},\nReceived chunk: {chunk}") + + async def aclose(self) -> None: + iterator = getattr(self, "async_response_iterator", self.streaming_response) + if iterator is not None and hasattr(iterator, "aclose"): + try: + await iterator.aclose() + except Exception as e: + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing iterator: %s", e + ) + if self.response is not None: + try: + await self.response.aclose() + except Exception as e: + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing response: %s", e + ) + + def close(self) -> None: + iterator = getattr(self, "response_iterator", self.streaming_response) + if iterator is not None and hasattr(iterator, "close"): + try: + iterator.close() + except Exception as e: + verbose_logger.debug( + "ModelResponseIterator.close: error closing iterator: %s", e + ) + if self.response is not None: + try: + self.response.close() + except Exception as e: + verbose_logger.debug( + "ModelResponseIterator.close: error closing response: %s", e + ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6558543370d..e8a4ca43ae8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -73,6 +73,20 @@ from litellm.types.utils import ModelResponse, ModelResponseStream, Usage _DD_STREAMING_TRACE_ENABLED = not isinstance(tracer, NullTracer) +async def _check_request_disconnection( + request: Request, llm_api_call_task: "asyncio.Task[Any]" +) -> None: + start_time = time.time() + while time.time() - start_time < 600: + await asyncio.sleep(1) + try: + if await request.is_disconnected(): + llm_api_call_task.cancel() + return + except Exception: + return + + def _serialize_http_exception_detail( detail: Any, ) -> Tuple[str, Optional[dict]]: @@ -1215,14 +1229,31 @@ class ProxyBaseLLMRequestProcessing: user_model=user_model, user_api_key_dict=user_api_key_dict, ) - tasks.append(llm_call) + llm_call_task = asyncio.create_task(llm_call) + tasks.append(llm_call_task) + + disconnect_task = asyncio.create_task( + _check_request_disconnection(request, llm_call_task) + ) # wait for call to end llm_responses = asyncio.gather( *tasks ) # run the moderation check in parallel to the actual llm api call - responses = await llm_responses + try: + responses = await llm_responses + except asyncio.CancelledError: + raise HTTPException( + status_code=499, + detail="Client disconnected the request", + ) + finally: + disconnect_task.cancel() + try: + await disconnect_task + except asyncio.CancelledError: + pass response = responses[1] diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 671d7355e8f..1a2d0d86810 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -4996,3 +4996,146 @@ def test_mid_stream_429_error_raises_during_iteration(): # Verify: 429 error is properly raised assert exc_info.value.status_code == 429 assert "RESOURCE_EXHAUSTED" in str(exc_info.value.message) + + +class TestModelResponseIteratorCleanup: + def _make_logging_obj(self): + from unittest.mock import Mock + + obj = Mock() + obj.optional_params = {} + return obj + + def test_aclose_closes_iterator_and_response(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_iterator.aclose.assert_awaited_once() + mock_response.aclose.assert_awaited_once() + + def test_close_closes_iterator_and_response(self): + from unittest.mock import MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_iterator = MagicMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=True, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.response_iterator = mock_iterator + + iterator.close() + + mock_iterator.close.assert_called_once() + mock_response.close.assert_called_once() + + def test_aclose_without_response_does_not_raise(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_iterator.aclose.assert_awaited_once() + + def test_aclose_tolerates_iterator_error(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock(side_effect=RuntimeError("transport error")) + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_response.aclose.assert_awaited_once() + + def test_custom_stream_wrapper_aclose_triggers_model_response_iterator_aclose(self): + """CustomStreamWrapper.aclose() must propagate to ModelResponseIterator.aclose().""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + model_response_iter = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + model_response_iter.async_response_iterator = mock_iterator + + wrapper = CustomStreamWrapper( + completion_stream=model_response_iter, + model="gemini-2.0-flash", + custom_llm_provider="vertex_ai", + logging_obj=MagicMock(), + ) + + asyncio.run(wrapper.aclose()) + + mock_iterator.aclose.assert_awaited_once() + mock_response.aclose.assert_awaited_once() diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 0f5a0cbe4b6..18e5b8c8a99 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2342,3 +2342,145 @@ class TestAsyncStreamingDataGeneratorFastPath: hook_spy.assert_awaited_once() ProxyLogging._callback_capabilities_cache.clear() + + +class TestCheckRequestDisconnection: + @pytest.mark.asyncio + async def test_cancels_task_when_client_disconnects(self, monkeypatch): + import asyncio + + import litellm.proxy.common_request_processing as cpr + from litellm.proxy.common_request_processing import _check_request_disconnection + + task_cancelled = False + + async def never_ending(): + nonlocal task_cancelled + try: + await asyncio.sleep(9999) + except asyncio.CancelledError: + task_cancelled = True + raise + + task = asyncio.create_task(never_ending()) + mock_request = MagicMock(spec=Request) + mock_request.is_disconnected = AsyncMock(return_value=True) + monkeypatch.setattr(cpr.asyncio, "sleep", AsyncMock()) + + await _check_request_disconnection(mock_request, task) + await asyncio.sleep(0) + + assert task.cancelled() or task_cancelled + + @pytest.mark.asyncio + async def test_does_not_cancel_task_when_client_stays_connected(self, monkeypatch): + import asyncio + + import litellm.proxy.common_request_processing as cpr + from litellm.proxy.common_request_processing import _check_request_disconnection + + call_count = 0 + + async def fake_sleep(_): + nonlocal call_count + call_count += 1 + if call_count >= 3: + raise asyncio.CancelledError + + task = asyncio.create_task(asyncio.sleep(9999)) + mock_request = MagicMock(spec=Request) + mock_request.is_disconnected = AsyncMock(return_value=False) + monkeypatch.setattr(cpr.asyncio, "sleep", fake_sleep) + + with pytest.raises(asyncio.CancelledError): + await _check_request_disconnection(mock_request, task) + + assert not task.cancelled() + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + @pytest.mark.asyncio + async def test_exits_gracefully_when_is_disconnected_raises(self, monkeypatch): + import asyncio + + import litellm.proxy.common_request_processing as cpr + from litellm.proxy.common_request_processing import _check_request_disconnection + + task = asyncio.create_task(asyncio.sleep(9999)) + mock_request = MagicMock(spec=Request) + mock_request.is_disconnected = AsyncMock( + side_effect=RuntimeError("transport closed") + ) + monkeypatch.setattr(cpr.asyncio, "sleep", AsyncMock()) + + await _check_request_disconnection(mock_request, task) + + assert not task.cancelled() + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + @pytest.mark.asyncio + async def test_base_process_llm_request_raises_499_on_client_disconnect( + self, monkeypatch + ): + """When _check_request_disconnection cancels the LLM task, base_process_llm_request + must raise HTTPException(499) instead of propagating CancelledError.""" + import asyncio + + import litellm.proxy.common_request_processing as cpr + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + async def slow_llm(): + await asyncio.sleep(9999) + + async def fake_route_request(**_kwargs): + return slow_llm() + + async def instant_disconnect(request, task): + task.cancel() + + mock_logging_obj = MagicMock() + mock_logging_obj.litellm_call_id = "test-call-id" + mock_logging_obj._defer_async_logging = False + + mock_proxy_logging = MagicMock(spec=ProxyLogging) + mock_proxy_logging.during_call_hook = AsyncMock(return_value=None) + mock_proxy_logging._callback_capabilities_cache = {} + + monkeypatch.setattr(cpr, "route_request", fake_route_request) + monkeypatch.setattr(cpr, "_check_request_disconnection", instant_disconnect) + + processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"}) + monkeypatch.setattr( + processing_obj, + "common_processing_pre_call_logic", + AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), + ) + monkeypatch.setattr( + processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) + ) + + mock_request = MagicMock(spec=Request) + mock_request.is_disconnected = AsyncMock(return_value=True) + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await processing_obj.base_process_llm_request( + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + proxy_logging_obj=mock_proxy_logging, + general_settings={}, + proxy_config=MagicMock(spec=ProxyConfig), + route_type="acompletion", + version=None, + ) + + assert exc_info.value.status_code == 499 + assert "disconnected" in exc_info.value.detail.lower()