mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
30f33a949b
commit
974da4c2c9
8 changed files with 476 additions and 13 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue