diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 22b04062534..2a9312dc0e3 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 88c4a7011ef..7bedb714338 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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() diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index e33897e3b63..239ee400503 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -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, } diff --git a/tests/integration/sdk/test_router_sync_stream_fallback_wire.py b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py index f5095c5a778..fc5fb028123 100644 --- a/tests/integration/sdk/test_router_sync_stream_fallback_wire.py +++ b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py @@ -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) diff --git a/tests/unit/proxy/test_route_llm_request.py b/tests/unit/proxy/test_route_llm_request.py index a891e1079d8..d5546da024b 100644 --- a/tests/unit/proxy/test_route_llm_request.py +++ b/tests/unit/proxy/test_route_llm_request.py @@ -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 diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index adfce2c34be..38c13decaa7 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -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 diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index e3e6326574b..cd490e127cf 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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