From 6db7fc319c9897789e75009e4a7f83e02fb3a9c5 Mon Sep 17 00:00:00 2001 From: wanna Date: Tue, 24 Feb 2026 21:46:08 +0800 Subject: [PATCH] Handle stream-required fallback --- .../handler.py | 93 ++++++- litellm/llms/custom_httpx/llm_http_handler.py | 257 ++++++++++++++++-- litellm/llms/openai/openai.py | 105 ++++++- 3 files changed, 420 insertions(+), 35 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 5c051797e8b..3087b9840d8 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -2,6 +2,7 @@ Handler for transforming /chat/completions api requests to litellm.responses requests """ +import json from typing import TYPE_CHECKING, Any, Coroutine, Optional, Union from typing_extensions import TypedDict @@ -37,6 +38,57 @@ class ResponsesToCompletionBridgeHandler: stream = litellm_params.get("stream", False) return bool(stream) + @staticmethod + def _contains_stream_required_text(value: Any) -> bool: + if value is None: + return False + if isinstance(value, (bytes, bytearray)): + try: + value = value.decode("utf-8", errors="ignore") + except Exception: + value = str(value) + if isinstance(value, str): + lowered = value.lower() + if "stream must be set to true" in lowered: + return True + try: + parsed = json.loads(value) + except Exception: + return False + return ResponsesToCompletionBridgeHandler._contains_stream_required_text( + parsed + ) + if isinstance(value, dict): + for key in ("detail", "message", "error"): + if key in value and ResponsesToCompletionBridgeHandler._contains_stream_required_text( + value[key] + ): + return True + return any( + ResponsesToCompletionBridgeHandler._contains_stream_required_text(v) + for v in value.values() + ) + if isinstance(value, list): + return any( + ResponsesToCompletionBridgeHandler._contains_stream_required_text(v) + for v in value + ) + return False + + @classmethod + def _is_stream_required_error(cls, e: Exception) -> bool: + for attr in ("body", "message", "text"): + if cls._contains_stream_required_text(getattr(e, attr, None)): + return True + response = getattr(e, "response", None) + if response is not None: + try: + if cls._contains_stream_required_text(response.text): + return True + except Exception: + return False + return cls._contains_stream_required_text(str(e)) + @staticmethod def _coerce_response_object( response_obj: Any, @@ -165,6 +217,7 @@ class ResponsesToCompletionBridgeHandler: logging_obj = validated_kwargs["logging_obj"] custom_llm_provider = validated_kwargs["custom_llm_provider"] + stream = self._resolve_stream_flag(optional_params, litellm_params) request_data = self.transformation_handler.transform_request( model=model, messages=messages, @@ -175,11 +228,21 @@ class ResponsesToCompletionBridgeHandler: client=kwargs.get("client"), ) - result = responses( - **request_data, - ) + try: + result = responses( + **request_data, + ) + except Exception as e: + if not stream and self._is_stream_required_error(e): + if hasattr(logging_obj, "model_call_details"): + logging_obj.model_call_details["forced_streaming_fallback"] = True + request_data = {**request_data, "stream": True} + result = responses( + **request_data, + ) + else: + raise - stream = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( model=model, @@ -239,6 +302,7 @@ class ResponsesToCompletionBridgeHandler: logging_obj = validated_kwargs["logging_obj"] custom_llm_provider = validated_kwargs["custom_llm_provider"] + stream = self._resolve_stream_flag(optional_params, litellm_params) try: request_data = self.transformation_handler.transform_request( model=model, @@ -251,12 +315,23 @@ class ResponsesToCompletionBridgeHandler: except Exception as e: raise e - result = await aresponses( - **request_data, - aresponses=True, - ) + try: + result = await aresponses( + **request_data, + aresponses=True, + ) + except Exception as e: + if not stream and self._is_stream_required_error(e): + if hasattr(logging_obj, "model_call_details"): + logging_obj.model_call_details["forced_streaming_fallback"] = True + request_data = {**request_data, "stream": True} + result = await aresponses( + **request_data, + aresponses=True, + ) + else: + raise - stream = self._resolve_stream_flag(optional_params, litellm_params) if isinstance(result, ResponsesAPIResponse): return self.transformation_handler.transform_response( model=model, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7267532933d..967b9f0d0f1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -151,6 +151,104 @@ else: class BaseLLMHTTPHandler: + @staticmethod + def _contains_stream_required_text(value: Any) -> bool: + if value is None: + return False + if isinstance(value, str): + if "stream must be set to true" in value.lower(): + return True + try: + parsed = json.loads(value) + except Exception: + return False + return BaseLLMHTTPHandler._contains_stream_required_text(parsed) + if isinstance(value, dict): + for key in ("detail", "message", "error"): + if key in value and BaseLLMHTTPHandler._contains_stream_required_text( + value[key] + ): + return True + return any( + BaseLLMHTTPHandler._contains_stream_required_text(v) + for v in value.values() + ) + if isinstance(value, list): + return any( + BaseLLMHTTPHandler._contains_stream_required_text(v) for v in value + ) + return False + + @classmethod + def _is_stream_required_error(cls, e: Exception) -> bool: + for attr in ("body", "message", "text"): + if cls._contains_stream_required_text(getattr(e, attr, None)): + return True + response = getattr(e, "response", None) + if response is not None: + try: + return cls._contains_stream_required_text(response.text) + except Exception: + return False + return False + + @staticmethod + def _merge_stream_hidden_params( + response: ModelResponse, streamwrapper: CustomStreamWrapper + ) -> None: + hidden = getattr(streamwrapper, "_hidden_params", None) + if not isinstance(hidden, dict): + return + response_hidden = getattr(response, "_hidden_params", None) + if response_hidden is None: + response._hidden_params = {} + response_hidden = response._hidden_params + response_hidden.update(hidden) + + def _build_complete_response_from_streaming( + self, + streamwrapper: CustomStreamWrapper, + messages: Optional[list], + provider_config: BaseConfig, + ) -> ModelResponse: + chunks = [] + for chunk in streamwrapper: + chunks.append(chunk) + complete_response = litellm.stream_chunk_builder( + chunks=chunks, messages=messages + ) + if complete_response is None: + raise provider_config.get_error_class( + error_message="Failed to assemble streaming response for forced stream.", + status_code=500, + headers={}, + ) + complete_response = cast(ModelResponse, complete_response) + self._merge_stream_hidden_params(complete_response, streamwrapper) + return complete_response + + async def _abuild_complete_response_from_streaming( + self, + streamwrapper: CustomStreamWrapper, + messages: Optional[list], + provider_config: BaseConfig, + ) -> ModelResponse: + chunks = [] + async for chunk in streamwrapper: + chunks.append(chunk) + complete_response = litellm.stream_chunk_builder( + chunks=chunks, messages=messages + ) + if complete_response is None: + raise provider_config.get_error_class( + error_message="Failed to assemble streaming response for forced stream.", + status_code=500, + headers={}, + ) + complete_response = cast(ModelResponse, complete_response) + self._merge_stream_hidden_params(complete_response, streamwrapper) + return complete_response + async def _make_common_async_call( self, async_httpx_client: AsyncHTTPHandler, @@ -304,18 +402,78 @@ class BaseLLMHTTPHandler: else: async_httpx_client = client - response = await self._make_common_async_call( - async_httpx_client=async_httpx_client, - provider_config=provider_config, - api_base=api_base, - headers=headers, - data=data, - timeout=timeout, - litellm_params=litellm_params, - stream=False, - logging_obj=logging_obj, - signed_json_body=signed_json_body, - ) + try: + response = await self._make_common_async_call( + async_httpx_client=async_httpx_client, + provider_config=provider_config, + api_base=api_base, + headers=headers, + data=data, + timeout=timeout, + litellm_params=litellm_params, + stream=False, + logging_obj=logging_obj, + signed_json_body=signed_json_body, + ) + except Exception as e: + if self._is_stream_required_error(e) and not provider_config.has_custom_stream_wrapper: + logging_obj.model_call_details["forced_streaming_fallback"] = True + stream_data = self._add_stream_param_to_request_body( + data=data.copy(), + provider_config=provider_config, + fake_stream=False, + ) + forced_headers, forced_signed_json_body = provider_config.sign_request( + headers=headers.copy(), + optional_params=optional_params, + request_data=stream_data, + api_base=api_base, + api_key=api_key, + stream=True, + fake_stream=False, + model=model, + ) + completion_stream, response_headers = await self.make_async_call_stream_helper( + model=model, + custom_llm_provider=custom_llm_provider, + provider_config=provider_config, + api_base=api_base, + headers=forced_headers, + data=stream_data, + messages=messages, + logging_obj=logging_obj, + timeout=timeout, + fake_stream=False, + client=client, + litellm_params=litellm_params, + optional_params=optional_params, + json_mode=json_mode, + signed_json_body=forced_signed_json_body, + ) + streamwrapper = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider=custom_llm_provider, + logging_obj=logging_obj, + _response_headers=dict(response_headers), + ) + complete_response = await self._abuild_complete_response_from_streaming( + streamwrapper=streamwrapper, + messages=messages, + provider_config=provider_config, + ) + agentic_response = await self._call_agentic_chat_completion_hooks( + response=complete_response, + model=model, + messages=messages, + optional_params=optional_params, + logging_obj=logging_obj, + stream=False, + custom_llm_provider=custom_llm_provider, + kwargs=litellm_params, + ) + return agentic_response if agentic_response is not None else complete_response + raise initial_response = provider_config.transform_response( model=model, raw_response=response, @@ -554,17 +712,70 @@ class BaseLLMHTTPHandler: else: sync_httpx_client = client - response = self._make_common_sync_call( - sync_httpx_client=sync_httpx_client, - provider_config=provider_config, - api_base=api_base, - headers=headers, - data=data, - signed_json_body=signed_json_body, - timeout=timeout, - litellm_params=litellm_params, - logging_obj=logging_obj, - ) + try: + response = self._make_common_sync_call( + sync_httpx_client=sync_httpx_client, + provider_config=provider_config, + api_base=api_base, + headers=headers, + data=data, + signed_json_body=signed_json_body, + timeout=timeout, + litellm_params=litellm_params, + logging_obj=logging_obj, + ) + except Exception as e: + if self._is_stream_required_error(e) and not provider_config.has_custom_stream_wrapper: + logging_obj.model_call_details["forced_streaming_fallback"] = True + stream_data = self._add_stream_param_to_request_body( + data=data.copy(), + provider_config=provider_config, + fake_stream=False, + ) + forced_headers, forced_signed_json_body = provider_config.sign_request( + headers=headers.copy(), + optional_params=optional_params, + request_data=stream_data, + api_base=api_base, + api_key=api_key, + stream=True, + fake_stream=False, + model=model, + ) + completion_stream, response_headers = self.make_sync_call( + provider_config=provider_config, + api_base=api_base, + headers=forced_headers, + data=stream_data, + signed_json_body=forced_signed_json_body, + original_data=stream_data, + model=model, + messages=messages, + logging_obj=logging_obj, + timeout=timeout, + fake_stream=False, + client=( + client + if client is not None and isinstance(client, HTTPHandler) + else None + ), + litellm_params=litellm_params, + json_mode=json_mode, + optional_params=optional_params, + ) + streamwrapper = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider=custom_llm_provider, + logging_obj=logging_obj, + _response_headers=dict(response_headers), + ) + return self._build_complete_response_from_streaming( + streamwrapper=streamwrapper, + messages=messages, + provider_config=provider_config, + ) + raise return provider_config.transform_response( model=model, raw_response=response, diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index c7524925bd0..e30d464ae97 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1,3 +1,5 @@ +import asyncio +import json import time import types from typing import ( @@ -339,6 +341,68 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): def __init__(self) -> None: super().__init__() + @staticmethod + def _is_stream_required_error(e: Exception) -> bool: + message = getattr(e, "message", None) or getattr(e, "text", None) or str(e) + if isinstance(message, dict): + message = json.dumps(message) + return "stream must be set to true" in str(message).lower() + + @staticmethod + def _merge_stream_hidden_params( + response: ModelResponse, streamwrapper: CustomStreamWrapper + ) -> None: + hidden = getattr(streamwrapper, "_hidden_params", None) + if not isinstance(hidden, dict): + return + response_hidden = getattr(response, "_hidden_params", None) + if response_hidden is None: + response._hidden_params = {} + response_hidden = response._hidden_params + response_hidden.update(hidden) + + def _build_complete_response_from_streaming( + self, streamwrapper: CustomStreamWrapper, messages: Optional[list] + ) -> ModelResponse: + chunks: List[ModelResponseStream] = [] + for chunk in streamwrapper: + chunks.append(chunk) + complete_response = litellm.stream_chunk_builder( + chunks=chunks, messages=messages + ) + if complete_response is None: + raise OpenAIError( + status_code=500, + message="Failed to assemble streaming response for forced stream.", + ) + complete_response = cast(ModelResponse, complete_response) + self._merge_stream_hidden_params(complete_response, streamwrapper) + return complete_response + + async def _abuild_complete_response_from_streaming( + self, + streamwrapper_or_coro: Union[CustomStreamWrapper, Coroutine], + messages: Optional[list], + ) -> ModelResponse: + if asyncio.iscoroutine(streamwrapper_or_coro): + streamwrapper = await streamwrapper_or_coro + else: + streamwrapper = streamwrapper_or_coro + chunks: List[ModelResponseStream] = [] + async for chunk in streamwrapper: + chunks.append(chunk) + complete_response = litellm.stream_chunk_builder( + chunks=chunks, messages=messages + ) + if complete_response is None: + raise OpenAIError( + status_code=500, + message="Failed to assemble streaming response for forced stream.", + ) + complete_response = cast(ModelResponse, complete_response) + self._merge_stream_hidden_params(complete_response, streamwrapper) + return complete_response + def _set_dynamic_params_on_client( self, client: Union[OpenAI, AsyncOpenAI], @@ -635,6 +699,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) stream: Optional[bool] = inference_params.pop("stream", False) provider_config: Optional[BaseConfig] = None + return_complete_response: bool = False if custom_llm_provider is not None and model is not None: try: @@ -652,7 +717,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): fake_stream = provider_config.should_fake_stream( model=model, custom_llm_provider=custom_llm_provider, stream=stream ) - if headers: inference_params["extra_headers"] = headers if model is None or messages is None: @@ -676,7 +740,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries = inference_params.pop("max_retries", 2) if acompletion is True: if stream is True and fake_stream is False: - return self.async_streaming( + streaming_response = self.async_streaming( logging_obj=logging_obj, headers=headers, messages=messages, @@ -695,6 +759,28 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): stream_options=stream_options, shared_session=shared_session, ) + if return_complete_response: + async def _finalize_forced_streaming(): + complete_response = await self._abuild_complete_response_from_streaming( + streaming_response, messages + ) + agentic_response = await self._call_agentic_completion_hooks_openai( + response=complete_response, + model=model, + messages=messages, + optional_params=inference_params, + logging_obj=logging_obj, + stream=False, + litellm_params=litellm_params, + ) + return ( + agentic_response + if agentic_response is not None + else complete_response + ) + + return _finalize_forced_streaming() + return streaming_response else: return self.acompletion( messages=messages, @@ -725,7 +811,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): headers=headers or {}, ) if stream is True and fake_stream is False: - return self.streaming( + streaming_response = self.streaming( logging_obj=logging_obj, headers=headers, data=data, @@ -739,6 +825,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, stream_options=stream_options, ) + if return_complete_response: + return self._build_complete_response_from_streaming( + streaming_response, messages + ) + return streaming_response else: if not isinstance(max_retries, int): raise OpenAIError( @@ -841,6 +932,14 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): e ): litellm.remove_index_from_tool_calls(messages=messages) + elif ( + stream is False + and return_complete_response is False + and self._is_stream_required_error(e) + ): + stream = True + return_complete_response = True + continue else: raise e except OpenAIError as e: