diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 0aeb83a64ba..e147be892f5 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -169,8 +169,11 @@ class CustomStreamWrapper: result = self.completion_stream.close() if result is not None: await result - except BaseException: - pass + except BaseException as e: + verbose_logger.debug( + "CustomStreamWrapper.aclose: error closing completion_stream: %s", + e, + ) def check_send_stream_usage(self, stream_options: Optional[dict]): return ( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1421a47dede..390d68813e6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5149,8 +5149,10 @@ async def async_data_generator( if hasattr(response, "aclose"): try: await response.aclose() - except Exception: - pass + except Exception as e: + verbose_proxy_logger.debug( + "async_data_generator: error closing response stream: %s", e + ) def select_data_generator( diff --git a/litellm/router.py b/litellm/router.py index 0c55d93066a..bbd7330c62a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1603,11 +1603,23 @@ class Router: # Shield from anyio cancellation so the awaits can complete. with anyio.CancelScope(shield=True): if hasattr(model_response, "aclose"): - await model_response.aclose() + try: + await model_response.aclose() + except BaseException as e: + verbose_router_logger.debug( + "stream_with_fallbacks: error closing model_response: %s", + e, + ) if fallback_response is not None and hasattr( fallback_response, "aclose" ): - await fallback_response.aclose() + try: + await fallback_response.aclose() + except BaseException as e: + verbose_router_logger.debug( + "stream_with_fallbacks: error closing fallback_response: %s", + e, + ) return FallbackStreamWrapper(stream_with_fallbacks()) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e57523ea058..c4d549b72be 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3433,3 +3433,50 @@ async def test_async_data_generator_cleanup_on_normal_completion(): assert any("[DONE]" in d for d in yielded_data) # aclose should still be called via finally block mock_response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_async_data_generator_cleanup_on_midstream_error(): + """ + Test that async_data_generator calls response.aclose() via finally block + even when an exception occurs mid-stream. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + + async def mock_streaming_iterator_with_error(*args, **kwargs): + yield {"choices": [{"delta": {"content": "Hello"}}]} + raise RuntimeError("upstream connection reset") + + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( + mock_streaming_iterator_with_error + ) + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( + side_effect=lambda **kwargs: kwargs.get("response") + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + yielded_data = [] + async for data in async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ): + yielded_data.append(data) + + # Should have yielded data chunk and then an error chunk + assert len(yielded_data) >= 2 + assert any("error" in d for d in yielded_data) + # aclose must still be called via finally block despite the error + mock_response.aclose.assert_awaited_once() diff --git a/tests/test_litellm/test_streaming_connection_cleanup.py b/tests/test_litellm/test_streaming_connection_cleanup.py index fb9a99436f2..677046bc66c 100644 --- a/tests/test_litellm/test_streaming_connection_cleanup.py +++ b/tests/test_litellm/test_streaming_connection_cleanup.py @@ -20,6 +20,9 @@ from litellm.llms.custom_httpx.aiohttp_transport import ( ) +# ── aiohttp transport layer tests ────────────────────────────── + + @pytest.mark.asyncio async def test_aiohttp_transport_response_uses_stream_not_content(): """handle_async_request must use stream= so aclose() propagates to AiohttpResponseStream.""" @@ -88,6 +91,9 @@ async def test_aiohttp_response_stream_aclose_releases_connection(): assert aexit_called +# ── CustomStreamWrapper.aclose() tests ───────────────────────── + + @pytest.mark.asyncio async def test_aclose_falls_back_to_close(): """OpenAI's AsyncStream has close() but not aclose(). Must fall back.""" @@ -161,27 +167,37 @@ async def test_aclose_completes_under_cancellation(): assert aclose_completed +# ── Router stream_with_fallbacks cleanup tests ────────────────── + + @pytest.mark.asyncio async def test_stream_with_fallbacks_closes_stream_on_generator_close(): - """Closing the generator from async_function_with_fallbacks must aclose() the stream.""" + """Closing the FallbackStreamWrapper must aclose() the underlying model_response + via stream_with_fallbacks' finally block.""" from litellm.router import Router stream_closed = False - class FakeStream: + class FakeStream(CustomStreamWrapper): def __init__(self): - self.chunks = ["chunk1", "chunk2", "chunk3"] - self.index = 0 + super().__init__( + completion_stream=None, + model="test-model", + logging_obj=MagicMock(), + custom_llm_provider="openai", + ) + self._items = ["chunk1", "chunk2", "chunk3"] + self._index = 0 def __aiter__(self): return self async def __anext__(self): - if self.index >= len(self.chunks): + if self._index >= len(self._items): raise StopAsyncIteration - chunk = self.chunks[self.index] - self.index += 1 - return chunk + item = self._items[self._index] + self._index += 1 + return item async def aclose(self): nonlocal stream_closed @@ -201,74 +217,53 @@ async def test_stream_with_fallbacks_closes_stream_on_generator_close(): fake_stream = FakeStream() - with patch.object(router, "acompletion", return_value=fake_stream): - result = await router.async_function_with_fallbacks( - original_function=router.acompletion, - model="test-model", - messages=[{"role": "user", "content": "hi"}], - stream=True, - num_retries=0, - ) + # Call _acompletion_streaming_iterator directly so we go through + # stream_with_fallbacks and its finally block + result = await router._acompletion_streaming_iterator( + model_response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + initial_kwargs={"model": "test-model"}, + ) - async for chunk in result: - break + # Consume one chunk then close (simulates client disconnect) + async for _ in result: + break + await result.aclose() - await result.aclose() - - assert stream_closed + assert stream_closed, "model_response stream was not closed by stream_with_fallbacks finally block" @pytest.mark.asyncio -async def test_stream_with_fallbacks_closes_fallback_response_on_disconnect(): - """When stream_with_fallbacks is closed during fallback iteration, - both model_response and fallback_response must be closed.""" +async def test_stream_with_fallbacks_closes_stream_on_normal_completion(): + """stream_with_fallbacks must aclose() model_response even on normal completion.""" from litellm.router import Router - model_closed = False - fallback_closed = False - - class FakeModelStream: - """Simulates a stream that fails mid-stream, triggering fallback.""" + stream_closed = False + class FakeStream(CustomStreamWrapper): def __init__(self): - self.chunks = [] - self.model = "test-model" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() + super().__init__( + completion_stream=None, + model="test-model", + logging_obj=MagicMock(), + custom_llm_provider="openai", + ) + self._items = ["chunk1"] + self._index = 0 def __aiter__(self): return self async def __anext__(self): - raise StopAsyncIteration - - async def aclose(self): - nonlocal model_closed - model_closed = True - - class FakeFallbackStream: - """Simulates a fallback stream that yields chunks.""" - - def __init__(self): - self.items = ["fb1", "fb2", "fb3"] - self.index = 0 - - def __aiter__(self): - return self - - async def __anext__(self): - if self.index >= len(self.items): + if self._index >= len(self._items): raise StopAsyncIteration - item = self.items[self.index] - self.index += 1 + item = self._items[self._index] + self._index += 1 return item async def aclose(self): - nonlocal fallback_closed - fallback_closed = True - - # Just verify the finally block closes model_response even on normal completion - fake_model_stream = FakeModelStream() + nonlocal stream_closed + stream_closed = True router = Router( model_list=[ @@ -282,18 +277,115 @@ async def test_stream_with_fallbacks_closes_fallback_response_on_disconnect(): ] ) - with patch.object(router, "acompletion", return_value=fake_model_stream): - result = await router.async_function_with_fallbacks( - original_function=router.acompletion, - model="test-model", + fake_stream = FakeStream() + + result = await router._acompletion_streaming_iterator( + model_response=fake_stream, + messages=[{"role": "user", "content": "hi"}], + initial_kwargs={"model": "test-model"}, + ) + + # Exhaust the stream fully + async for _ in result: + pass + await result.aclose() + + assert stream_closed, "model_response stream was not closed after normal completion" + + +@pytest.mark.asyncio +async def test_stream_with_fallbacks_closes_both_on_fallback_disconnect(): + """When a fallback is triggered and the client disconnects during fallback + iteration, both model_response and fallback_response must be closed.""" + from litellm.exceptions import MidStreamFallbackError + from litellm.router import Router + + model_closed = False + fallback_closed = False + + class FakeModelStream(CustomStreamWrapper): + """Stream that raises MidStreamFallbackError immediately to trigger fallback.""" + + def __init__(self): + super().__init__( + completion_stream=None, + model="test-model", + logging_obj=MagicMock(), + custom_llm_provider="openai", + ) + self.chunks = [] + + def __aiter__(self): + return self + + async def __anext__(self): + raise MidStreamFallbackError( + message="test mid-stream error", + model="test-model", + llm_provider="openai", + generated_content="", + ) + + async def aclose(self): + nonlocal model_closed + model_closed = True + + class FakeFallbackStream: + """Fallback stream that yields chunks.""" + + def __init__(self): + self._items = ["fb1", "fb2", "fb3"] + self._index = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + if self._index >= len(self._items): + raise StopAsyncIteration + item = self._items[self._index] + self._index += 1 + return item + + async def aclose(self): + nonlocal fallback_closed + fallback_closed = True + + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test", + "api_key": "fake", + }, + } + ] + ) + + fake_model_stream = FakeModelStream() + fake_fallback_stream = FakeFallbackStream() + + # Mock async_function_with_fallbacks_common_utils to return the fallback stream + # instead of actually calling through the full fallback machinery + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=fake_fallback_stream, + ): + result = await router._acompletion_streaming_iterator( + model_response=fake_model_stream, messages=[{"role": "user", "content": "hi"}], - stream=True, - num_retries=0, + initial_kwargs={ + "model": "test-model", + "fallbacks": ["other-model"], + }, ) - # Exhaust the stream then close + # Consume one fallback chunk then close (simulates client disconnect) async for _ in result: - pass + break await result.aclose() assert model_closed, "model_response stream was not closed" + assert fallback_closed, "fallback_response stream was not closed"