fix(router): honor a per-request fallbacks list on a mid-stream fallback (#45227)

* fix(router): honor a per-request fallbacks list on a mid-stream fallback

The completion entrypoints record the per-request fallbacks, context_window_fallbacks, and content_policy_fallbacks in the mid-stream controls carrier the Responses and Messages paths already use, and each completion attempt restores them into the kwargs its stream re-enters the fallback chain with, so a streaming request that dies before its first chunk fails over to the list the request named instead of the router-level list only. The sync vs async parity cell now asserts the backup's text on both twins.

* test(router): annotate the new fallback test locals as Final and flatten the leak check

* fix(proxy): skip null key and team router settings when merging per-request overrides

* test(proxy): annotate the null router settings test locals and drop a redundant comment
This commit is contained in:
Mateo Wang 2026-10-07 22:25:22 -07:00 • committed by GitHub
parent a7ee038592
commit 328f5a720c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 246 additions and 35 deletions

View file

@ -678,9 +678,8 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
"enable_tag_filtering",
]
# Merge override settings into data (only if not already set in request)
for key in per_request_settings:
if key in override_settings and key not in data:
if override_settings.get(key) is not None and key not in data:
data[key] = override_settings[key]
# Use main router with overridden kwargs

View file

@ -228,6 +228,7 @@ from litellm.router_utils.fallback_event_handlers import (
get_pre_routing_selection,
has_unattempted_fallback_target,
mid_stream_fallback_hop_kwargs,
mid_stream_fallback_snapshot_kwargs,
mid_stream_retry_kwargs,
per_request_fallback_controls,
record_disable_fallbacks,
@ -2661,6 +2662,8 @@ class Router:
kwargs["model"] = model
kwargs["messages"] = messages
kwargs["original_function"] = self._completion
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)
response: Final = self.function_with_fallbacks(**kwargs)
@ -2672,10 +2675,10 @@ class Router:
model_name = None
deployment = None
try:
# Capture kwargs before deployment selection so the streaming
# fallback iterator can re-dispatch with the original model group.
input_kwargs_for_streaming_fallback: Final = kwargs.copy()
input_kwargs_for_streaming_fallback["model"] = model
controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None)
input_kwargs_for_streaming_fallback: Final = mid_stream_fallback_snapshot_kwargs(
model=model, controls=controls, kwargs=kwargs
)
# pick the one that is available (lowest TPM/RPM)
deployment = self.get_available_deployment(
@ -2923,6 +2926,8 @@ class Router:
messages=messages,
kwargs=kwargs,
)
controls: Final = per_request_fallback_controls(kwargs)
kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop
if request_priority is not None and isinstance(request_priority, int):
response = await self.schedule_acompletion(**kwargs)
else:
@ -3818,8 +3823,10 @@ class Router:
deployment = None
_timeout_debug_deployment_dict = {} # this is a temporary dict to debug timeout issues
try:
input_kwargs_for_streaming_fallback: Final = kwargs.copy()
input_kwargs_for_streaming_fallback["model"] = model
controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None)
input_kwargs_for_streaming_fallback: Final = mid_stream_fallback_snapshot_kwargs(
model=model, controls=controls, kwargs=kwargs
)
parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs)
start_time: Final = time.time()

View file

