From 88f15e75720f9232ef9d1beef8dab5c997041c74 Mon Sep 17 00:00:00 2001 From: amarrtech Date: Wed, 7 Oct 2026 15:50:24 -0700 Subject: [PATCH] fix(router): resume sync streaming fallbacks without retrying primary (#43959) * fix(router): resume sync streaming fallback chain Signed-off-by: amarrtech <272048731+amarrtech@users.noreply.github.com> * test(router): cover the sync mid-stream fallback walking every configured target * test(router): cover the sync stream fallback walk against a wire upstream --------- Signed-off-by: amarrtech <272048731+amarrtech@users.noreply.github.com> Co-authored-by: amarrtech <272048731+amarrtech@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router.py | 10 +- .../test_router_sync_stream_fallback_wire.py | 271 ++++++++++++++++++ tests/unit/test_router/test_router.py | 144 +++++++++- 3 files changed, 415 insertions(+), 10 deletions(-) create mode 100644 tests/integration/sdk/test_router_sync_stream_fallback_wire.py diff --git a/litellm/router.py b/litellm/router.py index ab9d8884a1d..7b89d5eb5ac 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3713,11 +3713,17 @@ class Router: initial_kwargs["original_function"] = router_self._completion initial_kwargs["messages"] = messages router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) - fallback_response = router_self.function_with_fallbacks( - **initial_kwargs, + fallback_response = run_async_function( + router_self.async_function_with_fallbacks_common_utils, + e=e, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + args=(), + kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, ) if hasattr(fallback_response, "__iter__"): diff --git a/tests/integration/sdk/test_router_sync_stream_fallback_wire.py b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py new file mode 100644 index 00000000000..f5095c5a778 --- /dev/null +++ b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +import asyncio +import json +from collections.abc import Callable, Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from typing import Final + +import litellm +import pytest +from integration._support.wire import Reply, Request, Wire, wire_server +from litellm import Router +from litellm.integrations.custom_logger import CustomLogger + +_MODEL: Final = "gpt-5.6" +_API_KEY: Final = "synthetic-sync-fallback-key" +_PROMPT: Final = "which deployment answers when the primary dies before its first chunk?" +_ERROR_FRAME: Final = ( + b"data: " + json.dumps({"error": {"message": "overloaded", "type": "server_error", "code": 500}}).encode() + b"\n\n" +) +_DONE: Final = b"data: [DONE]\n\n" +_BURST: Final = 6 + + +def _delta(text: str, finish_reason: str | None) -> bytes: + chunk: Final = { + "id": "chatcmpl-sync-fallback-wire", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": _MODEL, + "choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": finish_reason}], + } + return b"data: " + json.dumps(chunk).encode() + b"\n\n" + + +def _serves(text: str) -> Reply: + return Reply(content_type="text/event-stream", chunks=(_delta(text, None), _delta("", "stop"), _DONE)) + + +def _dies_after(text: str) -> Reply: + return Reply(content_type="text/event-stream", chunks=(_delta(text, None), _ERROR_FRAME, _DONE)) + + +_DIES_BEFORE_CONTENT: Final = Reply(content_type="text/event-stream", chunks=(_ERROR_FRAME, _DONE)) +_DROPS_BEFORE_CONTENT: Final = Reply(content_type="text/event-stream", chunks=(_DONE,), abort_after=0) + + +def _peer(replies: Mapping[str, Reply]) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + deployment, _, route = request.target.lstrip("/").partition("/") + assert route == "chat/completions", request.target + return replies[deployment] + + return respond + + +def _deployments_hit(wire: Wire) -> tuple[str, ...]: + return tuple(request.target.lstrip("/").partition("/")[0] for request in wire.drain()) + + +def _router(wire: Wire, deployments: tuple[str, ...], **settings: object) -> Router: + return Router( + model_list=[ + { + "model_name": name, + "litellm_params": {"model": f"openai/{_MODEL}", "api_base": f"{wire.url}/{name}", "api_key": _API_KEY}, + } + for name in deployments + ], + num_retries=0, + disable_cooldowns=True, + **settings, + ) + + +class _FallbackRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.successes: tuple[str, ...] = () + self.failures: tuple[str, ...] = () + + async def log_success_fallback_event( + self, original_model_group: str, kwargs: dict, original_exception: Exception + ) -> None: + self.successes = (*self.successes, original_model_group) + + async def log_failure_fallback_event( + self, original_model_group: str, kwargs: dict, original_exception: Exception + ) -> None: + self.failures = (*self.failures, original_model_group) + + +@dataclass(frozen=True, slots=True) +class _Streamed: + text: str + attempted_fallbacks: object + + +def _text_of(chunk: object) -> str: + choices: Final = getattr(chunk, "choices", None) or () + return "".join(str(choice.delta.content or "") for choice in choices) + + +def _attempted_fallbacks(stream: object) -> object: + hidden: Final = getattr(stream, "_hidden_params", None) or {} + return (hidden.get("additional_headers") or {}).get("x-litellm-attempted-fallbacks") + + +def _stream_sync(router: Router, **request: object) -> _Streamed: + stream: Final = router.completion(model="primary", messages=[{"role": "user", "content": _PROMPT}], stream=True, **request) + text: Final = "".join(_text_of(chunk) for chunk in stream) + return _Streamed(text=text, attempted_fallbacks=_attempted_fallbacks(stream)) + + +async def _stream_async(router: Router, **request: object) -> _Streamed: + stream: Final = await router.acompletion( + model="primary", messages=[{"role": "user", "content": _PROMPT}], stream=True, **request + ) + parts: Final = [_text_of(chunk) async for chunk in stream] + return _Streamed(text="".join(parts), attempted_fallbacks=_attempted_fallbacks(stream)) + + +def _stream(client: str, router: Router, **request: object) -> _Streamed: + if client == "async": + return asyncio.run(_stream_async(router, **request)) + return _stream_sync(router, **request) + + +_CLIENTS: Final = ("sync", "async") +_PRIMARY_DIES: Final = {"primary": _DIES_BEFORE_CONTENT, "backup": _serves("answered by the backup")} +_PRIMARY_AND_FB1_DIE: Final = {"primary": _DIES_BEFORE_CONTENT, "fb1": _DIES_BEFORE_CONTENT, "fb2": _serves("answered by fb2")} +_PRIMARY_TO_BACKUP: Final = [{"primary": ["backup"]}] + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_primary_dies_before_content(client: str, monkeypatch: pytest.MonkeyPatch) -> None: + recorder: Final = _FallbackRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + with wire_server(_peer(_PRIMARY_DIES)) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + streamed: Final = _stream(client, router) + assert streamed.text == "answered by the backup", streamed + assert streamed.attempted_fallbacks == 1, streamed + assert _deployments_hit(wire) == ("primary", "backup") + assert recorder.successes == ("primary",), recorder.successes + assert recorder.failures == (), recorder.failures + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_walks_every_configured_fallback(client: str) -> None: + with wire_server(_peer(_PRIMARY_AND_FB1_DIE)) as wire: + router: Final = _router(wire, ("primary", "fb1", "fb2"), fallbacks=[{"primary": ["fb1", "fb2"]}]) + streamed: Final = _stream(client, router) + assert streamed.text == "answered by fb2", streamed + assert streamed.attempted_fallbacks == 2, streamed + assert _deployments_hit(wire) == ("primary", "fb1", "fb2") + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_every_target_dies(client: str) -> None: + with wire_server(_peer({"primary": _DIES_BEFORE_CONTENT, "backup": _DIES_BEFORE_CONTENT})) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + with pytest.raises(litellm.APIConnectionError, match="overloaded"): + _stream(client, router) + assert _deployments_hit(wire) == ("primary", "backup") + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_fallbacks_disabled(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, disable_fallbacks=True) + assert _deployments_hit(wire) == ("primary",) + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_dies_after_first_chunk(client: str) -> None: + with wire_server(_peer({"primary": _dies_after("partial "), "backup": _serves("never asked")})) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + with pytest.raises(litellm.APIConnectionError, match="overloaded"): + _stream(client, router) + assert _deployments_hit(wire) == ("primary",) + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_router_retries_configured(client: str) -> None: + with wire_server(_peer(_PRIMARY_DIES)) as wire: + router: Final = Router( + model_list=_router(wire, ("primary", "backup")).model_list, + fallbacks=_PRIMARY_TO_BACKUP, + num_retries=2, + disable_cooldowns=True, + ) + streamed: Final = _stream(client, router) + assert streamed.text == "answered by the backup", streamed + 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: + 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 + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_max_fallbacks_caps_the_walk(client: str) -> None: + with wire_server(_peer(_PRIMARY_AND_FB1_DIE)) as wire: + router: Final = _router(wire, ("primary", "fb1", "fb2"), fallbacks=[{"primary": ["fb1", "fb2"]}], max_fallbacks=1) + with pytest.raises(litellm.APIConnectionError, match="overloaded"): + _stream(client, router) + assert _deployments_hit(wire) == ("primary", "fb1") + + +def test_called_inside_a_running_loop() -> None: + with wire_server(_peer(_PRIMARY_DIES)) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + + async def inside_a_loop() -> _Streamed: + return _stream_sync(router) + + streamed: Final = asyncio.run(inside_a_loop()) + assert streamed.text == "answered by the backup", streamed + assert _deployments_hit(wire) == ("primary", "backup") + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_concurrent_burst(client: str) -> None: + with wire_server(_peer(_PRIMARY_DIES)) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + if client == "async": + + async def burst() -> tuple[_Streamed, ...]: + return tuple(await asyncio.gather(*(_stream_async(router) for _ in range(_BURST)))) + + streamed: tuple[_Streamed, ...] = asyncio.run(burst()) + else: + with ThreadPoolExecutor(max_workers=_BURST) as pool: + streamed = tuple(pool.map(lambda _: _stream_sync(router), range(_BURST))) + assert [item.text for item in streamed] == ["answered by the backup"] * _BURST, streamed + hit: Final = _deployments_hit(wire) + assert (hit.count("primary"), hit.count("backup"), len(hit)) == (_BURST, _BURST, 2 * _BURST), hit + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_primary_drops_the_connection_before_content(client: str) -> None: + with wire_server(_peer({"primary": _DROPS_BEFORE_CONTENT, "backup": _serves("answered by the backup")})) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + streamed: Final = _stream(client, router) + assert streamed.text == "answered by the backup", streamed + assert _deployments_hit(wire) == ("primary", "backup") diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 461b0049af3..8f3133dcaa4 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -3650,7 +3650,11 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste } return chunk - with patch.object(router, "function_with_fallbacks", return_value=NestedFallbackStream()): + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=NestedFallbackStream()), + ): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3667,6 +3671,126 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste assert result._hidden_params["model_id"] == "served-deployment" +def test_completion_streaming_fallback_resumes_chain_without_retrying_primary(): + class FailingStream(CustomStreamWrapper): + def __init__(self, model: str): + super().__init__( + completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock() + ) + + def __iter__(self): + return self + + def __next__(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}"}}])] + ) + + def __next__(self): + return next(self._chunks) + + router = 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"}}, + ], + fallbacks=[{"primary": ["backup"]}], + num_retries=0, + ) + primary_calls: Final = iter(range(2)) + + def fake_completion(**kwargs): + model_group: Final = kwargs["metadata"]["model_group"] + if model_group == "backup": + return OkStream(kwargs["model"]) + if next(primary_calls) > 0: + raise RuntimeError("primary group retried") + return FailingStream(kwargs["model"]) + + with patch("litellm.completion", side_effect=fake_completion) as provider_calls: + response: Final = router.completion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True) + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response) + + assert content == "ok-from-openai/backup-model" + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == [ + "primary", + "backup", + ] + + +def test_completion_mid_stream_fallback_walks_every_entry_of_the_configured_list(): + class FailingStream(CustomStreamWrapper): + def __init__(self, model: str): + super().__init__( + completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock() + ) + + def __iter__(self): + return self + + def __next__(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}"}}])] + ) + + def __next__(self): + return next(self._chunks) + + def fake_completion(**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.completion", side_effect=fake_completion) as provider_calls: + response: Final = router.completion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True) + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response if chunk is not None) + + assert content == "ok-from-openai/fb2-model" + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == [ + "primary", + "fb1", + "fb2", + ] + + @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. @@ -3828,7 +3952,11 @@ def test_completion_streaming_iterator_adopts_fallback_response_headers(): def __iter__(self): return iter([]) - with patch.object(router, "function_with_fallbacks", return_value=FallbackStream()): + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=FallbackStream()), + ): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3895,8 +4023,8 @@ def test_completion_streaming_iterator_fallback_on_429(): with patch.object( router, - "function_with_fallbacks", - return_value=mock_fallback_response, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=mock_fallback_response), ) as mock_fallback: result = router._completion_streaming_iterator( model_response=mock_response, @@ -3906,12 +4034,12 @@ def test_completion_streaming_iterator_fallback_on_429(): collected_chunks = list(result) - assert mock_fallback.called - call_kwargs = mock_fallback.call_args + mock_fallback.assert_awaited_once() + call_kwargs = mock_fallback.await_args.kwargs["kwargs"] # Pre-first-chunk: should use original messages, no continuation prompt - assert call_kwargs.kwargs.get("messages") == messages + assert call_kwargs.get("messages") == messages # Verify original_function is _completion (sync) - assert call_kwargs.kwargs.get("original_function") == router._completion + assert call_kwargs.get("original_function") == router._completion def test_completion_streaming_iterator_preserves_hidden_params():