diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 3bae8a95ef6..10ebccf00ef 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -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: diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 9d33f86d841..9705983ad76 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a935ffdb2dd..4c01a0b3077 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 24554e61516..0b75a422a45 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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", ] 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..2bde8ca920d --- /dev/null +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -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) diff --git a/litellm/utils.py b/litellm/utils.py index 35fea4f5f4a..d437d9e28d0 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a935ffdb2dd..4c01a0b3077 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index f88e082d577..03a2fc5fa43 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -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 diff --git a/tests/unit/router_utils/pre_call_checks/test_continuation_prefill_check.py b/tests/unit/router_utils/pre_call_checks/test_continuation_prefill_check.py new file mode 100644 index 00000000000..77a3db58336 --- /dev/null +++ b/tests/unit/router_utils/pre_call_checks/test_continuation_prefill_check.py @@ -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 == [] diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 96dddf15869..86778b54158 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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