Deduplicate stream-required error detection

This commit is contained in:
wanna 2026-02-24 22:14:03 +08:00
parent 6db7fc319c
commit 3d8bcbd5d6
4 changed files with 61 additions and 105 deletions

View file

@ -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}

View file

@ -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))

View file

@ -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(),

View file

@ -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