mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
a7ee038592
commit
328f5a720c
7 changed files with 246 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue