mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Deduplicate stream-required error detection
This commit is contained in:
parent
6db7fc319c
commit
3d8bcbd5d6
4 changed files with 61 additions and 105 deletions
|
|
@ -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}
|
||||
|
|
|
|||
47
litellm/litellm_core_utils/error_utils.py
Normal file
47
litellm/litellm_core_utils/error_utils.py
Normal 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))
|
||||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue