mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Handle stream-required fallback
This commit is contained in:
parent
f55fe7afdc
commit
6db7fc319c
3 changed files with 420 additions and 35 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue