mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 4fa01f9702 into 5724117116
This commit is contained in:
commit
729b28643a
10 changed files with 785 additions and 49 deletions
|
|
@ -1134,6 +1134,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
|
||||
|
|
@ -1146,6 +1147,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:
|
||||
|
|
|
|||
|
|
@ -250,6 +250,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(
|
||||
|
|
@ -1804,9 +1805,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)
|
||||
|
|
@ -2009,9 +2008,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:
|
||||
|
|
@ -2075,9 +2072,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)
|
||||
|
|
@ -2259,6 +2254,35 @@ class CustomStreamWrapper:
|
|||
recover_error,
|
||||
)
|
||||
|
||||
# Delta fields a text-only prefill continuation cannot carry, so a stream
|
||||
# that emitted any of them is not eligible for mid-stream continuation.
|
||||
_CONTINUATION_DISQUALIFYING_DELTA_FIELDS: Final = (
|
||||
"tool_calls",
|
||||
"function_call",
|
||||
"reasoning_content",
|
||||
"thinking_blocks",
|
||||
"reasoning_items",
|
||||
"audio",
|
||||
"images",
|
||||
"annotations",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _delta_disqualifies_continuation(cls, delta: object) -> bool:
|
||||
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:
|
||||
# Shared by the sync, async, and non-aiohttp iteration sites so answer
|
||||
# text and the disqualifying-content latch stay in step across all three.
|
||||
get: Final = getattr(delta, "get", None)
|
||||
content: Final = get("content") if callable(get) else None
|
||||
self.response_uptil_now += content if isinstance(content, str) else ""
|
||||
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__.
|
||||
|
|
@ -2330,6 +2354,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
|
||||
|
|
|
|||
|
|
@ -2753,7 +2753,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2788,7 +2788,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2823,7 +2823,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2858,7 +2858,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2893,7 +2893,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2928,7 +2928,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -3716,7 +3716,7 @@
|
|||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -15088,7 +15088,7 @@
|
|||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -19346,7 +19346,7 @@
|
|||
"output_cost_per_token": 2.5000010000000002e-05,
|
||||
"output_dbu_cost_per_token": 0.000357143,
|
||||
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -19561,7 +19561,7 @@
|
|||
"output_cost_per_token": 1.5000020000000002e-05,
|
||||
"output_dbu_cost_per_token": 0.000214286,
|
||||
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -41877,7 +41877,7 @@
|
|||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
"source": "https://openrouter.ai/api/v1/models",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -50091,7 +50091,7 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_batches": 7.5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -57800,7 +57800,7 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_batches": 7.5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -78639,7 +78639,7 @@
|
|||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ from collections.abc import (
|
|||
from datetime import datetime, timezone
|
||||
from functools import lru_cache, partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias, TypeVar, Union, cast
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
|
|
@ -473,7 +473,17 @@ def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream])
|
|||
|
||||
_NO_SESSION_KWARGS: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType({})
|
||||
_SESSION_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_UNCONSTRAINED_RESPONSE_FORMAT: Final[Mapping[str, str]] = MappingProxyType({"type": "text"})
|
||||
_SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str])
|
||||
_CONTENT_BLOCKS_ADAPTER: Final = TypeAdapter(list[dict[str, object]])
|
||||
|
||||
|
||||
def _extend_trailing_text_block(blocks: Sequence[Mapping[str, object]], text: str) -> Sequence[Mapping[str, object]]:
|
||||
"""Continue the last text block in place so the prefill reads exactly as the text the caller already received."""
|
||||
last: Final = blocks[-1] if blocks else None
|
||||
if last is not None and last.get("type") == "text" and isinstance(last_text := last.get("text"), str):
|
||||
return [*blocks[:-1], {**last, "text": f"{last_text}{text}"}]
|
||||
return [*blocks, {"type": "text", "text": text}]
|
||||
|
||||
|
||||
def _as_retry_skipped_deployment_ids(value: object) -> tuple[str, ...]:
|
||||
|
|
@ -832,6 +842,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,
|
||||
fallback_budget_check: FallbackBudgetCheck | None = None,
|
||||
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
|
||||
|
|
@ -873,6 +884,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).
|
||||
fallback_budget_check (Optional[FallbackBudgetCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects as over budget is skipped. Defaults to None (budget is not re-checked on fallback).
|
||||
Returns:
|
||||
|
|
@ -1061,6 +1073,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)
|
||||
|
|
@ -1255,6 +1268,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:
|
||||
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
|
||||
ContinuationPrefillDeploymentCheck,
|
||||
)
|
||||
|
||||
# Process-global on purpose: Router.discard() must not drop a filter another router still needs.
|
||||
litellm.logging_callback_manager.add_litellm_callback(ContinuationPrefillDeploymentCheck())
|
||||
|
||||
def discard(self):
|
||||
"""
|
||||
Pseudo-destructor to be invoked to clean up global data structures when router is no longer used.
|
||||
|
|
@ -2543,7 +2564,7 @@ class Router:
|
|||
try:
|
||||
# Capture kwargs before deployment selection so the streaming
|
||||
# fallback iterator can re-dispatch with the original model group.
|
||||
input_kwargs_for_streaming_fallback: Final = kwargs.copy()
|
||||
input_kwargs_for_streaming_fallback: Final[dict[str, Any]] = kwargs.copy()
|
||||
input_kwargs_for_streaming_fallback["model"] = model
|
||||
|
||||
# pick the one that is available (lowest TPM/RPM)
|
||||
|
|
@ -2884,7 +2905,7 @@ class Router:
|
|||
self,
|
||||
model_response: CustomStreamWrapper,
|
||||
messages: list[dict[str, str]],
|
||||
initial_kwargs: dict,
|
||||
initial_kwargs: dict[str, Any],
|
||||
deployment_slot: contextlib.AsyncExitStack | None = None,
|
||||
) -> CustomStreamWrapper:
|
||||
"""
|
||||
|
|
@ -2943,12 +2964,15 @@ 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)
|
||||
):
|
||||
if e.original_exception is not None:
|
||||
raise e.original_exception from e
|
||||
raise
|
||||
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
|
||||
)
|
||||
if committed and not continue_after_content:
|
||||
self._raise_original_mid_stream_error(e)
|
||||
|
||||
from litellm.main import stream_chunk_builder
|
||||
|
||||
|
|
@ -2968,7 +2992,30 @@ class Router:
|
|||
"content_policy_fallbacks", self.content_policy_fallbacks
|
||||
)
|
||||
initial_kwargs["original_function"] = self._acompletion
|
||||
initial_kwargs["messages"] = messages
|
||||
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
|
||||
MID_STREAM_CONTINUATION_KWARG,
|
||||
MID_STREAM_CONTINUATION_MARKER,
|
||||
ContinuationOutputRules,
|
||||
)
|
||||
|
||||
continuation_rules: Final = (
|
||||
ContinuationOutputRules(e.generated_content, model_group) if continue_after_content else None
|
||||
)
|
||||
if continue_after_content:
|
||||
emitted_tokens: Final = int(
|
||||
getattr(complete_response_object_usage, "completion_tokens", 0) or 0
|
||||
)
|
||||
reduced_ceilings: Final = self._continuation_output_ceilings(initial_kwargs, emitted_tokens)
|
||||
continuation_messages: Final = self._build_completion_continuation_input(
|
||||
messages, e.generated_content
|
||||
)
|
||||
if reduced_ceilings is None or continuation_messages is None:
|
||||
self._raise_original_mid_stream_error(e)
|
||||
initial_kwargs.update(reduced_ceilings)
|
||||
initial_kwargs["messages"] = continuation_messages
|
||||
initial_kwargs[MID_STREAM_CONTINUATION_KWARG] = MID_STREAM_CONTINUATION_MARKER
|
||||
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,
|
||||
|
|
@ -3002,6 +3049,8 @@ class Router:
|
|||
and hasattr(fallback_item, "usage")
|
||||
):
|
||||
self._combine_fallback_usage(fallback_item, complete_response_object_usage)
|
||||
if continuation_rules is not None and isinstance(fallback_item, ModelResponseStream):
|
||||
continuation_rules.observe(fallback_item)
|
||||
yield fallback_item
|
||||
else:
|
||||
# If fallback returns a non-streaming response, yield None
|
||||
|
|
@ -3158,6 +3207,70 @@ class Router:
|
|||
total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _raise_original_mid_stream_error(e: "MidStreamFallbackError") -> NoReturn:
|
||||
"""Surface the provider error the stream wrapper carried instead of the internal MidStreamFallbackError."""
|
||||
if e.original_exception is not None:
|
||||
raise e.original_exception from e
|
||||
raise e
|
||||
|
||||
def _mid_stream_continuation_eligible(
|
||||
self,
|
||||
e: "MidStreamFallbackError",
|
||||
request_kwargs: Mapping[str, object],
|
||||
) -> bool:
|
||||
"""Whether a stream that broke after plain assistant text may continue via prefill."""
|
||||
if not self.enable_mid_stream_fallback_continuation:
|
||||
return False
|
||||
if not e.generated_content or e.emitted_disqualifying_content:
|
||||
return False
|
||||
response_format: Final = request_kwargs.get("response_format")
|
||||
if response_format is not None and response_format != _UNCONSTRAINED_RESPONSE_FORMAT:
|
||||
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 _continuation_output_ceilings(
|
||||
request_kwargs: Mapping[str, object],
|
||||
emitted_tokens: int,
|
||||
) -> Mapping[str, int] | None:
|
||||
"""The caller's output ceilings minus the tokens already emitted, or None once nothing is left."""
|
||||
ceilings: Final = MappingProxyType(
|
||||
{
|
||||
key: value - emitted_tokens
|
||||
for key in ("max_tokens", "max_completion_tokens")
|
||||
if isinstance(value := request_kwargs.get(key), int)
|
||||
}
|
||||
)
|
||||
if ceilings and min(ceilings.values()) <= 0:
|
||||
return None
|
||||
return ceilings
|
||||
|
||||
@staticmethod
|
||||
def _build_completion_continuation_input(
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
generated_content: str,
|
||||
) -> Sequence[Mapping[str, object]] | None:
|
||||
"""Append the partial output as an assistant prefill, or extend a trailing assistant turn in place so the
|
||||
request never ends in two assistant messages. None when that turn's content has a shape this cannot extend."""
|
||||
last: Final = messages[-1] if messages else None
|
||||
if last is None or last.get("role") != "assistant":
|
||||
return [*messages, {"role": "assistant", "content": generated_content, "prefix": True}]
|
||||
content: Final = last.get("content")
|
||||
if content is None or isinstance(content, str):
|
||||
return [*messages[:-1], {**last, "content": f"{content or ''}{generated_content}", "prefix": True}]
|
||||
try:
|
||||
blocks: Final = _CONTENT_BLOCKS_ADAPTER.validate_python(content)
|
||||
except ValidationError:
|
||||
return None
|
||||
extended: Final = _extend_trailing_text_block(blocks, generated_content)
|
||||
return [*messages[:-1], {**last, "content": extended, "prefix": True}]
|
||||
|
||||
@staticmethod
|
||||
def _build_responses_continuation_input(
|
||||
input_val: Union[str, "ResponseInputParam"] | None,
|
||||
|
|
@ -3509,7 +3622,7 @@ class Router:
|
|||
self,
|
||||
model_response: CustomStreamWrapper,
|
||||
messages: list[dict[str, str]],
|
||||
initial_kwargs: dict,
|
||||
initial_kwargs: dict[str, Any],
|
||||
) -> CustomStreamWrapper:
|
||||
"""
|
||||
Sync equivalent of _acompletion_streaming_iterator.
|
||||
|
|
@ -3672,7 +3785,7 @@ class Router:
|
|||
deployment = None
|
||||
_timeout_debug_deployment_dict = {} # this is a temporary dict to debug timeout issues
|
||||
try:
|
||||
input_kwargs_for_streaming_fallback: Final = kwargs.copy()
|
||||
input_kwargs_for_streaming_fallback: Final[dict[str, Any]] = kwargs.copy()
|
||||
input_kwargs_for_streaming_fallback["model"] = model
|
||||
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
|
|
@ -3744,8 +3857,13 @@ class Router:
|
|||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
|
||||
MID_STREAM_CONTINUATION_KWARG,
|
||||
)
|
||||
|
||||
input_kwargs.pop("silent_model", None)
|
||||
input_kwargs.pop("include_fallback_errors", None)
|
||||
input_kwargs.pop(MID_STREAM_CONTINUATION_KWARG, None)
|
||||
|
||||
logging_obj: Final[LiteLLMLogging | None] = kwargs.get("litellm_logging_obj", None)
|
||||
|
||||
|
|
@ -12186,6 +12304,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,88 @@
|
|||
"""Keeps a mid-stream continuation on deployments whose model supports assistant prefill."""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.litellm_core_utils.rules import Rules
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
from litellm.utils import supports_assistant_prefill
|
||||
|
||||
MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation"
|
||||
|
||||
|
||||
class _ContinuationMarker:
|
||||
"""Only the router constructs one, so a JSON request body cannot forge the continuation flag."""
|
||||
|
||||
|
||||
MID_STREAM_CONTINUATION_MARKER: Final = _ContinuationMarker()
|
||||
|
||||
_STR_KEYED_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _declared_prefill_support(deployment_map: Mapping[str, object]) -> bool | None:
|
||||
try:
|
||||
model_info: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment_map.get("model_info"))
|
||||
except ValidationError:
|
||||
return None
|
||||
declared: Final = model_info.get("supports_assistant_prefill")
|
||||
return declared if isinstance(declared, bool) else None
|
||||
|
||||
|
||||
def _deployment_supports_prefill(deployment: object) -> bool:
|
||||
try:
|
||||
deployment_map: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment)
|
||||
except ValidationError:
|
||||
return False
|
||||
declared: Final = _declared_prefill_support(deployment_map)
|
||||
if declared is not None:
|
||||
return declared
|
||||
try:
|
||||
litellm_params: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment_map.get("litellm_params"))
|
||||
except ValidationError:
|
||||
return False
|
||||
model: Final = litellm_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]], # mutable-ok: CustomLogger deployment-list contract
|
||||
messages: Sequence[AllMessageValues] | None,
|
||||
request_kwargs: Mapping[str, object] | None = None,
|
||||
parent_otel_span: Span | None = None,
|
||||
) -> list[dict[str, object]]: # mutable-ok: returns a mutable deployment list
|
||||
marker: Final = (request_kwargs or {}).get(MID_STREAM_CONTINUATION_KWARG)
|
||||
if not isinstance(marker, _ContinuationMarker):
|
||||
return healthy_deployments
|
||||
eligible: Final = (deployment for deployment in healthy_deployments if _deployment_supports_prefill(deployment))
|
||||
return list(eligible)
|
||||
|
||||
|
||||
def _delta_text(chunk: ModelResponseStream) -> str:
|
||||
if not chunk.choices:
|
||||
return ""
|
||||
delta: Final = chunk.choices[0].delta
|
||||
content: Final = delta.content if delta is not None else None
|
||||
return content if isinstance(content, str) else ""
|
||||
|
||||
|
||||
class ContinuationOutputRules:
|
||||
"""Runs litellm.post_call_rules over the primary's text plus the continuation's, since each stream wrapper only sees its own fragment."""
|
||||
|
||||
def __init__(self, emitted_text: str, model: str) -> None:
|
||||
self._text = emitted_text
|
||||
self._model: Final = model
|
||||
self._rules: Final = Rules()
|
||||
|
||||
def observe(self, chunk: ModelResponseStream) -> None:
|
||||
if not litellm.post_call_rules:
|
||||
return
|
||||
self._text += _delta_text(chunk)
|
||||
self._rules.post_call_rules(input=self._text, model=self._model)
|
||||
|
|
@ -2969,6 +2969,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_thinking_cache_preservation(model: str, custom_llm_provider: str | None = None) -> bool:
|
||||
return _supports_factory(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -2753,7 +2753,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2788,7 +2788,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2823,7 +2823,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2858,7 +2858,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2893,7 +2893,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -2928,7 +2928,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -3716,7 +3716,7 @@
|
|||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -15088,7 +15088,7 @@
|
|||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -19346,7 +19346,7 @@
|
|||
"output_cost_per_token": 2.5000010000000002e-05,
|
||||
"output_dbu_cost_per_token": 0.000357143,
|
||||
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -19561,7 +19561,7 @@
|
|||
"output_cost_per_token": 1.5000020000000002e-05,
|
||||
"output_dbu_cost_per_token": 0.000214286,
|
||||
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_legacy_thinking": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -41877,7 +41877,7 @@
|
|||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
"source": "https://openrouter.ai/api/v1/models",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -50091,7 +50091,7 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_batches": 7.5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -57800,7 +57800,7 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_batches": 7.5e-06,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
@ -78639,7 +78639,7 @@
|
|||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
|
|
|
|||
|
|
@ -4938,6 +4938,59 @@ class TestStableStreamingResponseId:
|
|||
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",
|
||||
"reasoning_content",
|
||||
"thinking_blocks",
|
||||
"reasoning_items",
|
||||
"audio",
|
||||
"images",
|
||||
"annotations",
|
||||
],
|
||||
)
|
||||
def test_disqualifying_fields_flagged(self, field):
|
||||
# reasoning_content included: it reaches the caller as visible reasoning a
|
||||
# text-only prefill cannot carry, so a reasoning fallback would re-derive it
|
||||
assert CustomStreamWrapper._delta_disqualifies_continuation({field: [{"x": 1}]}) is True
|
||||
|
||||
@pytest.mark.parametrize("delta", [{"content": "hi"}, {}, {"role": "assistant"}])
|
||||
def test_plain_text_not_flagged(self, delta):
|
||||
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}]})
|
||||
wrapper._accumulate_streamed_delta({"content": "there"})
|
||||
|
||||
assert wrapper.response_uptil_now == "Hithere"
|
||||
assert wrapper._emitted_disqualifying_content is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_stream_without_usage_counts_tokens_off_the_event_loop():
|
||||
from tests.large_text import text
|
||||
|
|
|
|||
|
|
@ -0,0 +1,123 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
|
||||
MID_STREAM_CONTINUATION_KWARG,
|
||||
MID_STREAM_CONTINUATION_MARKER,
|
||||
ContinuationPrefillDeploymentCheck,
|
||||
_deployment_supports_prefill,
|
||||
)
|
||||
|
||||
PREFILL_MODEL = "anthropic/prefill-capable-test-model"
|
||||
NON_PREFILL_MODEL = "openai/prefill-unknown-test-model"
|
||||
UNMAPPED_MODEL = "openai/unmapped-test-model"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _synthetic_cost_map_entries() -> None:
|
||||
litellm.register_model(
|
||||
{
|
||||
PREFILL_MODEL: {"litellm_provider": "anthropic", "mode": "chat", "supports_assistant_prefill": True},
|
||||
NON_PREFILL_MODEL: {"litellm_provider": "openai", "mode": "chat"},
|
||||
},
|
||||
persist_across_reloads=False,
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
def _deployment(model: str, dep_id: str) -> dict[str, object]:
|
||||
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
|
||||
assert _deployment_supports_prefill({"litellm_params": {"model": PREFILL_MODEL}}) is True
|
||||
assert _deployment_supports_prefill({"litellm_params": {"model": PREFILL_MODEL}, "model_info": "bogus"}) is True
|
||||
|
||||
|
||||
def test_deployment_model_info_override_wins_over_cost_map():
|
||||
# model_info True opts in a model that is not in the cost map
|
||||
assert (
|
||||
_deployment_supports_prefill(
|
||||
{"litellm_params": {"model": "vendor/custom-model"}, "model_info": {"supports_assistant_prefill": True}}
|
||||
)
|
||||
is True
|
||||
)
|
||||
# model_info False opts out a model the cost map would otherwise allow
|
||||
assert (
|
||||
_deployment_supports_prefill(
|
||||
{"litellm_params": {"model": PREFILL_MODEL}, "model_info": {"supports_assistant_prefill": False}}
|
||||
)
|
||||
is False
|
||||
)
|
||||
# model_info without the key falls through to the cost map
|
||||
assert _deployment_supports_prefill(_deployment(PREFILL_MODEL, "z")) is True
|
||||
|
||||
|
||||
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_ignores_forged_client_flag():
|
||||
"""A client cannot steer routing: a plain truthy value under the marker key
|
||||
(which the proxy could forward from the request body) is not the internal
|
||||
sentinel, so the filter leaves the deployment list untouched."""
|
||||
check = ContinuationPrefillDeploymentCheck()
|
||||
deployments = [_deployment(PREFILL_MODEL, "a"), _deployment(NON_PREFILL_MODEL, "b")]
|
||||
|
||||
for forged in (True, "true", 1, {"any": "json"}):
|
||||
result = await check.async_filter_deployments(
|
||||
model="group",
|
||||
healthy_deployments=deployments,
|
||||
messages=None,
|
||||
request_kwargs={MID_STREAM_CONTINUATION_KWARG: forged},
|
||||
)
|
||||
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: MID_STREAM_CONTINUATION_MARKER},
|
||||
)
|
||||
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(UNMAPPED_MODEL, "c")]
|
||||
|
||||
result = await check.async_filter_deployments(
|
||||
model="group",
|
||||
healthy_deployments=deployments,
|
||||
messages=None,
|
||||
request_kwargs={MID_STREAM_CONTINUATION_KWARG: MID_STREAM_CONTINUATION_MARKER},
|
||||
)
|
||||
assert result == []
|
||||
|
|
@ -2662,6 +2662,318 @@ async def test_acompletion_streaming_iterator():
|
|||
print("\n=== All tests passed! ===")
|
||||
|
||||
|
||||
def _make_midstream_source(error, chunks=None):
|
||||
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,
|
||||
MID_STREAM_CONTINUATION_MARKER,
|
||||
)
|
||||
|
||||
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 MID_STREAM_CONTINUATION_MARKER
|
||||
assert passed_kwargs["messages"][-1] == {"role": "assistant", "content": "Hello", "prefix": True}
|
||||
assert fallback_chunk in collected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"continuation_text,joined_output_trips_rule",
|
||||
[("6789", True), ("6788", False)],
|
||||
ids=["rule_trips_only_on_the_joined_text", "clean_continuation_streams"],
|
||||
)
|
||||
async def test_acompletion_streaming_iterator_continuation_runs_post_call_rules_on_the_joined_text(
|
||||
monkeypatch, continuation_text, joined_output_trips_rule
|
||||
):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
monkeypatch.setattr(litellm, "post_call_rules", [lambda output: "123-45-6789" not in output])
|
||||
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="123-45-",
|
||||
is_pre_first_chunk=False, emitted_disqualifying_content=False,
|
||||
)
|
||||
fallback_chunk = litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": continuation_text}}])
|
||||
|
||||
with patch.object(
|
||||
router,
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
new=AsyncMock(return_value=_FakeFallbackStream(fallback_chunk)),
|
||||
):
|
||||
result = await router._acompletion_streaming_iterator(
|
||||
model_response=_make_midstream_source(error),
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
initial_kwargs={"model": "gpt-4", "stream": True},
|
||||
)
|
||||
if joined_output_trips_rule:
|
||||
with pytest.raises(litellm.APIResponseValidationError):
|
||||
async for _ in result:
|
||||
pass
|
||||
else:
|
||||
assert fallback_chunk in [chunk async for chunk in result]
|
||||
|
||||
|
||||
@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"}),
|
||||
({"emitted_disqualifying_content": False}, {"merge_reasoning_content_in_choices": True}),
|
||||
],
|
||||
ids=["tool_or_thinking_emitted", "json_mode", "forced_tool_choice", "merged_reasoning"],
|
||||
)
|
||||
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}
|
||||
|
||||
|
||||
def test_build_completion_continuation_input_folds_into_existing_prefill():
|
||||
"""A nested break must not leave two trailing assistant turns: the new partial
|
||||
folds into the prior prefill so a non-merging provider still gets one."""
|
||||
once = litellm.Router._build_completion_continuation_input([{"role": "user", "content": "hi"}], "part one ")
|
||||
twice = litellm.Router._build_completion_continuation_input(list(once), "part two")
|
||||
assert [m["role"] for m in twice] == ["user", "assistant"]
|
||||
assert twice[-1] == {"role": "assistant", "content": "part one part two", "prefix": True}
|
||||
|
||||
|
||||
def test_build_completion_continuation_input_folds_into_trailing_plain_assistant_turn():
|
||||
"""A request that already ends in an assistant message takes the partial as its prefill
|
||||
instead of gaining a second assistant turn, which prefill providers reject."""
|
||||
messages = [{"role": "user", "content": "hi"}, {"role": "assistant", "content": "Sure, "}]
|
||||
built = litellm.Router._build_completion_continuation_input(messages, "here it is")
|
||||
assert [m["role"] for m in built] == ["user", "assistant"]
|
||||
assert built[-1] == {"role": "assistant", "content": "Sure, here it is", "prefix": True}
|
||||
|
||||
|
||||
def test_build_completion_continuation_input_keeps_structured_assistant_content():
|
||||
"""Content blocks on a trailing assistant turn stay blocks: the partial continues the last text
|
||||
block in place (a block boundary would let a provider drop the space between "Sure, " and "here"),
|
||||
lands as a new text block only after a non-text block, and a content shape that cannot be extended declines."""
|
||||
from litellm.router import _extend_trailing_text_block
|
||||
|
||||
messages = [{"role": "user", "content": "hi"}, {"role": "assistant", "content": [{"type": "text", "text": "Sure, "}]}]
|
||||
built = litellm.Router._build_completion_continuation_input(messages, "here it is")
|
||||
assert built is not None
|
||||
assert [m["role"] for m in built] == ["user", "assistant"]
|
||||
assert built[-1] == {"role": "assistant", "content": [{"type": "text", "text": "Sure, here it is"}], "prefix": True}
|
||||
image_block = {"type": "image_url", "image_url": {"url": "https://example.test/a.png"}}
|
||||
assert _extend_trailing_text_block([image_block], "here it is") == [image_block, {"type": "text", "text": "here it is"}]
|
||||
assert _extend_trailing_text_block([], "here it is") == [{"type": "text", "text": "here it is"}]
|
||||
assert litellm.Router._build_completion_continuation_input([{"role": "assistant", "content": 42}], "x") is None
|
||||
|
||||
|
||||
def test_continuation_output_ceilings_reduces_by_emitted_tokens():
|
||||
"""A continuation must complete within the caller's original allowance, so each
|
||||
output ceiling is reduced by the tokens already emitted."""
|
||||
assert litellm.Router._continuation_output_ceilings({"max_tokens": 100}, 30) == {"max_tokens": 70}
|
||||
assert litellm.Router._continuation_output_ceilings({"max_tokens": 100, "max_completion_tokens": 40}, 25) == {
|
||||
"max_tokens": 75,
|
||||
"max_completion_tokens": 15,
|
||||
}
|
||||
# no ceiling configured -> nothing to reduce, continuation proceeds as before
|
||||
assert litellm.Router._continuation_output_ceilings({}, 50) == {}
|
||||
|
||||
|
||||
def test_continuation_output_ceilings_none_when_allowance_exhausted():
|
||||
"""When the emitted tokens already meet or exceed a ceiling, there is no budget
|
||||
left to continue, so the helper signals a decline rather than a fresh allowance."""
|
||||
assert litellm.Router._continuation_output_ceilings({"max_tokens": 20}, 20) is None
|
||||
assert litellm.Router._continuation_output_ceilings({"max_tokens": 20}, 25) is None
|
||||
assert litellm.Router._continuation_output_ceilings({"max_tokens": 100, "max_completion_tokens": 10}, 10) is None
|
||||
|
||||
|
||||
def test_mid_stream_continuation_eligible_allows_text_response_format():
|
||||
"""response_format={"type": "text"} is the unconstrained default and must stay
|
||||
eligible, unlike json_object / json_schema."""
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "k"}}],
|
||||
enable_mid_stream_fallback_continuation=True,
|
||||
)
|
||||
e = MidStreamFallbackError(
|
||||
message="boom", model="gpt-4", llm_provider="openai", generated_content="Hello",
|
||||
is_pre_first_chunk=False, emitted_disqualifying_content=False,
|
||||
)
|
||||
assert router._mid_stream_continuation_eligible(e=e, request_kwargs={"response_format": {"type": "text"}}) is True
|
||||
assert router._mid_stream_continuation_eligible(e=e, request_kwargs={"response_format": {"type": "json_object"}}) is False
|
||||
|
||||
|
||||
def test_raise_original_mid_stream_error_surfaces_the_provider_exception():
|
||||
from litellm.exceptions import MidStreamFallbackError, RateLimitError
|
||||
|
||||
provider_error = RateLimitError(message="rate limited", llm_provider="openai", model="gpt-4")
|
||||
wrapped = MidStreamFallbackError(
|
||||
message="rate limited", model="gpt-4", llm_provider="openai",
|
||||
original_exception=provider_error, generated_content="Hello",
|
||||
)
|
||||
with pytest.raises(RateLimitError) as raised:
|
||||
litellm.Router._raise_original_mid_stream_error(wrapped)
|
||||
assert raised.value is provider_error
|
||||
assert raised.value.__cause__ is wrapped
|
||||
|
||||
bare = MidStreamFallbackError(message="boom", model="gpt-4", llm_provider="openai", generated_content="Hello")
|
||||
with pytest.raises(MidStreamFallbackError) as bare_raised:
|
||||
litellm.Router._raise_original_mid_stream_error(bare)
|
||||
assert bare_raised.value is bare
|
||||
|
||||
|
||||
@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