fix(router): sync streaming fallback re-ran the failing group instead of falling back

_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
This commit is contained in:
JingHao-Leon 2026-10-01 11:15:09 +08:00
parent c168199e33
commit aeb88738e8
2 changed files with 108 additions and 8 deletions

View file

@ -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__"):

View file

@ -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,