mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge pull request #42283 from BerriAI/litellm_mid_stream_fallback_walks_full_list
fix(router): walk every entry of a fallback list after a mid-stream failure
This commit is contained in:
commit
3d26a29a1a
5 changed files with 570 additions and 112 deletions
|
|
@ -184,13 +184,17 @@ from litellm.router_utils.cooldown_handlers import (
|
|||
is_caller_timeout_408,
|
||||
)
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
MID_STREAM_FALLBACK_CONTROLS_KEY,
|
||||
AttemptedFallbackTargets,
|
||||
_check_non_standard_fallback_format,
|
||||
carry_over_pre_routing_selection,
|
||||
clear_pre_routing_selection,
|
||||
fallback_lookup_groups,
|
||||
fallbacks_disabled_for_request,
|
||||
get_fallback_model_group_for_lookup_groups,
|
||||
get_pre_routing_selection,
|
||||
has_unattempted_fallback_target,
|
||||
mid_stream_fallback_hop_kwargs,
|
||||
per_request_fallback_controls,
|
||||
record_disable_fallbacks,
|
||||
record_pre_routing_selection,
|
||||
run_async_fallback,
|
||||
|
|
@ -3304,12 +3308,7 @@ class Router:
|
|||
content_policy_fallbacks: Final[list | None] = initial_kwargs.get(
|
||||
"content_policy_fallbacks", self.content_policy_fallbacks
|
||||
)
|
||||
# Re-enter via the per-attempt helper so the fallback chain
|
||||
# picks deployments through
|
||||
# _ageneric_api_call_with_fallbacks_helper.
|
||||
# original_generic_function is preserved by the caller so
|
||||
# the helper knows what underlying API to invoke per attempt.
|
||||
initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper
|
||||
initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_responses_attempt
|
||||
if e.is_pre_first_chunk or not e.generated_content:
|
||||
# No content generated before the error — retry with the
|
||||
# original input. Adding a continuation prompt would
|
||||
|
|
@ -5140,22 +5139,28 @@ class Router:
|
|||
request_kwargs=None,
|
||||
)
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs):
|
||||
async def _ageneric_api_call_with_fallbacks(
|
||||
self, model: str, original_function: Callable, attempt_function: Callable | None = None, **kwargs
|
||||
):
|
||||
"""
|
||||
Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
|
||||
|
||||
attempt_function runs every attempt of the chain instead of the plain helper, so a streaming
|
||||
endpoint can wrap each attempt's stream with its own mid-stream fallback handling.
|
||||
"""
|
||||
try:
|
||||
kwargs["model"] = model
|
||||
kwargs["original_generic_function"] = original_function
|
||||
kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper
|
||||
kwargs["original_function"] = attempt_function or self._ageneric_api_call_with_fallbacks_helper
|
||||
if attempt_function is not None:
|
||||
controls: Final = per_request_fallback_controls(kwargs)
|
||||
kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop
|
||||
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs, metadata_variable_name="litellm_metadata")
|
||||
verbose_router_logger.debug(
|
||||
"Inside ageneric_api_call_with_fallbacks() - model: %s; kwargs: %s", model, kwargs
|
||||
)
|
||||
response: Final = await self.async_function_with_fallbacks(**kwargs)
|
||||
return response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
asyncio.create_task(
|
||||
send_llm_exception_alert(
|
||||
|
|
@ -5276,61 +5281,42 @@ class Router:
|
|||
self, original_function: Callable, **kwargs: Any
|
||||
) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]:
|
||||
"""
|
||||
_ageneric_api_call_with_fallbacks for the Responses API, with the
|
||||
addition of mid-stream fallback handling.
|
||||
|
||||
When stream=True and the underlying call returns a
|
||||
BaseResponsesAPIStreamingIterator, wrap it with
|
||||
_aresponses_streaming_iterator so MidStreamFallbackError raised
|
||||
during iteration triggers the Router's cross-provider fallback chain.
|
||||
_ageneric_api_call_with_fallbacks for the Responses API, with every attempt's stream
|
||||
carrying its own mid-stream fallback handling
|
||||
(see _ageneric_api_call_with_fallbacks_responses_attempt).
|
||||
"""
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
original_function=original_function,
|
||||
attempt_function=self._ageneric_api_call_with_fallbacks_responses_attempt,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks_responses_attempt(
|
||||
self,
|
||||
model: str,
|
||||
original_generic_function: Callable,
|
||||
**kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site
|
||||
) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]:
|
||||
"""
|
||||
One attempt of the Responses API fallback chain. A streaming result is wrapped with
|
||||
_aresponses_streaming_iterator over this attempt's own kwargs, so a fallback hop that
|
||||
fails mid-stream resumes the original group's chain instead of re-raising; the name keeps
|
||||
_get_router_metadata_variable_name resolving to litellm_metadata for every hop.
|
||||
"""
|
||||
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
|
||||
from litellm.responses.streaming_iterator import (
|
||||
BaseResponsesAPIStreamingIterator,
|
||||
)
|
||||
|
||||
# Snapshot the request kwargs before _ageneric_api_call_with_fallbacks
|
||||
# mutates them. A shallow copy alone is not enough: the primary
|
||||
# attempt mutates nested dicts in place — notably `litellm_metadata`,
|
||||
# which `_update_kwargs_with_deployment` populates with
|
||||
# deployment-specific fields (`deployment`, `model_info`, `api_base`,
|
||||
# tags, etc.). Without an explicit copy of that dict, the shallow
|
||||
# copy would still share its reference, leaking primary-deployment
|
||||
# metadata into the mid-stream fallback request.
|
||||
#
|
||||
# We avoid deep-copying the full kwargs because it can contain
|
||||
# non-deepcopyable objects (logging handles, async clients, etc.);
|
||||
# `safe_deep_copy` deep-copies the metadata dicts key-by-key with a
|
||||
# fallback to the original reference for any non-picklable value.
|
||||
# The original_generic_function is preserved so the per-attempt
|
||||
# helper knows which underlying API to call on fallback.
|
||||
# The pre-routing hook stamps its tier selection into this bucket during the primary
|
||||
# attempt; seeding it before the snapshot gives both the live kwargs and the copy a
|
||||
# bucket, so the post-call carry-over below always has somewhere to read and write.
|
||||
kwargs.setdefault("litellm_metadata", {}) # mutable-ok: shared bucket # rebind-ok: stamp must be readable here
|
||||
|
||||
fallback_kwargs: Final[dict[str, object]] = kwargs.copy()
|
||||
if isinstance(fallback_kwargs.get("litellm_metadata"), dict):
|
||||
fallback_kwargs["litellm_metadata"] = safe_deep_copy(fallback_kwargs["litellm_metadata"])
|
||||
if isinstance(fallback_kwargs.get("metadata"), dict):
|
||||
fallback_kwargs["metadata"] = safe_deep_copy(fallback_kwargs["metadata"])
|
||||
fallback_kwargs["original_generic_function"] = original_function
|
||||
|
||||
response: Final = await self._ageneric_api_call_with_fallbacks(original_function=original_function, **kwargs)
|
||||
|
||||
# The snapshot predates the pre-routing hook, so the tier it stamped into the live kwargs
|
||||
# is carried over write-or-clear: a stale or caller-supplied selection left in the copy
|
||||
# would key the mid-stream fallback lookup off a tier this attempt never routed to.
|
||||
clear_pre_routing_selection(fallback_kwargs)
|
||||
live_pre_routing_selection: Final = get_pre_routing_selection(kwargs)
|
||||
if live_pre_routing_selection is not None:
|
||||
record_pre_routing_selection(fallback_kwargs, live_pre_routing_selection)
|
||||
|
||||
controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None)
|
||||
hop_kwargs: Final = mid_stream_fallback_hop_kwargs(
|
||||
model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs
|
||||
)
|
||||
response: Final = await self._ageneric_api_call_with_fallbacks_helper(
|
||||
model=model, original_generic_function=original_generic_function, **kwargs
|
||||
)
|
||||
carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs)
|
||||
if kwargs.get("stream") and isinstance(response, BaseResponsesAPIStreamingIterator):
|
||||
return await self._aresponses_streaming_iterator(
|
||||
response=response,
|
||||
initial_kwargs=fallback_kwargs,
|
||||
)
|
||||
return await self._aresponses_streaming_iterator(response=response, initial_kwargs=hop_kwargs)
|
||||
return response
|
||||
|
||||
async def _aanthropic_messages_streaming_iterator(
|
||||
|
|
@ -5559,7 +5545,7 @@ class Router:
|
|||
content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below
|
||||
"content_policy_fallbacks", self.content_policy_fallbacks
|
||||
)
|
||||
initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper
|
||||
initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt
|
||||
self._update_kwargs_before_fallbacks(
|
||||
model=model_group,
|
||||
kwargs=initial_kwargs,
|
||||
|
|
@ -5613,46 +5599,41 @@ class Router:
|
|||
**kwargs: object, # kwargs-ok: forwarded verbatim to original_function, shape varies per call site
|
||||
) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]:
|
||||
"""
|
||||
_ageneric_api_call_with_fallbacks for anthropic_messages, with the
|
||||
addition of mid-stream fallback handling (see
|
||||
_aanthropic_messages_streaming_iterator). Parity with
|
||||
_ageneric_api_call_with_fallbacks for anthropic_messages, with every attempt's stream
|
||||
carrying its own mid-stream fallback handling
|
||||
(see _ageneric_api_call_with_fallbacks_anthropic_messages_attempt). Parity with
|
||||
_aresponses_with_streaming_fallbacks for the Responses API.
|
||||
"""
|
||||
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
|
||||
|
||||
# Snapshot the request kwargs before the primary attempt mutates them
|
||||
# in place: _update_kwargs_with_deployment writes deployment-specific
|
||||
# fields (deployment, model_info, api_base, tags, ...) into the
|
||||
# SAME litellm_metadata/metadata dicts a shallow .copy() would still
|
||||
# share, leaking primary-deployment metadata into the mid-stream
|
||||
# fallback request. safe_deep_copy avoids deep-copying the full
|
||||
# kwargs (which can hold non-deepcopyable logging handles/clients).
|
||||
# The pre-routing hook stamps its tier selection into this bucket during the primary
|
||||
# attempt; seeding it before the snapshot gives both the live kwargs and the copy a
|
||||
# bucket, so the post-call carry-over below always has somewhere to read and write.
|
||||
kwargs.setdefault("litellm_metadata", {}) # mutable-ok: shared bucket # rebind-ok: stamp must be readable here
|
||||
|
||||
fallback_kwargs: Final[dict[str, object]] = kwargs.copy() # mutable-ok: mutated below before re-entry
|
||||
if isinstance(fallback_kwargs.get("litellm_metadata"), dict):
|
||||
fallback_kwargs["litellm_metadata"] = safe_deep_copy(fallback_kwargs["litellm_metadata"])
|
||||
if isinstance(fallback_kwargs.get("metadata"), dict):
|
||||
fallback_kwargs["metadata"] = safe_deep_copy(fallback_kwargs["metadata"])
|
||||
fallback_kwargs["original_generic_function"] = original_function
|
||||
|
||||
response: Final = await self._ageneric_api_call_with_fallbacks(original_function=original_function, **kwargs)
|
||||
|
||||
# The snapshot predates the pre-routing hook, so the tier it stamped into the live kwargs
|
||||
# is carried over write-or-clear: a stale or caller-supplied selection left in the copy
|
||||
# would key the mid-stream fallback lookup off a tier this attempt never routed to.
|
||||
clear_pre_routing_selection(fallback_kwargs)
|
||||
live_pre_routing_selection: Final = get_pre_routing_selection(kwargs)
|
||||
if live_pre_routing_selection is not None:
|
||||
record_pre_routing_selection(fallback_kwargs, live_pre_routing_selection)
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
original_function=original_function,
|
||||
attempt_function=self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks_anthropic_messages_attempt(
|
||||
self,
|
||||
model: str,
|
||||
original_generic_function: Callable,
|
||||
**kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site
|
||||
) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]:
|
||||
"""
|
||||
One attempt of the anthropic_messages fallback chain. A streaming result is wrapped with
|
||||
_aanthropic_messages_streaming_iterator over this attempt's own kwargs, so a fallback hop
|
||||
that fails mid-stream resumes the original group's chain instead of re-raising; the name
|
||||
keeps _get_router_metadata_variable_name resolving to litellm_metadata for every hop.
|
||||
"""
|
||||
controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None)
|
||||
hop_kwargs: Final = mid_stream_fallback_hop_kwargs(
|
||||
model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs
|
||||
)
|
||||
response: Final = await self._ageneric_api_call_with_fallbacks_helper(
|
||||
model=model, original_generic_function=original_generic_function, **kwargs
|
||||
)
|
||||
carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs)
|
||||
if kwargs.get("stream") and hasattr(response, "__aiter__"):
|
||||
return await self._aanthropic_messages_streaming_iterator(
|
||||
response=cast("AsyncIterator[bytes]", response), # cast-ok: stream=True always returns a byte iterator
|
||||
initial_kwargs=fallback_kwargs,
|
||||
initial_kwargs=hop_kwargs,
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -8338,12 +8319,12 @@ class Router:
|
|||
"""
|
||||
content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks)
|
||||
if content_policy_fallbacks is not None:
|
||||
return (
|
||||
return has_unattempted_fallback_target(
|
||||
self._get_fallback_model_group_for_lookup_groups(
|
||||
fallbacks=content_policy_fallbacks,
|
||||
lookup_groups=fallback_lookup_groups(kwargs, model_group),
|
||||
)
|
||||
is not None
|
||||
),
|
||||
kwargs,
|
||||
)
|
||||
if self._has_default_fallbacks():
|
||||
return True
|
||||
|
|
@ -8375,7 +8356,7 @@ class Router:
|
|||
fallbacks=fallbacks,
|
||||
lookup_groups=fallback_lookup_groups(kwargs, model_group),
|
||||
)
|
||||
return resolved is not None
|
||||
return has_unattempted_fallback_target(resolved, kwargs)
|
||||
|
||||
def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
|
|
@ -10,7 +10,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
|
||||
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs, safe_deep_copy
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_structure
|
||||
from litellm.router_utils.add_retry_fallback_headers import (
|
||||
add_fallback_headers_to_response,
|
||||
|
|
@ -199,6 +199,18 @@ class AttemptedFallbackTargets:
|
|||
self.keys = self.keys | frozenset((key,))
|
||||
|
||||
|
||||
def has_unattempted_fallback_target(
|
||||
fallback_model_group: Sequence[object] | None, kwargs: Mapping[str, object]
|
||||
) -> bool:
|
||||
"""Whether a resolved chain still holds an entry this request has not tried."""
|
||||
if fallback_model_group is None:
|
||||
return False
|
||||
attempted: Final = kwargs.get("attempted_targets")
|
||||
if not isinstance(attempted, AttemptedFallbackTargets):
|
||||
return True
|
||||
return any((key := fallback_attempt_key(target)) is None or key not in attempted for target in fallback_model_group)
|
||||
|
||||
|
||||
def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
|
||||
"""
|
||||
Handles wildcard routing scenario
|
||||
|
|
@ -272,6 +284,76 @@ def get_pre_routing_selection(kwargs: Mapping[str, object]) -> str | None:
|
|||
return next((selected for selected in selections if isinstance(selected, str) and selected), None)
|
||||
|
||||
|
||||
def carry_over_pre_routing_selection(live_kwargs: Mapping[str, object], snapshot: Mapping[str, object]) -> None:
|
||||
"""
|
||||
Replace whatever selection the snapshot carries with the one the pre-routing hook stamped
|
||||
into the live kwargs while routing this attempt, so a mid-stream fallback keys its lookup
|
||||
off the tier this attempt actually routed to.
|
||||
"""
|
||||
clear_pre_routing_selection(snapshot)
|
||||
live_selection: Final = get_pre_routing_selection(live_kwargs)
|
||||
if live_selection is not None:
|
||||
record_pre_routing_selection(snapshot, live_selection)
|
||||
|
||||
|
||||
MID_STREAM_FALLBACK_CONTROLS_KEY: Final = "_mid_stream_fallback_controls"
|
||||
_PER_REQUEST_FALLBACK_CONTROL_KEYS: Final = (
|
||||
"fallbacks",
|
||||
"context_window_fallbacks",
|
||||
"content_policy_fallbacks",
|
||||
"num_retries",
|
||||
"model_group_retry_policy",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MidStreamFallbackControls:
|
||||
"""
|
||||
The per-request fallback and retry overrides every streaming attempt must see again.
|
||||
|
||||
async_function_with_retries pops them before the attempt function runs, so without this
|
||||
carrier a fallback hop's own mid-stream re-entry would fall back to the router-level settings.
|
||||
"""
|
||||
|
||||
overrides: Mapping[str, object]
|
||||
|
||||
|
||||
_NO_FALLBACK_CONTROLS: Final = MidStreamFallbackControls(MappingProxyType({}))
|
||||
|
||||
|
||||
def per_request_fallback_controls(kwargs: Mapping[str, object]) -> MidStreamFallbackControls:
|
||||
return MidStreamFallbackControls(
|
||||
MappingProxyType({key: kwargs[key] for key in _PER_REQUEST_FALLBACK_CONTROL_KEYS if key in kwargs})
|
||||
)
|
||||
|
||||
|
||||
def mid_stream_fallback_hop_kwargs(
|
||||
model: str,
|
||||
original_generic_function: Callable[..., object],
|
||||
controls: object,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> dict[str, object]: # mutable-ok: the streaming iterators rewrite it in place when they re-enter the chain
|
||||
"""
|
||||
The kwargs one streaming attempt re-enters the fallback chain with if its stream fails.
|
||||
|
||||
A shallow copy keeps ``attempted_targets`` shared with the outer chain, so entries this
|
||||
request already tried are never retried; the metadata buckets are copied key by key because
|
||||
the attempt writes deployment-specific fields into them in place.
|
||||
"""
|
||||
hop_controls: Final = controls if isinstance(controls, MidStreamFallbackControls) else _NO_FALLBACK_CONTROLS
|
||||
copied_buckets: Final = MappingProxyType(
|
||||
{name: safe_deep_copy(kwargs[name]) for name in _ROUTER_METADATA_BUCKETS if isinstance(kwargs.get(name), dict)}
|
||||
)
|
||||
return { # mutable-ok: handed to the streaming iterator as its initial_kwargs, which it rewrites on re-entry
|
||||
**kwargs,
|
||||
**copied_buckets,
|
||||
**hop_controls.overrides,
|
||||
MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls,
|
||||
"model": model,
|
||||
"original_generic_function": original_generic_function,
|
||||
}
|
||||
|
||||
|
||||
DISABLE_FALLBACKS_METADATA_KEY: Final = "_disable_fallbacks"
|
||||
|
||||
|
||||
|
|
@ -307,13 +389,19 @@ def fallbacks_disabled_for_request(kwargs: Mapping[str, object]) -> bool:
|
|||
def fallback_lookup_groups(kwargs: Mapping[str, object], model_group: str | None) -> tuple[str, ...]:
|
||||
"""
|
||||
Ordered keys for resolving a fallback chain: the tier a pre-routing hook selected wins,
|
||||
then the routed group, then the requested group. The routed group differs when Claude Code
|
||||
session affinity remaps a subagent's concrete model to its bound router.
|
||||
then the routed group, then the requested group, then the group the request was
|
||||
originally for. The routed group differs when Claude Code session affinity remaps a
|
||||
subagent's concrete model to its bound router. The original group differs on a fallback
|
||||
hop that fails after `run_async_fallback` already returned its stream: the hop has no
|
||||
chain of its own, so it resumes the original group's chain, and `attempted_targets` keeps
|
||||
the entries already tried from being repeated.
|
||||
"""
|
||||
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
|
||||
routed_group_value: Final = metadata.get("model_group") if isinstance(metadata, Mapping) else None
|
||||
routed_group: Final = routed_group_value if isinstance(routed_group_value, str) else None
|
||||
ordered: Final = (get_pre_routing_selection(kwargs), routed_group, model_group)
|
||||
original_group_value: Final = metadata.get("original_model_group") if isinstance(metadata, Mapping) else None
|
||||
original_group: Final = original_group_value if isinstance(original_group_value, str) else None
|
||||
ordered: Final = (get_pre_routing_selection(kwargs), routed_group, model_group, original_group)
|
||||
return tuple(dict.fromkeys(group for group in ordered if group))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -251,7 +251,7 @@ async def test_aresponses_with_streaming_fallbacks_non_streaming_passthrough():
|
|||
|
||||
with patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(return_value=plain_response),
|
||||
):
|
||||
out = await router._aresponses_with_streaming_fallbacks(
|
||||
|
|
@ -278,7 +278,7 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
|
|||
|
||||
with patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(return_value=streaming_iter),
|
||||
), patch.object(
|
||||
router,
|
||||
|
|
@ -294,6 +294,173 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
|
|||
mock_wrap.assert_awaited_once()
|
||||
|
||||
|
||||
# -------- every fallback entry stays reachable across hops --------
|
||||
|
||||
|
||||
def _make_three_tier_router(**router_kwargs) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "sk-test"}},
|
||||
{"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "sk-test"}},
|
||||
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "sk-test"}},
|
||||
],
|
||||
num_retries=0,
|
||||
**router_kwargs,
|
||||
)
|
||||
|
||||
|
||||
def _mid_stream_failure(model: str):
|
||||
import litellm
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
return MidStreamFallbackError(
|
||||
message="stream dropped",
|
||||
model=model,
|
||||
llm_provider="openai",
|
||||
original_exception=litellm.InternalServerError(message="stream dropped", llm_provider="openai", model=model),
|
||||
is_pre_first_chunk=True,
|
||||
)
|
||||
|
||||
|
||||
def _scripted_responses_stream(events: list, error: Exception | None = None):
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
class _ScriptedStream(BaseResponsesAPIStreamingIterator):
|
||||
def __init__(self) -> None:
|
||||
self._events = list(events)
|
||||
self._hidden_params: dict = {}
|
||||
self.completed_response = None
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
if self._events:
|
||||
return self._events.pop(0)
|
||||
if error is not None:
|
||||
raise error
|
||||
raise StopAsyncIteration
|
||||
|
||||
async def aclose(self) -> None:
|
||||
return None
|
||||
|
||||
return _ScriptedStream()
|
||||
|
||||
|
||||
def _three_tier_original(calls: list, primary_fails_pre_stream: bool):
|
||||
import litellm
|
||||
|
||||
completed_event = _make_completed_event(1, 1, 2)
|
||||
|
||||
async def fake_original(**kwargs):
|
||||
model = kwargs["model"]
|
||||
calls.append(model)
|
||||
if model == "openai/primary-model":
|
||||
if primary_fails_pre_stream:
|
||||
raise litellm.InternalServerError(message="primary down", llm_provider="openai", model=model)
|
||||
return _scripted_responses_stream([], _mid_stream_failure(model))
|
||||
if model == "openai/fb1-model":
|
||||
return _scripted_responses_stream([], _mid_stream_failure(model))
|
||||
return _scripted_responses_stream([completed_event])
|
||||
|
||||
return fake_original, completed_event
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_pre_stream_primary_failure_then_hop_stream_failure_reaches_second_entry():
|
||||
"""Regression: fallbacks=[{"primary": ["fb1", "fb2"]}]. The primary fails before streaming,
|
||||
fb1 is reached through the regular fallback chain and then fails mid-stream. Only the
|
||||
primary's stream used to be wrapped, so fb1's mid-stream failure either re-raised or
|
||||
re-tried fb1 itself; fb2 was unreachable."""
|
||||
router = _make_three_tier_router(fallbacks=[{"primary": ["fb1", "fb2"]}])
|
||||
calls: list = []
|
||||
fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=True)
|
||||
|
||||
stream = await router._aresponses_with_streaming_fallbacks(
|
||||
original_function=fake_original, model="primary", stream=True, input="hi"
|
||||
)
|
||||
collected = [event async for event in stream]
|
||||
|
||||
assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"]
|
||||
assert collected == [completed_event]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_two_consecutive_mid_stream_failures_reach_second_entry():
|
||||
"""Regression: the primary and fb1 both fail mid-stream; fb2 must still be tried."""
|
||||
router = _make_three_tier_router(fallbacks=[{"primary": ["fb1", "fb2"]}])
|
||||
calls: list = []
|
||||
fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=False)
|
||||
|
||||
stream = await router._aresponses_with_streaming_fallbacks(
|
||||
original_function=fake_original, model="primary", stream=True, input="hi"
|
||||
)
|
||||
collected = [event async for event in stream]
|
||||
|
||||
assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"]
|
||||
assert collected == [completed_event]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_per_request_fallbacks_survive_into_hop_streams():
|
||||
"""Regression: a request-level fallbacks list (key or team router_settings) is popped
|
||||
before each attempt runs, so a hop's mid-stream re-entry used to see only the router's
|
||||
own (empty) list and gave up after fb1."""
|
||||
router = _make_three_tier_router()
|
||||
calls: list = []
|
||||
fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=False)
|
||||
|
||||
stream = await router._aresponses_with_streaming_fallbacks(
|
||||
original_function=fake_original,
|
||||
model="primary",
|
||||
stream=True,
|
||||
input="hi",
|
||||
fallbacks=[{"primary": ["fb1", "fb2"]}],
|
||||
)
|
||||
collected = [event async for event in stream]
|
||||
|
||||
assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"]
|
||||
assert collected == [completed_event]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_attempt_strips_the_controls_carrier_and_wraps_every_hop_stream():
|
||||
"""Each attempt of the chain, not only the primary's, comes back wrapped for mid-stream
|
||||
failover, and the per-request controls carrier rides into the wrapper's re-entry kwargs
|
||||
without ever reaching the provider call."""
|
||||
from types import MappingProxyType
|
||||
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
MID_STREAM_FALLBACK_CONTROLS_KEY,
|
||||
MidStreamFallbackControls,
|
||||
)
|
||||
|
||||
router = _make_three_tier_router()
|
||||
completed_event = _make_completed_event(1, 1, 2)
|
||||
hop_stream = _scripted_responses_stream([completed_event])
|
||||
seen: dict = {}
|
||||
|
||||
async def fake_original(**kwargs):
|
||||
seen.update(kwargs)
|
||||
return hop_stream
|
||||
|
||||
controls = MidStreamFallbackControls(MappingProxyType({"fallbacks": [{"primary": ["fb1", "fb2"]}]}))
|
||||
stream = await router._ageneric_api_call_with_fallbacks_responses_attempt(
|
||||
model="fb1",
|
||||
original_generic_function=fake_original,
|
||||
stream=True,
|
||||
input="hi",
|
||||
**{MID_STREAM_FALLBACK_CONTROLS_KEY: controls},
|
||||
)
|
||||
collected = [event async for event in stream]
|
||||
|
||||
assert seen["model"] == "openai/fb1-model"
|
||||
assert MID_STREAM_FALLBACK_CONTROLS_KEY not in seen
|
||||
assert "fallbacks" not in seen
|
||||
assert stream is not hop_stream
|
||||
assert collected == [completed_event]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_fallback_on_in_stream_error_event():
|
||||
"""A retriable in-stream error event (429) must trigger the router's mid-stream
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
from datetime import datetime, timedelta
|
||||
from typing import NoReturn
|
||||
from typing import Final, NoReturn
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -1305,6 +1305,14 @@ class TestOrderedFallbackLookupGroups:
|
|||
"requested-model",
|
||||
)
|
||||
|
||||
def test_fallback_hop_resumes_the_original_groups_chain_last(self):
|
||||
from litellm.router_utils.fallback_event_handlers import fallback_lookup_groups
|
||||
|
||||
kwargs = {"metadata": {"model_group": "fb1", "original_model_group": "primary"}}
|
||||
|
||||
assert fallback_lookup_groups(kwargs, "fb1") == ("fb1", "primary")
|
||||
assert fallback_lookup_groups({"metadata": {"original_model_group": 42}}, "fb1") == ("fb1",)
|
||||
|
||||
def test_first_resolving_group_wins_and_generic_idx_survives_a_miss(self):
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
get_fallback_model_group_for_lookup_groups,
|
||||
|
|
@ -1315,3 +1323,20 @@ class TestOrderedFallbackLookupGroups:
|
|||
assert get_fallback_model_group_for_lookup_groups(fallbacks, ("tier9", "smart-router")) == (["backup-b"], None)
|
||||
assert get_fallback_model_group_for_lookup_groups(fallbacks, ("tier9", "no-such")) == (["backup-c"], 2)
|
||||
assert get_fallback_model_group_for_lookup_groups([{"tier1": ["backup-a"]}], ("no", "nope")) == (None, None)
|
||||
|
||||
|
||||
class TestHasUnattemptedFallbackTarget:
|
||||
def test_exhausted_chain_is_not_recoverable_but_a_fresh_entry_is(self):
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
has_unattempted_fallback_target,
|
||||
)
|
||||
|
||||
attempted: Final = AttemptedFallbackTargets()
|
||||
attempted.record("primary")
|
||||
attempted.record("fb1")
|
||||
attempted.record("fb2")
|
||||
|
||||
assert has_unattempted_fallback_target(["fb1", "fb2"], {"attempted_targets": attempted}) is False
|
||||
assert has_unattempted_fallback_target(["fb1", "fb3"], {"attempted_targets": attempted}) is True
|
||||
assert has_unattempted_fallback_target(["fb1"], {}) is True
|
||||
assert has_unattempted_fallback_target(None, {}) is False
|
||||
|
|
|
|||
|
|
@ -3447,6 +3447,114 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste
|
|||
assert result._hidden_params["model_id"] == "served-deployment"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configured_list():
|
||||
"""LIT-7400: fallbacks=[{primary: [fb1, fb2]}] must reach fb2 when fb1 dies before its first chunk.
|
||||
|
||||
run_async_fallback returns as soon as fb1's stream wrapper exists, so fb1's failure surfaces
|
||||
inside the streaming iterator, where the lookup is keyed by fb1. That key has no chain of its
|
||||
own, so the iterator has to resume the chain of the group the request was originally for.
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
class FailingStream(CustomStreamWrapper):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(
|
||||
completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()
|
||||
)
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise MidStreamFallbackError(
|
||||
message=f"provider 500 from {self.model}",
|
||||
model=self.model,
|
||||
llm_provider="openai",
|
||||
generated_content="",
|
||||
is_pre_first_chunk=True,
|
||||
original_exception=litellm.InternalServerError(
|
||||
message=f"provider 500 from {self.model}", model=self.model, llm_provider="openai"
|
||||
),
|
||||
)
|
||||
|
||||
class OkStream(FailingStream):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(model)
|
||||
self._chunks = iter(
|
||||
[litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])]
|
||||
)
|
||||
|
||||
async def __anext__(self):
|
||||
try:
|
||||
return next(self._chunks)
|
||||
except StopIteration:
|
||||
raise StopAsyncIteration from None
|
||||
|
||||
async def fake_acompletion(**kwargs):
|
||||
if "fb2" in kwargs["model"]:
|
||||
return OkStream(kwargs["model"])
|
||||
return FailingStream(kwargs["model"])
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "fake-key"}},
|
||||
],
|
||||
fallbacks=[{"primary": ["fb1", "fb2"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
with patch("litellm.acompletion", side_effect=fake_acompletion) as mock_acompletion:
|
||||
response = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True)
|
||||
content: Final = "".join(
|
||||
[chunk.choices[0].delta.content or "" async for chunk in response if chunk is not None]
|
||||
)
|
||||
|
||||
assert content == "ok-from-openai/fb2-model"
|
||||
assert [c.kwargs["metadata"]["model_group"] for c in mock_acompletion.call_args_list] == [
|
||||
"primary",
|
||||
"fb1",
|
||||
"fb2",
|
||||
]
|
||||
|
||||
|
||||
def test_refusal_on_the_last_fallback_hop_is_returned_instead_of_raised():
|
||||
"""LIT-7400 follow-up: a refusal on the final hop of an exhausted list passes through."""
|
||||
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "fake-key"}},
|
||||
],
|
||||
fallbacks=[{"primary": ["fb1", "fb2"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
attempted: Final = AttemptedFallbackTargets()
|
||||
attempted.record("primary")
|
||||
attempted.record("fb1")
|
||||
attempted.record("fb2")
|
||||
kwargs: Final = {
|
||||
"attempted_targets": attempted,
|
||||
"metadata": {"model_group": "fb2", "original_model_group": "primary"},
|
||||
}
|
||||
|
||||
assert router._refusal_fallback_available("fb2", kwargs) is False
|
||||
assert (
|
||||
router._refusal_fallback_available(
|
||||
"fb1", {"metadata": {"model_group": "fb1", "original_model_group": "primary"}}
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_completion_streaming_iterator_adopts_fallback_response_headers():
|
||||
"""LIT-6767, sync counterpart of the fallback-adoption test."""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
|
@ -4126,7 +4234,7 @@ async def test_aresponses_streaming_iterator_fallback():
|
|||
call_kwargs = mock_fallback_utils.call_args.kwargs
|
||||
fbk = call_kwargs["kwargs"]
|
||||
# Bound methods compare equal when they share the same instance + __func__.
|
||||
assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_helper
|
||||
assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_responses_attempt
|
||||
assert fbk["original_generic_function"] is litellm.aresponses
|
||||
assert call_kwargs["model_group"] == "anthropic/claude-sonnet-4-6"
|
||||
assert call_kwargs["disable_fallbacks"] is False
|
||||
|
|
@ -13713,7 +13821,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_non_streaming_passth
|
|||
|
||||
with patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(return_value=plain_response),
|
||||
):
|
||||
out = await router._aanthropic_messages_with_streaming_fallbacks(
|
||||
|
|
@ -13737,7 +13845,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_wraps_streaming_iter
|
|||
with (
|
||||
patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(return_value=streaming_iter),
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -14022,7 +14130,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_nested_m
|
|||
):
|
||||
with patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(side_effect=fake_original),
|
||||
):
|
||||
await router._aanthropic_messages_with_streaming_fallbacks(
|
||||
|
|
@ -14056,7 +14164,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_metadata
|
|||
):
|
||||
with patch.object(
|
||||
router,
|
||||
"_ageneric_api_call_with_fallbacks",
|
||||
"_ageneric_api_call_with_fallbacks_helper",
|
||||
new=AsyncMock(side_effect=fake_original),
|
||||
):
|
||||
await router._aanthropic_messages_with_streaming_fallbacks(
|
||||
|
|
@ -14071,6 +14179,95 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_metadata
|
|||
assert "deployment" not in fallback_kwargs["metadata"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_hop_stream_failure_reaches_second_fallback_entry():
|
||||
"""Regression: fallbacks=[{"primary": ["fb1", "fb2"]}]. The primary fails before
|
||||
streaming, fb1 is reached through the regular fallback chain and then sends an
|
||||
error frame mid-stream. Only the primary's stream used to be wrapped, so the outer
|
||||
wrapper re-tried fb1 with a fresh attempted set and forwarded fb1's error frame to
|
||||
the client on an HTTP 200; fb2 was unreachable."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "anthropic/primary-model", "api_key": "sk-test"}},
|
||||
{"model_name": "fb1", "litellm_params": {"model": "anthropic/fb1-model", "api_key": "sk-test"}},
|
||||
{"model_name": "fb2", "litellm_params": {"model": "anthropic/fb2-model", "api_key": "sk-test"}},
|
||||
],
|
||||
num_retries=0,
|
||||
fallbacks=[{"primary": ["fb1", "fb2"]}],
|
||||
)
|
||||
calls: list = []
|
||||
|
||||
async def fake_original(**kwargs):
|
||||
model = kwargs["model"]
|
||||
calls.append(model)
|
||||
if model == "anthropic/primary-model":
|
||||
raise litellm.InternalServerError(message="primary down", llm_provider="anthropic", model=model)
|
||||
if model == "anthropic/fb1-model":
|
||||
return _AnthropicMessagesFakeByteStream(
|
||||
[_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()]
|
||||
)
|
||||
return _AnthropicMessagesFakeByteStream(
|
||||
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb2")]
|
||||
)
|
||||
|
||||
stream = await router._aanthropic_messages_with_streaming_fallbacks(
|
||||
original_function=fake_original,
|
||||
model="primary",
|
||||
stream=True,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_tokens=10,
|
||||
)
|
||||
body = b"".join([chunk async for chunk in stream])
|
||||
|
||||
assert calls == ["anthropic/primary-model", "anthropic/fb1-model", "anthropic/fb2-model"]
|
||||
assert b"from fb2" in body
|
||||
assert b"overloaded_error" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_attempt_strips_the_controls_carrier_and_wraps_every_hop_stream():
|
||||
"""Each attempt of the chain, not only the primary's, comes back wrapped for mid-stream
|
||||
failover, and the per-request controls carrier never reaches the provider call."""
|
||||
from types import MappingProxyType
|
||||
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
MID_STREAM_FALLBACK_CONTROLS_KEY,
|
||||
MidStreamFallbackControls,
|
||||
)
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "fb1", "litellm_params": {"model": "anthropic/fb1-model", "api_key": "sk-test"}},
|
||||
],
|
||||
num_retries=0,
|
||||
)
|
||||
hop_stream = _AnthropicMessagesFakeByteStream(
|
||||
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb1")]
|
||||
)
|
||||
seen: dict = {}
|
||||
|
||||
async def fake_original(**kwargs):
|
||||
seen.update(kwargs)
|
||||
return hop_stream
|
||||
|
||||
controls = MidStreamFallbackControls(MappingProxyType({"fallbacks": [{"primary": ["fb1", "fb2"]}]}))
|
||||
stream = await router._ageneric_api_call_with_fallbacks_anthropic_messages_attempt(
|
||||
model="fb1",
|
||||
original_generic_function=fake_original,
|
||||
stream=True,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_tokens=10,
|
||||
**{MID_STREAM_FALLBACK_CONTROLS_KEY: controls},
|
||||
)
|
||||
body = b"".join([chunk async for chunk in stream])
|
||||
|
||||
assert seen["model"] == "anthropic/fb1-model"
|
||||
assert MID_STREAM_FALLBACK_CONTROLS_KEY not in seen
|
||||
assert "fallbacks" not in seen
|
||||
assert stream is not hop_stream
|
||||
assert b"from fb1" in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_fallback_triggers_after_lifecycle_only_frame():
|
||||
"""Regression: Anthropic routinely sends a message_start lifecycle frame
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue