diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 3f22a4b2dcd..dfda31b11c0 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -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: diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 6a5a8832cc6..f62a23d2c42 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 8865543badd..a10acc63fe8 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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", ] diff --git a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py new file mode 100644 index 00000000000..75d1ecb225c --- /dev/null +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -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)] diff --git a/litellm/utils.py b/litellm/utils.py index 7732cd88cb5..72c4c673547 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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. diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 37e2031fdf4..f233d6e9bb2 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -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 diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py new file mode 100644 index 00000000000..28f220100fa --- /dev/null +++ b/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py @@ -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 == [] diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index b8a0d70f5bc..1a1ca08b9b2 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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