feat(router): continue chat streams on fallback after a mid-stream break

A chat-completions stream that breaks after content has been delivered was
re-raised and the fallback deployment never ran (#40404), while the
Responses-API path already continues via a prefilled assistant turn. Add the
same for chat completions behind `enable_mid_stream_fallback_continuation`
(opt-in, async-only), so a post-content break re-enters the fallback chain with
the partial text as an assistant prefill instead of surfacing the error.

Continuation is only attempted when it is safe: the stream emitted plain
assistant text (no tool/function calls, thinking blocks, reasoning items,
audio or images), the request is not constrained output (response_format or a
forced tool_choice) and not merge-reasoning mode, and the fallback target's
model supports assistant prefill. The target check runs as a deployment
pre-call filter, so a chain with no prefill-capable deployment empties and the
original error is surfaced rather than a duplicated or rejected request being
sent. Plain reasoning_content is treated as out-of-band and does not block a
continuation, matching the Responses-API path; Anthropic thinking is excluded
because its signed thinking blocks cannot ride a text-only prefill.

CustomStreamWrapper now latches whether a disqualifying delta was streamed and
carries it on MidStreamFallbackError so the router can decide without
re-scanning chunks. Default behavior is unchanged.
This commit is contained in:
Ayush 2026-09-14 17:31:05 +05:30
parent 30f33a949b
commit 974da4c2c9
8 changed files with 476 additions and 13 deletions

View file

@ -1124,6 +1124,7 @@ class MidStreamFallbackError(ServiceUnavailableError):
num_retries: int | None = None,
generated_content: str = "",
is_pre_first_chunk: bool = False,
emitted_disqualifying_content: bool = False,
):
original_status: Final = getattr(original_exception, "status_code", None)
self.status_code = int(original_status) if original_status is not None else 503
@ -1136,6 +1137,7 @@ class MidStreamFallbackError(ServiceUnavailableError):
self.num_retries = num_retries
self.generated_content = generated_content
self.is_pre_first_chunk = is_pre_first_chunk
self.emitted_disqualifying_content = emitted_disqualifying_content
# Create a response if one wasn't provided
if response is None:

View file

