mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(proxy): enhance fallback handling in streaming mode
This commit is contained in:
parent
ec2f3aadb8
commit
8665cddf44
5 changed files with 86 additions and 2 deletions
|
|
@ -1255,6 +1255,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# streaming generators. Pre-call processing can rewrite `self.data["model"]` for
|
||||
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
|
||||
# what the client sent.
|
||||
attempted_fallbacks = additional_headers.get(
|
||||
"x-litellm-attempted-fallbacks", 0
|
||||
)
|
||||
try:
|
||||
attempted_fallbacks = int(attempted_fallbacks or 0)
|
||||
except (TypeError, ValueError):
|
||||
attempted_fallbacks = 0
|
||||
if attempted_fallbacks > 0:
|
||||
self.data["_litellm_attempted_fallbacks"] = attempted_fallbacks
|
||||
|
||||
if requested_model_from_client:
|
||||
self.data["_litellm_client_requested_model"] = (
|
||||
requested_model_from_client
|
||||
|
|
|
|||
|
|
@ -6709,6 +6709,22 @@ def _restamp_streaming_chunk_model(
|
|||
if request_data.get("fastest_response", False):
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
attempted_fallbacks = request_data.get("_litellm_attempted_fallbacks", 0)
|
||||
if attempted_fallbacks is None:
|
||||
attempted_fallbacks = 0
|
||||
hidden_params = getattr(chunk, "_hidden_params", {}) or {}
|
||||
if isinstance(hidden_params, dict):
|
||||
additional_headers = hidden_params.get("additional_headers", {}) or {}
|
||||
attempted_fallbacks = additional_headers.get(
|
||||
"x-litellm-attempted-fallbacks", attempted_fallbacks
|
||||
)
|
||||
try:
|
||||
attempted_fallbacks = int(attempted_fallbacks or 0)
|
||||
except (TypeError, ValueError):
|
||||
attempted_fallbacks = 0
|
||||
if attempted_fallbacks > 0:
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
downstream_model = (
|
||||
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2062,6 +2062,35 @@ class Router:
|
|||
)
|
||||
setattr(fallback_item, "usage", combined_usage)
|
||||
|
||||
@staticmethod
|
||||
def _copy_stream_wrapper_additional_headers_to_chunk(
|
||||
stream_response: Any,
|
||||
stream_chunk: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
Copy wrapper-level headers onto streamed chunks.
|
||||
|
||||
Fallback headers are attached to the stream wrapper after a successful
|
||||
fallback. The proxy only sees individual chunks during SSE serialization,
|
||||
so the chunk needs the same hidden header metadata.
|
||||
"""
|
||||
response_hidden_params = getattr(stream_response, "_hidden_params", {}) or {}
|
||||
if not isinstance(response_hidden_params, dict):
|
||||
return stream_chunk
|
||||
|
||||
response_headers = response_hidden_params.get("additional_headers", {}) or {}
|
||||
if not response_headers or not hasattr(stream_chunk, "_hidden_params"):
|
||||
return stream_chunk
|
||||
|
||||
chunk_hidden_params = getattr(stream_chunk, "_hidden_params", {}) or {}
|
||||
if not isinstance(chunk_hidden_params, dict):
|
||||
chunk_hidden_params = {}
|
||||
|
||||
chunk_hidden_params.setdefault("additional_headers", {})
|
||||
chunk_hidden_params["additional_headers"].update(response_headers)
|
||||
setattr(stream_chunk, "_hidden_params", chunk_hidden_params)
|
||||
return stream_chunk
|
||||
|
||||
async def _acompletion_streaming_iterator(
|
||||
self,
|
||||
model_response: CustomStreamWrapper,
|
||||
|
|
@ -2160,6 +2189,12 @@ class Router:
|
|||
# If fallback returns a streaming response, iterate over it
|
||||
if hasattr(fallback_response, "__aiter__"):
|
||||
async for fallback_item in fallback_response: # type: ignore
|
||||
fallback_item = (
|
||||
self._copy_stream_wrapper_additional_headers_to_chunk(
|
||||
stream_response=fallback_response,
|
||||
stream_chunk=fallback_item,
|
||||
)
|
||||
)
|
||||
if (
|
||||
fallback_item
|
||||
and isinstance(fallback_item, ModelResponseStream)
|
||||
|
|
@ -2296,6 +2331,10 @@ class Router:
|
|||
|
||||
if hasattr(fallback_response, "__iter__"):
|
||||
for fallback_item in fallback_response:
|
||||
fallback_item = router_self._copy_stream_wrapper_additional_headers_to_chunk(
|
||||
stream_response=fallback_response,
|
||||
stream_chunk=fallback_item,
|
||||
)
|
||||
if (
|
||||
fallback_item
|
||||
and isinstance(fallback_item, ModelResponseStream)
|
||||
|
|
@ -8844,7 +8883,7 @@ class Router:
|
|||
# - else return the model's rate limit headers
|
||||
"""
|
||||
if (
|
||||
isinstance(response, BaseModel)
|
||||
(isinstance(response, BaseModel) or hasattr(response, "_hidden_params"))
|
||||
and hasattr(response, "_hidden_params")
|
||||
and isinstance(response._hidden_params, dict) # type: ignore
|
||||
):
|
||||
|
|
|
|||
|
|
@ -9,7 +9,9 @@ def _add_headers_to_response(response: Any, headers: dict) -> Any:
|
|||
"""
|
||||
Helper function to add headers to a response's hidden params
|
||||
"""
|
||||
if response is None or not isinstance(response, BaseModel):
|
||||
if response is None or not (
|
||||
isinstance(response, BaseModel) or hasattr(response, "_hidden_params")
|
||||
):
|
||||
return response
|
||||
|
||||
hidden_params: Optional[Union[dict, HiddenParams]] = getattr(
|
||||
|
|
|
|||
|
|
@ -86,6 +86,23 @@ def test_restamp_streaming_chunk_skips_matching_model():
|
|||
assert result.model == "client-model"
|
||||
assert model_mismatch_logged is False
|
||||
|
||||
def test_restamp_streaming_chunk_preserves_fallback_model():
|
||||
from litellm.proxy.proxy_server import _restamp_streaming_chunk_model
|
||||
|
||||
chunk = _make_model_response_stream_chunk("fallback-model")
|
||||
chunk._hidden_params.setdefault("additional_headers", {})
|
||||
chunk._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] = 1
|
||||
|
||||
result, model_mismatch_logged = _restamp_streaming_chunk_model(
|
||||
chunk=chunk,
|
||||
requested_model_from_client="client-model",
|
||||
request_data={"litellm_call_id": "test-call-id"},
|
||||
model_mismatch_logged=False,
|
||||
)
|
||||
|
||||
assert result is chunk
|
||||
assert result.model == "fallback-model"
|
||||
assert model_mismatch_logged is False
|
||||
|
||||
def test_fast_serialize_simple_streaming_chunk_matches_model_dump_json():
|
||||
from litellm.proxy.proxy_server import _serialize_streaming_chunk
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue