From aeb88738e8ff6c34766e7b80c211476bfceafbbb Mon Sep 17 00:00:00 2001 From: JingHao-Leon Date: Thu, 1 Oct 2026 11:15:09 +0800 Subject: [PATCH] fix(router): sync streaming fallback re-ran the failing group instead of falling back MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _completion_streaming_iterator.stream_with_fallbacks() re-entered the fallback chain via function_with_fallbacks(), which tries the original group first. The retry returns a fresh stream that succeeds at creation time (HTTP 200) and only fails while being iterated; each nested Router.completion() wraps its own stream in this same iterator, so every retry fails again during iteration and re-enters the chain — recursing until the stack runs out (measured: 478 requests to the failing group, ~116 s, then InternalServerError, with a healthy fallback configured). Mirror the async twin: hand the MidStreamFallbackError to async_function_with_fallbacks_common_utils() via run_async_function, which cools down the failed deployment and walks the fallback list directly. Sync tests that patched function_with_fallbacks are updated to the new re-entry point, plus a regression test asserting the triggering error reaches the common utils and the original group is not re-run. Fixes #43945 --- litellm/router.py | 15 +++- tests/unit/test_router/test_router.py | 101 ++++++++++++++++++++++++-- 2 files changed, 108 insertions(+), 8 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index bb118639839..26e1ff16152 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3582,11 +3582,22 @@ 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, + # Pass the MidStreamFallbackError through the common fallback utils, like the + # async twin does. Calling function_with_fallbacks() here instead re-runs the + # original (failing) group first, and because each nested Router.completion() + # wraps its stream in this same iterator, every retry fails again only when + # the caller iterates it — recursing until the stack runs out. + fallback_response = run_async_function( + router_self.async_function_with_fallbacks_common_utils, + 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/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 96dddf15869..34eb200e7cd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -3438,7 +3438,7 @@ 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", return_value=NestedFallbackStream()): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3616,7 +3616,7 @@ 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", return_value=FallbackStream()): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3683,7 +3683,7 @@ def test_completion_streaming_iterator_fallback_on_429(): with patch.object( router, - "function_with_fallbacks", + "async_function_with_fallbacks_common_utils", return_value=mock_fallback_response, ) as mock_fallback: result = router._completion_streaming_iterator( @@ -3696,10 +3696,99 @@ def test_completion_streaming_iterator_fallback_on_429(): assert mock_fallback.called call_kwargs = mock_fallback.call_args + assert mock_fallback.call_args.args[0] is rate_limit_error # Pre-first-chunk: should use original messages, no continuation prompt - assert call_kwargs.kwargs.get("messages") == messages + assert call_kwargs.kwargs.get("kwargs", {}).get("messages") == messages # Verify original_function is _completion (sync) - assert call_kwargs.kwargs.get("original_function") == router._completion + assert call_kwargs.kwargs.get("kwargs", {}).get("original_function") == router._completion + + +def test_completion_streaming_iterator_routes_mid_stream_fallback_through_common_utils(): + """Regression (#43945): the sync mid-stream fallback re-entry must hand the + MidStreamFallbackError to async_function_with_fallbacks_common_utils, like the async + twin does. Calling function_with_fallbacks() instead re-runs the original (failing) + group first, and because each nested Router.completion() wraps its own stream in this + same iterator, every retry fails again only while being iterated — recursing until + the stack runs out (measured: 478 requests to the failing group, then + InternalServerError, with a healthy fallback configured).""" + from unittest.mock import MagicMock + + from litellm.exceptions import MidStreamFallbackError + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + ) + + messages = [{"role": "user", "content": "Test"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + pre_first_chunk_error = MidStreamFallbackError( + message="upstream died before the first chunk", + model="gpt-4", + llm_provider="openai", + generated_content="", + is_pre_first_chunk=True, + ) + + class SyncIteratorImmediateError: + def __init__(self): + self.model = "gpt-4" + self.custom_llm_provider = "openai" + self.logging_obj = MagicMock() + self.chunks = [] + + def __iter__(self): + return self + + def __next__(self): + raise pre_first_chunk_error + + class FallbackStream: + def __init__(self): + self._chunks = iter( + [ + litellm.ModelResponseStream( + choices=[{"index": 0, "delta": {"content": "from the fallback"}}] + ) + ] + ) + + def __iter__(self): + return self + + def __next__(self): + return next(self._chunks) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=FallbackStream(), + ) as mock_utils: + with patch.object(router, "function_with_fallbacks") as mock_function_with_fallbacks: + result = router._completion_streaming_iterator( + model_response=SyncIteratorImmediateError(), + messages=messages, + initial_kwargs=initial_kwargs, + ) + + collected_chunks = list(result) + + assert mock_utils.called + # the triggering error must reach the common utils, so cooldowns apply and + # the walk starts from the fallback list, not the failing group + assert mock_utils.call_args.args[0] is pre_first_chunk_error + assert mock_utils.call_args.kwargs.get("kwargs", {}).get("messages") == messages + assert not mock_function_with_fallbacks.called, ( + "re-running the original group is what recurses; common utils already " + "excludes the deployment that raised" + ) + + assert len(collected_chunks) == 1 def test_completion_streaming_iterator_preserves_hidden_params(): @@ -3911,7 +4000,7 @@ def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_con mock_response = SyncIteratorNoTextChunkError() - with patch.object(router, "function_with_fallbacks") as mock_fallback: + with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback: result = router._completion_streaming_iterator( model_response=mock_response, messages=messages,