@ -272,6 +272,7 @@ class CustomStreamWrapper:
self.holding_chunk = ""
self.complete_response = ""
self.response_uptil_now = ""
self._emitted_disqualifying_content = False
_model_info: Final[dict] = litellm_params.model_info or {}
_api_base: Final = get_api_base(
@ -1950,9 +1951,7 @@ class CustomStreamWrapper:
if response.choices:
choice = response.choices[0]
if isinstance(choice, StreamingChoices):
self.response_uptil_now += choice.delta.get("content", "") or ""
else:
self.response_uptil_now += ""
self._accumulate_streamed_delta(choice.delta)
self.rules.post_call_rules(input=self.response_uptil_now, model=self.model)
# HANDLE STREAM OPTIONS
self.chunks.append(response)
@ -2150,9 +2149,7 @@ class CustomStreamWrapper:
if processed_chunk.choices:
choice = processed_chunk.choices[0]
if isinstance(choice, StreamingChoices):
self.response_uptil_now += choice.delta.get("content", "") or ""
else:
self.response_uptil_now += ""
self._accumulate_streamed_delta(choice.delta)
self.rules.post_call_rules(input=self.response_uptil_now, model=self.model)
# Add mcp_list_tools to first chunk if present
if not self.sent_first_chunk and processed_chunk.choices:
@ -2216,9 +2213,7 @@ class CustomStreamWrapper:
choice = processed_chunk.choices[0]
if isinstance(choice, StreamingChoices):
self.response_uptil_now += choice.delta.get("content", "") or ""
else:
self.response_uptil_now += ""
self._accumulate_streamed_delta(choice.delta)
self.rules.post_call_rules(input=self.response_uptil_now, model=self.model)
# RETURN RESULT
self.chunks.append(processed_chunk)
@ -2395,6 +2390,41 @@ class CustomStreamWrapper:
recover_error,
)
_CONTINUATION_DISQUALIFYING_DELTA_FIELDS: Final = (
"tool_calls",
"function_call",
"thinking_blocks",
"reasoning_items",
"audio",
"images",
"annotations",
)
@classmethod
def _delta_disqualifies_continuation(cls, delta: object) -> bool:
"""
True when a streamed delta carries output a text-only prefill
continuation cannot represent: tool/function calls, signed Anthropic
thinking blocks, structured reasoning items, audio or image parts, or
annotations. Plain ``reasoning_content`` is deliberately not here - it
is out-of-band, never reaches the caller as answer text, and so does
not block a continuation (parity with the Responses-API path).
"""
get: Final = getattr(delta, "get", None)
if not callable(get):
return False
return any(get(field) for field in cls._CONTINUATION_DISQUALIFYING_DELTA_FIELDS)
def _accumulate_streamed_delta(self, delta: object) -> None:
"""Grow the running answer text and latch whether anything a
continuation cannot carry has been streamed. One home for both so the
three iteration sites (sync, async, non-aiohttp) stay in step."""
get: Final = getattr(delta, "get", None)
content: Final = get("content", "") if callable(get) else ""
self.response_uptil_now += content or ""
if not self._emitted_disqualifying_content and self._delta_disqualifies_continuation(delta):
self._emitted_disqualifying_content = True
def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn":
"""
Common error handling for both __next__ and __anext__.
@ -2466,6 +2496,7 @@ class CustomStreamWrapper:
original_exception=mapped_exception,
generated_content=self.response_uptil_now,
is_pre_first_chunk=not self.sent_first_chunk,
emitted_disqualifying_content=self._emitted_disqualifying_content,
)
@staticmethod

View file

@ -189,6 +189,10 @@ from litellm.router_utils.handle_error import (
send_llm_exception_alert,
)
from litellm.router_utils.health_state_cache import DeploymentHealthCache
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
MID_STREAM_CONTINUATION_KWARG,
ContinuationPrefillDeploymentCheck,
)
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
DeploymentAffinityCheck,
warn_on_unknown_model_group_affinity_flags,
@ -764,6 +768,7 @@ class Router:
health_check_ignore_transient_errors: bool = False,
background_health_check_model_groups: Sequence[str] | None = None,
enable_weighted_failover: bool = False,
enable_mid_stream_fallback_continuation: bool = False,
fallback_access_check: FallbackAccessCheck | None = None,
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
) -> None:
@ -802,6 +807,7 @@ class Router:
deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600.
ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error.
enable_weighted_failover (bool): When True and the routing strategy is "simple-shuffle", a retryable failure on one deployment causes the request to re-pick (weighted) across the other deployments in the same model group before any cross-group fallback runs. Bounded by `max_fallbacks`. Async-only: currently honored by `router.acompletion()` and other async entrypoints. The sync `router.completion()` path falls back to the regular fallback flow. Defaults to False.
enable_mid_stream_fallback_continuation (bool): When True, a chat-completions stream that breaks after plain assistant text has been delivered continues on a fallback deployment via assistant prefill instead of surfacing the error. Only deployments whose model supports assistant prefill are eligible, so the partial text is continued, not regenerated; if none is, the original error is surfaced. Streams that emitted tool calls, thinking blocks, audio/images, or constrained (JSON / forced tool_choice) output are never continued. Async-only. Defaults to False.
fallback_access_check (Optional[FallbackAccessCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects is skipped. Defaults to None (every configured fallback is attempted).
Returns:
Router: An instance of the litellm.Router class.
@ -984,6 +990,7 @@ class Router:
self.disable_cooldowns = disable_cooldowns
self.enable_health_check_routing = enable_health_check_routing
self.enable_weighted_failover = enable_weighted_failover
self.enable_mid_stream_fallback_continuation = enable_mid_stream_fallback_continuation
self.health_check_ignore_transient_errors = health_check_ignore_transient_errors
self.background_health_check_model_groups: frozenset[str] | None = (
frozenset(background_health_check_model_groups)
@ -1178,6 +1185,14 @@ class Router:
default_pre_call_checks: Final[OptionalPreCallChecks] = []
self.add_optional_pre_call_checks(default_pre_call_checks)
if self.enable_mid_stream_fallback_continuation:
if self.optional_callbacks is None:
self.optional_callbacks = []
if not any(isinstance(cb, ContinuationPrefillDeploymentCheck) for cb in self.optional_callbacks):
continuation_check: Final = ContinuationPrefillDeploymentCheck()
self.optional_callbacks.append(continuation_check)
litellm.logging_callback_manager.add_litellm_callback(continuation_check)
def discard(self):
"""
Pseudo-destructor to be invoked to clean up global data structures when router is no longer used.
@ -2802,9 +2817,17 @@ class Router:
with anyio.CancelScope(shield=True):
await close_model_response()
await held_slot.aclose()
if not e.is_pre_first_chunk and (
e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)
):
committed: Final = bool(
not e.is_pre_first_chunk
and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks))
)
continue_after_content: Final = committed and self._mid_stream_continuation_eligible(
e=e, request_kwargs=initial_kwargs
)
# Content already reached the caller and we cannot safely
# continue it (feature off, or tool/thinking/constrained output):
# surface the real error rather than restart into the same stream.
if committed and not continue_after_content:
if e.original_exception is not None:
raise e.original_exception from e
raise
@ -2827,7 +2850,13 @@ class Router:
"content_policy_fallbacks", self.content_policy_fallbacks
)
initial_kwargs["original_function"] = self._acompletion
initial_kwargs["messages"] = messages
if continue_after_content:
initial_kwargs["messages"] = self._build_completion_continuation_input(
messages, e.generated_content
)
initial_kwargs[MID_STREAM_CONTINUATION_KWARG] = True
else:
initial_kwargs["messages"] = messages
self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs)
fallback_response = await self.async_function_with_fallbacks_common_utils(
e=e,
@ -3017,6 +3046,51 @@ class Router:
total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0),
)
def _mid_stream_continuation_eligible(
self,
e: "MidStreamFallbackError",
request_kwargs: Mapping[str, object],
) -> bool:
"""
Whether a chat-completions stream that broke after content may be
continued on a fallback deployment via assistant prefill, instead of
re-raising. Only plain assistant text is safe: a continuation built from
``generated_content`` (text-only) cannot carry tool calls, signed
thinking blocks, audio or images, and a constrained (JSON / forced
tool_choice) or merged-reasoning response cannot be resumed from an
arbitrary cut point. The fallback target's prefill support is enforced
separately at deployment selection.
"""
if not self.enable_mid_stream_fallback_continuation:
return False
if not e.generated_content or e.emitted_disqualifying_content:
return False
# Any structured-output request (response_format, or a forced tool call)
# produces a partial that cannot be resumed from an arbitrary cut point.
if request_kwargs.get("response_format") is not None:
return False
tool_choice: Final = request_kwargs.get("tool_choice")
if tool_choice == "required" or isinstance(tool_choice, Mapping):
return False
if request_kwargs.get("merge_reasoning_content_in_choices") is True:
return False
return True
@staticmethod
def _build_completion_continuation_input(
messages: list[dict[str, str]],
generated_content: str,
) -> Sequence[Mapping[str, object]]:
"""
Append the partial assistant output as a prefill so a prefill-capable
fallback continues where the broken stream stopped instead of
regenerating text already delivered to the caller. The deployment filter
guarantees the target supports ``prefix: True`` (parity with
``_build_responses_continuation_input`` for the Responses-API path).
"""
prefill: dict[str, object] = {"role": "assistant", "content": generated_content, "prefix": True}
return [*messages, prefill]
@staticmethod
def _build_responses_continuation_input(
input_val: Union[str, "ResponseInputParam"] | None,
@ -3566,6 +3640,7 @@ class Router:
}
input_kwargs.pop("silent_model", None)
input_kwargs.pop("include_fallback_errors", None)
input_kwargs.pop(MID_STREAM_CONTINUATION_KWARG, None)
_response: Final = litellm.acompletion(**input_kwargs)
@ -11959,6 +12034,7 @@ class Router:
"retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_mid_stream_fallback_continuation",
"enable_tag_filtering",
"tag_routing_prefix",
]

View file

@ -0,0 +1,51 @@
"""
Mid-stream fallback continuation: keep the fallback on a deployment that can
actually continue a prefilled assistant message.
When a chat-completions stream breaks after content and the Router re-enters the
fallback chain to continue it (the request is marked with
``MID_STREAM_CONTINUATION_KWARG``), only a deployment whose model supports
assistant prefill can pick up the partial text without regenerating it.
Deployments that cannot are dropped, so selection lands on a
continuation-capable one. If a group has none it empties and the fallback chain
moves on, surfacing the original error rather than sending a request the target
would reject or duplicate. A request without the marker is passed through
untouched.
"""
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.types.llms.openai import AllMessageValues
from litellm.utils import supports_assistant_prefill
MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation"
_STR_KEYED_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _deployment_supports_prefill(deployment: object) -> bool:
try:
params: Final = _STR_KEYED_DICT_ADAPTER.validate_python(
_STR_KEYED_DICT_ADAPTER.validate_python(deployment).get("litellm_params")
)
except ValidationError:
return False
model: Final = params.get("model")
return isinstance(model, str) and bool(model) and supports_assistant_prefill(model=model)
class ContinuationPrefillDeploymentCheck(CustomLogger):
async def async_filter_deployments(
self,
model: str,
healthy_deployments: list[dict[str, object]],
messages: list[AllMessageValues] | None,
request_kwargs: dict[str, object] | None = None,
parent_otel_span: Span | None = None,
) -> list[dict[str, object]]:
if not (request_kwargs or {}).get(MID_STREAM_CONTINUATION_KWARG):
return healthy_deployments
return [deployment for deployment in healthy_deployments if _deployment_supports_prefill(deployment)]