@ -335,6 +335,28 @@ def per_request_fallback_controls(kwargs: Mapping[str, object]) -> MidStreamFall
)
def mid_stream_fallback_snapshot_kwargs(
model: str,
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 a completion attempt's stream re-enters the fallback chain with if it fails.
async_function_with_retries popped the per-request fallback lists before the attempt ran, so
the carrier restores them here and rides along into every hop this re-entry opens. A shallow
copy keeps the metadata buckets shared with the live kwargs, the way the attempt's own
in-place bucket writes expect.
"""
hop_controls: Final = controls if isinstance(controls, MidStreamFallbackControls) else _NO_FALLBACK_CONTROLS
return {
**kwargs,
**hop_controls.overrides,
MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls,
"model": model,
}
def mid_stream_fallback_hop_kwargs(
model: str,
original_generic_function: Callable[..., object],
@ -342,22 +364,18 @@ def mid_stream_fallback_hop_kwargs(
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.
The kwargs one generic-endpoint 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 {
**kwargs,
**mid_stream_fallback_snapshot_kwargs(model=model, controls=controls, kwargs=kwargs),
**copied_buckets,
**hop_controls.overrides,
MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls,
"model": model,
"original_generic_function": original_generic_function,
}

View file

@ -199,28 +199,23 @@ def test_router_retries_configured(client: str) -> None:
assert _deployments_hit(wire) == ("primary", "backup")
@dataclass(frozen=True, slots=True)
class _Outcome:
text: str | None
error: str | None
hit: tuple[str, ...]
def _outcome(client: str, wire: Wire, router: Router, **request: object) -> _Outcome:
try:
streamed: Final = _stream(client, router, **request)
except litellm.APIConnectionError as error:
return _Outcome(text=None, error=type(error).__name__, hit=_deployments_hit(wire))
return _Outcome(text=streamed.text, error=None, hit=_deployments_hit(wire))
def test_per_request_fallback_list_behaves_like_the_async_twin() -> None:
@pytest.mark.parametrize("client", _CLIENTS)
def test_per_request_fallback_list(client: str) -> None:
with wire_server(_peer(_PRIMARY_DIES)) as wire:
router: Final = _router(wire, ("primary", "backup"))
twin: Final = _outcome("async", wire, router, fallbacks=_PRIMARY_TO_BACKUP)
observed: Final = _outcome("sync", wire, router, fallbacks=_PRIMARY_TO_BACKUP)
assert observed == twin, (observed, twin)
assert observed.hit[:1] == ("primary",), observed
streamed: Final = _stream(client, router, fallbacks=_PRIMARY_TO_BACKUP)
assert streamed.text == "answered by the backup", streamed
assert streamed.attempted_fallbacks == 1, streamed
assert _deployments_hit(wire) == ("primary", "backup")
@pytest.mark.parametrize("client", _CLIENTS)
def test_per_request_fallbacks_none_turns_the_router_list_off(client: str) -> None:
with wire_server(_peer(_PRIMARY_DIES)) as wire:
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
with pytest.raises(litellm.APIConnectionError, match="overloaded"):
_stream(client, router, fallbacks=None)
assert _deployments_hit(wire) == ("primary",)
@pytest.mark.parametrize("client", _CLIENTS)

View file

@ -1728,3 +1728,37 @@ async def test_route_request_without_model_on_model_routed_endpoint_is_a_400():
assert exc_info.value.code == "400"
assert exc_info.value.param == "model"
@pytest.mark.asyncio
async def test_route_request_router_settings_override_skips_null_fields():
"""
A key or team saved from the dashboard stores every unset router setting as null. Those nulls
must not reach the router as explicit per-request values, or they switch the router-level
fallbacks and retries off for that key.
"""
data: Final = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
"stream": True,
"router_settings_override": {
"fallbacks": None,
"context_window_fallbacks": None,
"num_retries": None,
"model_group_retry_policy": None,
"timeout": 600,
},
}
llm_router: Final = MagicMock()
llm_router.acompletion.return_value = "success"
response: Final = await route_request(data, llm_router, None, "acompletion")
assert response == "success"
call_kwargs: Final = llm_router.acompletion.call_args[1]
assert call_kwargs["timeout"] == 600
assert "fallbacks" not in call_kwargs
assert "context_window_fallbacks" not in call_kwargs
assert "num_retries" not in call_kwargs
assert "model_group_retry_policy" not in call_kwargs

View file

@ -30,6 +30,7 @@ from litellm.router_utils.fallback_event_handlers import (
get_pre_routing_selection,
log_failure_fallback_event,
log_success_fallback_event,
mid_stream_fallback_snapshot_kwargs,
mid_stream_retry_kwargs,
record_pre_routing_selection,
record_retry_attempt,
@ -1471,6 +1472,30 @@ def test_get_fallback_model_group_never_resolves_a_provider_without_a_prefixed_k
resolver.assert_not_called()
def test_mid_stream_fallback_snapshot_kwargs_restores_the_popped_lists_and_shares_the_buckets():
controls: Final = MidStreamFallbackControls(
MappingProxyType({"fallbacks": [{"primary": ["backup"]}], "context_window_fallbacks": None})
)
metadata: Final = {"model_group": "primary"}
kwargs: Final = {"messages": [{"role": "user", "content": "hi"}], "stream": True, "metadata": metadata}
snapshot: Final = mid_stream_fallback_snapshot_kwargs(model="primary", controls=controls, kwargs=kwargs)
assert snapshot == {
**kwargs,
"fallbacks": [{"primary": ["backup"]}],
"context_window_fallbacks": None,
MID_STREAM_FALLBACK_CONTROLS_KEY: controls,
"model": "primary",
}
assert snapshot["metadata"] is metadata
assert "fallbacks" not in kwargs
bare: Final = mid_stream_fallback_snapshot_kwargs(model="primary", controls=None, kwargs=kwargs)
assert "fallbacks" not in bare
assert bare[MID_STREAM_FALLBACK_CONTROLS_KEY] == MidStreamFallbackControls(MappingProxyType({}))
def test_mid_stream_retry_kwargs_strips_what_the_retry_wrapper_pops_and_keeps_the_controls_carrier():
def generic_function(**kwargs) -> None:
return None

View file

@ -63,7 +63,10 @@ from litellm.router_utils.cooldown_handlers import (
async_get_cooldown_deployments,
get_cooldown_deployments,
)
from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY
from litellm.router_utils.fallback_event_handlers import (
DISABLE_FALLBACKS_METADATA_KEY,
MID_STREAM_FALLBACK_CONTROLS_KEY,
)
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
from litellm.scheduler import FlowItem
from litellm.types.llms.openai import ChatCompletionRequest
@ -3868,6 +3871,136 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur
]
class _DiesBeforeFirstChunk(CustomStreamWrapper):
def __init__(self, model: str):
super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock())
def _mid_stream_error(self) -> MidStreamFallbackError:
return 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"
),
)
def __iter__(self):
return self
def __next__(self):
raise self._mid_stream_error()
def __aiter__(self):
return self
async def __anext__(self):
raise self._mid_stream_error()
class _Answers(_DiesBeforeFirstChunk):
def __init__(self, model: str):
super().__init__(model)
self._chunks = iter(
[litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])]
)
def __next__(self):
return next(self._chunks)
async def __anext__(self):
try:
return next(self._chunks)
except StopIteration:
raise StopAsyncIteration from None
def _primary_and_backup_router(**settings: object) -> litellm.Router:
return litellm.Router(
model_list=[
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}},
{"model_name": "backup", "litellm_params": {"model": "openai/backup-model", "api_key": "fake-key"}},
],
num_retries=0,
**settings,
)
def _stream_for(**kwargs: object) -> CustomStreamWrapper:
model: Final = str(kwargs["model"])
return _Answers(model) if "backup" in model else _DiesBeforeFirstChunk(model)
def _groups_called(provider_calls: MagicMock) -> list[str]:
return [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list]
def _router_internals_reached_the_provider(provider_calls: MagicMock) -> bool:
leaked: Final = frozenset(
("fallbacks", "context_window_fallbacks", "content_policy_fallbacks", MID_STREAM_FALLBACK_CONTROLS_KEY)
)
return any(leaked & call.kwargs.keys() for call in provider_calls.call_args_list)
def test_completion_mid_stream_fallback_honors_the_per_request_list():
router: Final = _primary_and_backup_router()
with patch("litellm.completion", side_effect=_stream_for) as provider_calls:
response: Final = router.completion(
model="primary",
messages=[{"role": "user", "content": "hi"}],
stream=True,
fallbacks=[{"primary": ["backup"]}],
)
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response if chunk is not None)
assert content == "ok-from-openai/backup-model"
assert _groups_called(provider_calls) == ["primary", "backup"]
assert not _router_internals_reached_the_provider(provider_calls)
@pytest.mark.asyncio
async def test_acompletion_mid_stream_fallback_honors_the_per_request_list():
router: Final = _primary_and_backup_router()
async def fake_acompletion(**kwargs):
return _stream_for(**kwargs)
with patch("litellm.acompletion", side_effect=fake_acompletion) as provider_calls:
response: Final = await router.acompletion(
model="primary",
messages=[{"role": "user", "content": "hi"}],
stream=True,
fallbacks=[{"primary": ["backup"]}],
)
content: Final = "".join(
[chunk.choices[0].delta.content or "" async for chunk in response if chunk is not None]
)
assert content == "ok-from-openai/backup-model"
assert _groups_called(provider_calls) == ["primary", "backup"]
assert not _router_internals_reached_the_provider(provider_calls)
@pytest.mark.asyncio
async def test_acompletion_mid_stream_fallback_honors_a_per_request_fallbacks_none():
router: Final = _primary_and_backup_router(fallbacks=[{"primary": ["backup"]}])
async def fake_acompletion(**kwargs):
return _stream_for(**kwargs)
with patch("litellm.acompletion", side_effect=fake_acompletion) as provider_calls:
response: Final = await router.acompletion(
model="primary", messages=[{"role": "user", "content": "hi"}], stream=True, fallbacks=None
)
with pytest.raises(litellm.InternalServerError, match="provider 500 from openai/primary-model"):
[chunk async for chunk in response]
assert _groups_called(provider_calls) == ["primary"]
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