mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
c168199e33
commit
aeb88738e8
2 changed files with 108 additions and 8 deletions
|
|
@ -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__"):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue