From 3d8bcbd5d6faf5d19386ab50aea6de7e0a743d68 Mon Sep 17 00:00:00 2001 From: wanna Date: Tue, 24 Feb 2026 22:14:03 +0800 Subject: [PATCH] Deduplicate stream-required error detection --- .../handler.py | 57 +------------------ litellm/litellm_core_utils/error_utils.py | 47 +++++++++++++++ litellm/llms/custom_httpx/llm_http_handler.py | 52 +++-------------- litellm/llms/openai/openai.py | 10 +--- 4 files changed, 61 insertions(+), 105 deletions(-) create mode 100644 litellm/litellm_core_utils/error_utils.py diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 3087b9840d8..38e54a717c8 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -2,12 +2,12 @@ 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 from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.litellm_core_utils.error_utils import is_stream_required_error if TYPE_CHECKING: from litellm import CustomStreamWrapper, LiteLLMLoggingObj, ModelResponse @@ -38,57 +38,6 @@ 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, @@ -233,7 +182,7 @@ class ResponsesToCompletionBridgeHandler: **request_data, ) except Exception as e: - if not stream and self._is_stream_required_error(e): + if not stream and 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} @@ -321,7 +270,7 @@ class ResponsesToCompletionBridgeHandler: aresponses=True, ) except Exception as e: - if not stream and self._is_stream_required_error(e): + if not stream and 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} diff --git a/litellm/litellm_core_utils/error_utils.py b/litellm/litellm_core_utils/error_utils.py new file mode 100644 index 00000000000..4f442524d8a --- /dev/null +++ b/litellm/litellm_core_utils/error_utils.py @@ -0,0 +1,47 @@ +"""Shared error-detection helpers.""" + +import json +from typing import Any + +_STREAM_REQUIRED_TEXT = "stream must be set to true" + + +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_REQUIRED_TEXT in lowered: + return True + try: + parsed = json.loads(value) + except Exception: + return False + return _contains_stream_required_text(parsed) + if isinstance(value, dict): + for key in ("detail", "message", "error"): + if key in value and _contains_stream_required_text(value[key]): + return True + return any(_contains_stream_required_text(v) for v in value.values()) + if isinstance(value, list): + return any(_contains_stream_required_text(v) for v in value) + return False + + +def is_stream_required_error(err: Exception) -> bool: + for attr in ("body", "message", "text"): + if _contains_stream_required_text(getattr(err, attr, None)): + return True + response = getattr(err, "response", None) + if response is not None: + try: + if _contains_stream_required_text(getattr(response, "text", None)): + return True + except Exception: + return False + return _contains_stream_required_text(str(err)) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 967b9f0d0f1..4a16c7a4049 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -25,6 +25,7 @@ from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, ) from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.litellm_core_utils.error_utils import is_stream_required_error from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, @@ -151,47 +152,6 @@ 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 @@ -416,7 +376,10 @@ class BaseLLMHTTPHandler: 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: + if ( + 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(), @@ -725,7 +688,10 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, ) except Exception as e: - if self._is_stream_required_error(e) and not provider_config.has_custom_stream_wrapper: + if ( + 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(), diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index e30d464ae97..de2293546aa 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -34,6 +34,7 @@ import litellm from litellm import LlmProviders from litellm._logging import verbose_logger from litellm.constants import DEFAULT_MAX_RETRIES +from litellm.litellm_core_utils.error_utils import is_stream_required_error from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_utils import track_llm_api_timing from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -341,13 +342,6 @@ 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 @@ -935,7 +929,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): elif ( stream is False and return_complete_response is False - and self._is_stream_required_error(e) + and is_stream_required_error(e) ): stream = True return_complete_response = True