diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index 795d1a61bc2..0a310116f96 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -74,7 +74,7 @@ async def prefetch_proxy_stream_for_timing(completion_response: object) -> None: not is_proxy_stream_header_prefetch.get() or is_internal_call.get() or not isinstance(completion_response, CustomStreamWrapper) - or completion_response.custom_llm_provider != "vertex_ai_beta" + or not completion_response.prefetch_for_proxy_stream_headers or completion_response.completion_stream is not None or completion_response.make_call is None ): diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 1e0b778d244..6197420056e 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -235,9 +235,11 @@ class CustomStreamWrapper: stream_options=None, make_call: Callable | None = None, _response_headers: dict | httpx.Headers | None = None, + prefetch_for_proxy_stream_headers: bool = False, ): self.model = model self.make_call = make_call + self.prefetch_for_proxy_stream_headers = prefetch_for_proxy_stream_headers self.custom_llm_provider = custom_llm_provider self.logging_obj: LiteLLMLoggingObject = logging_obj self.completion_stream = completion_stream 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 d8b1e7ba17c..8f2a91457cc 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 @@ -2744,6 +2744,7 @@ class VertexLLM(VertexBase): model=model, custom_llm_provider="vertex_ai_beta", logging_obj=logging_obj, + prefetch_for_proxy_stream_headers=True, ) return streaming_response @@ -3026,6 +3027,7 @@ class VertexLLM(VertexBase): model=model, custom_llm_provider="vertex_ai_beta", logging_obj=logging_obj, + prefetch_for_proxy_stream_headers=True, ) return streaming_response diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 2cadb2fdb23..485d4a561cf 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -410,6 +410,17 @@ async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None pass +async def _close_unowned_stream(response: object) -> None: + aclose: Final = getattr(response, "aclose", None) + if aclose is None: + return + with anyio.CancelScope(shield=True): + try: + await aclose() + except Exception as exc: # noqa: BLE001 + verbose_proxy_logger.debug("Error closing unowned response stream: %s", exc) + + @lru_cache(maxsize=512) def _litellm_model_supports_stream_options(litellm_model: str) -> bool: try: @@ -2320,8 +2331,6 @@ class ProxyBaseLLMRequestProcessing: await _cancel_pending_gather_tasks(tasks) response = responses[1] - response_ownership_transferred = False - response_requires_cleanup = False _exception_raised = False try: @@ -2350,7 +2359,6 @@ class ProxyBaseLLMRequestProcessing: if self._is_streaming_request( data=self.data, is_streaming_request=is_streaming_request ) or self._is_streaming_response(response): # use generate_responses to stream responses - response_requires_cleanup = self._is_streaming_response(response) custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, call_id=logging_obj.litellm_call_id, @@ -2434,7 +2442,6 @@ class ProxyBaseLLMRequestProcessing: if route_type == "allm_passthrough_route": # Check if response is an async generator if self._is_streaming_response(response): - response_ownership_transferred = True if asyncio.iscoroutine(response): generator = await response else: @@ -2505,7 +2512,6 @@ class ProxyBaseLLMRequestProcessing: headers=custom_headers, request=request, ) - response_ownership_transferred = True return anthropic_stream_response # Non-streaming response - fall through to normal response handling elif select_data_generator: @@ -2538,7 +2544,6 @@ class ProxyBaseLLMRequestProcessing: headers=custom_headers, request=request, ) - response_ownership_transferred = True return responses_stream_response ### CALL HOOKS ### - modify outgoing data @@ -2580,18 +2585,16 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, response=response, ) + except asyncio.CancelledError: + if self._is_streaming_response(response): + await _close_unowned_stream(response) + raise except Exception: _exception_raised = True + if self._is_streaming_response(response): + await _close_unowned_stream(response) raise finally: - if response_requires_cleanup and not response_ownership_transferred: - aclose = getattr(response, "aclose", None) - if aclose is not None: - with anyio.CancelScope(shield=True): - try: - await aclose() - except Exception as exc: # noqa: BLE001 - verbose_proxy_logger.debug("Error closing unowned response stream: %s", exc) ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, exception_raised=_exception_raised, diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/test_response_metadata.py b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_response_metadata.py index c146e41acd2..a4421344249 100644 --- a/tests/test_litellm/litellm_core_utils/llm_response_utils/test_response_metadata.py +++ b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_response_metadata.py @@ -281,6 +281,7 @@ class TestPrefetchedStreamTiming: logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", make_call=AsyncMock(return_value=empty_stream()), + prefetch_for_proxy_stream_headers=True, ) token = is_proxy_stream_header_prefetch.set(True) @@ -293,14 +294,14 @@ class TestPrefetchedStreamTiming: assert dict(logging_obj.set_response_timing_metrics.call_args.args[0])["litellm_overhead_time_ms"] > 0 @pytest.mark.asyncio - async def test_prefetch_does_not_open_non_vertex_stream(self): + async def test_prefetch_does_not_open_stream_without_capability(self): logging_obj = _messages_logging_obj() logging_obj.start_time = datetime.datetime.now() - datetime.timedelta(milliseconds=1000) stream = litellm.CustomStreamWrapper( completion_stream=None, - model="other-model", + model="gemini-3.5-flash", logging_obj=logging_obj, - custom_llm_provider="openai", + custom_llm_provider="vertex_ai_beta", make_call=AsyncMock(return_value=object()), ) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py index c0e07ea0830..20471f4e5cb 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py @@ -74,6 +74,7 @@ def _deferred_stream(provider: str = "vertex_ai_beta") -> litellm.CustomStreamWr logging_obj=logging_obj, custom_llm_provider=provider, make_call=AsyncMock(return_value=_EmptyAsyncStream()), + prefetch_for_proxy_stream_headers=provider == "vertex_ai_beta", ) diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py b/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py index bde3767512d..e3eeeb482cb 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py @@ -96,6 +96,7 @@ async def test_async_responses_bridge_keeps_sdk_deferred_gemini_stream_lazy(): model="gemini-3.5-flash", logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", + prefetch_for_proxy_stream_headers=True, make_call=AsyncMock(return_value=_EmptyAsyncStream()), ) @@ -123,6 +124,7 @@ async def test_async_responses_bridge_prefetches_deferred_gemini_stream_for_prox model="gemini-3.5-flash", logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", + prefetch_for_proxy_stream_headers=True, make_call=AsyncMock(return_value=_EmptyAsyncStream()), ) token = is_proxy_stream_header_prefetch.set(True) @@ -153,6 +155,7 @@ async def test_async_responses_bridge_does_not_prefetch_already_connected_gemini model="gemini-3.5-flash", logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", + prefetch_for_proxy_stream_headers=True, make_call=AsyncMock(return_value=_EmptyAsyncStream()), ) token = is_proxy_stream_header_prefetch.set(True) @@ -182,6 +185,7 @@ async def test_async_responses_bridge_closes_prefetched_gemini_stream_on_early_c model="gemini-3.5-flash", logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", + prefetch_for_proxy_stream_headers=True, make_call=AsyncMock(return_value=upstream), ) token = is_proxy_stream_header_prefetch.set(True) @@ -217,6 +221,7 @@ async def test_async_responses_bridge_propagates_initial_fetch_failure(): model="gemini-3.5-flash", logging_obj=logging_obj, custom_llm_provider="vertex_ai_beta", + prefetch_for_proxy_stream_headers=True, make_call=AsyncMock(side_effect=expected_error), ) token = is_proxy_stream_header_prefetch.set(True)