Handle stream-required fallback

This commit is contained in:
wanna 2026-02-24 21:46:08 +08:00
parent f55fe7afdc
commit 6db7fc319c
3 changed files with 420 additions and 35 deletions

View file

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

View file

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

View file

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