View file

@ -2796,6 +2796,20 @@ def supports_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None
)
def supports_assistant_prefill(model: str, custom_llm_provider: str | None = None) -> bool:
"""
Whether the model can continue a prefilled assistant message (Anthropic's
``prefix: True`` trick and equivalents). Missing metadata reads as False,
so a mid-stream fallback continuation is only routed to a model known to
support it.
"""
return _supports_factory(
model=model,
custom_llm_provider=custom_llm_provider,
key="supports_assistant_prefill",
)
def supports_computer_use(model: str, custom_llm_provider: str | None = None) -> bool:
"""
Check if the given model supports computer use and return a boolean value.

View file

@ -4917,3 +4917,47 @@ class TestStableStreamingResponseId:
)
wrapper.response_id = "chatcmpl-from-provider"
assert wrapper.model_response_creator().id == "chatcmpl-from-provider"
class TestContinuationDisqualifiers:
"""The mid-stream continuation eligibility hinges on classifying which deltas
carry output a text-only prefill cannot represent."""
@pytest.mark.parametrize(
"field",
["tool_calls", "function_call", "thinking_blocks", "reasoning_items", "audio", "images", "annotations"],
)
def test_disqualifying_fields_flagged(self, field):
assert CustomStreamWrapper._delta_disqualifies_continuation({field: [{"x": 1}]}) is True
@pytest.mark.parametrize("delta", [{"content": "hi"}, {"reasoning_content": "thinking"}, {}, {"role": "assistant"}])
def test_plain_text_and_reasoning_content_not_flagged(self, delta):
# plain reasoning_content is out-of-band and must NOT block a continuation
assert CustomStreamWrapper._delta_disqualifies_continuation(delta) is False
def test_non_mapping_delta_is_safe(self):
assert CustomStreamWrapper._delta_disqualifies_continuation(object()) is False
def test_accumulate_grows_text_and_leaves_flag_clear_for_plain_text(self):
wrapper = object.__new__(CustomStreamWrapper)
wrapper.response_uptil_now = ""
wrapper._emitted_disqualifying_content = False
wrapper._accumulate_streamed_delta({"content": "Hel"})
wrapper._accumulate_streamed_delta({"content": "lo"})
assert wrapper.response_uptil_now == "Hello"
assert wrapper._emitted_disqualifying_content is False
def test_accumulate_latches_flag_on_disqualifying_delta(self):
wrapper = object.__new__(CustomStreamWrapper)
wrapper.response_uptil_now = ""
wrapper._emitted_disqualifying_content = False
wrapper._accumulate_streamed_delta({"content": "Hi"})
wrapper._accumulate_streamed_delta({"tool_calls": [{"index": 0}]})
# a later plain-text delta must not clear the latch
wrapper._accumulate_streamed_delta({"content": "there"})
assert wrapper.response_uptil_now == "Hithere"
assert wrapper._emitted_disqualifying_content is True

View file

@ -0,0 +1,69 @@
import pytest
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
MID_STREAM_CONTINUATION_KWARG,
ContinuationPrefillDeploymentCheck,
_deployment_supports_prefill,
)
PREFILL_MODEL = "anthropic/claude-3-opus-20240229" # supports_assistant_prefill: True in the cost map
NON_PREFILL_MODEL = "openai/gpt-4o" # capability absent -> treated as unsupported
def _deployment(model: str, dep_id: str) -> dict:
return {"litellm_params": {"model": model}, "model_info": {"id": dep_id}}
def test_deployment_supports_prefill_reads_capability():
assert _deployment_supports_prefill(_deployment(PREFILL_MODEL, "a")) is True
assert _deployment_supports_prefill(_deployment(NON_PREFILL_MODEL, "b")) is False
def test_deployment_supports_prefill_rejects_malformed_deployments():
assert _deployment_supports_prefill({}) is False
assert _deployment_supports_prefill({"litellm_params": {}}) is False
assert _deployment_supports_prefill("not-a-dict") is False
@pytest.mark.asyncio
async def test_filter_is_noop_without_continuation_marker():
"""A normal (non-continuation) request must be passed through untouched, even
if some deployments cannot prefill."""
check = ContinuationPrefillDeploymentCheck()
deployments = [_deployment(PREFILL_MODEL, "a"), _deployment(NON_PREFILL_MODEL, "b")]
for request_kwargs in ({}, None, {MID_STREAM_CONTINUATION_KWARG: False}):
result = await check.async_filter_deployments(
model="group", healthy_deployments=deployments, messages=None, request_kwargs=request_kwargs
)
assert result == deployments
@pytest.mark.asyncio
async def test_filter_keeps_only_prefill_capable_on_continuation():
check = ContinuationPrefillDeploymentCheck()
deployments = [_deployment(PREFILL_MODEL, "a"), _deployment(NON_PREFILL_MODEL, "b")]
result = await check.async_filter_deployments(
model="group",
healthy_deployments=deployments,
messages=None,
request_kwargs={MID_STREAM_CONTINUATION_KWARG: True},
)
assert [d["model_info"]["id"] for d in result] == ["a"]
@pytest.mark.asyncio
async def test_filter_empties_group_when_no_prefill_capable_deployment():
"""No prefill-capable deployment -> empty result, so the router advances the
fallback chain and ultimately surfaces the original error."""
check = ContinuationPrefillDeploymentCheck()
deployments = [_deployment(NON_PREFILL_MODEL, "b"), _deployment("openai/gpt-4.1", "c")]
result = await check.async_filter_deployments(
model="group",
healthy_deployments=deployments,
messages=None,
request_kwargs={MID_STREAM_CONTINUATION_KWARG: True},
)
assert result == []

View file

@ -2176,6 +2176,182 @@ async def test_acompletion_streaming_iterator():
print("\n=== All tests passed! ===")
def _make_midstream_source(error, chunks=None):
"""A minimal stand-in for the CustomStreamWrapper: yields one content chunk,
then raises ``error`` on the next pull. Carries the attributes
_acompletion_streaming_iterator reads."""
from unittest.mock import MagicMock
first_chunk = MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))])
class _Source:
def __init__(self):
self.index = 0
self.chunks = chunks if chunks is not None else []
self.model = "gpt-4"
self.custom_llm_provider = "openai"
self.logging_obj = MagicMock()
def __aiter__(self):
return self
async def __anext__(self):
if self.index == 0:
self.index += 1
return first_chunk
raise error
return _Source()
class _FakeFallbackStream:
def __init__(self, item):
self._item = item
self._done = False
def __aiter__(self):
return self
async def __anext__(self):
if self._done:
raise StopAsyncIteration
self._done = True
return self._item
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_continues_after_content_when_eligible():
"""Flag on + plain-text break: the router re-enters the fallback chain with an
assistant-prefill continuation and the mid-stream marker, then streams the
fallback's output instead of re-raising."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.exceptions import MidStreamFallbackError
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
MID_STREAM_CONTINUATION_KWARG,
)
router = litellm.Router(
model_list=[
{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "k1"}},
{"model_name": "backup", "litellm_params": {"model": "anthropic/claude-3-opus-20240229", "api_key": "k2"}},
],
fallbacks=[{"gpt-4": ["backup"]}],
enable_mid_stream_fallback_continuation=True,
)
error = MidStreamFallbackError(
message="Connection lost",
model="gpt-4",
llm_provider="openai",
generated_content="Hello",
is_pre_first_chunk=False,
emitted_disqualifying_content=False,
)
source = _make_midstream_source(error)
fallback_chunk = litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": "world"}}])
initial_kwargs = {"model": "gpt-4", "stream": True}
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(return_value=_FakeFallbackStream(fallback_chunk)),
) as mock_fallback:
result = await router._acompletion_streaming_iterator(
model_response=source, messages=[{"role": "user", "content": "Hi"}], initial_kwargs=initial_kwargs
)
collected = [chunk async for chunk in result]
mock_fallback.assert_awaited_once()
passed_kwargs = mock_fallback.await_args.kwargs["kwargs"]
assert passed_kwargs[MID_STREAM_CONTINUATION_KWARG] is True
assert passed_kwargs["messages"][-1] == {"role": "assistant", "content": "Hello", "prefix": True}
assert fallback_chunk in collected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error_kwargs,request_kwargs",
[
({"emitted_disqualifying_content": True}, {}),
({"emitted_disqualifying_content": False}, {"response_format": {"type": "json_object"}}),
({"emitted_disqualifying_content": False}, {"tool_choice": "required"}),
],
ids=["tool_or_thinking_emitted", "json_mode", "forced_tool_choice"],
)
async def test_acompletion_streaming_iterator_declines_ineligible_after_content(error_kwargs, request_kwargs):
"""Flag on but the break is not continuation-safe: the router re-raises and
never enters the fallback chain, so no duplicated/rejected request is sent."""
from unittest.mock import AsyncMock, patch
from litellm.exceptions import MidStreamFallbackError
router = litellm.Router(
model_list=[
{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "k1"}},
{"model_name": "backup", "litellm_params": {"model": "anthropic/claude-3-opus-20240229", "api_key": "k2"}},
],
fallbacks=[{"gpt-4": ["backup"]}],
enable_mid_stream_fallback_continuation=True,
)
error = MidStreamFallbackError(
message="boom", model="gpt-4", llm_provider="openai", generated_content="Hello",
is_pre_first_chunk=False, **error_kwargs,
)
source = _make_midstream_source(error)
initial_kwargs = {"model": "gpt-4", "stream": True, **request_kwargs}
with patch.object(router, "async_function_with_fallbacks_common_utils", new=AsyncMock()) as mock_fallback:
result = await router._acompletion_streaming_iterator(
model_response=source, messages=[{"role": "user", "content": "Hi"}], initial_kwargs=initial_kwargs
)
with pytest.raises(MidStreamFallbackError):
async for _ in result:
pass
mock_fallback.assert_not_awaited()
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_flag_off_declines_after_content():
"""Regression guard for the opt-in default: with the flag off, a post-content
break re-raises and never falls back, exactly as before this feature."""
from unittest.mock import AsyncMock, patch
from litellm.exceptions import MidStreamFallbackError
router = litellm.Router(
model_list=[
{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "k1"}},
{"model_name": "backup", "litellm_params": {"model": "anthropic/claude-3-opus-20240229", "api_key": "k2"}},
],
fallbacks=[{"gpt-4": ["backup"]}],
)
error = MidStreamFallbackError(
message="boom", model="gpt-4", llm_provider="openai", generated_content="Hello",
is_pre_first_chunk=False, emitted_disqualifying_content=False,
)
source = _make_midstream_source(error)
with patch.object(router, "async_function_with_fallbacks_common_utils", new=AsyncMock()) as mock_fallback:
result = await router._acompletion_streaming_iterator(
model_response=source, messages=[{"role": "user", "content": "Hi"}],
initial_kwargs={"model": "gpt-4", "stream": True},
)
with pytest.raises(MidStreamFallbackError):
async for _ in result:
pass
mock_fallback.assert_not_awaited()
def test_build_completion_continuation_input_appends_assistant_prefill():
messages = [{"role": "user", "content": "hi"}]
built = litellm.Router._build_completion_continuation_input(messages, "partial answer")
assert built[:-1] == messages
assert built[-1] == {"role": "assistant", "content": "partial answer", "prefix": True}
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_reraises_original_exception_when_available():
"""Async: when the mid-stream MidStreamFallbackError wraps a real provider