diff --git a/litellm/router.py b/litellm/router.py index 0fd97ebccad..d479519c18e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -473,8 +473,6 @@ 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]) -# The only response_format that leaves a stream continuation-eligible: the rest -# ask for structured output that cannot resume from an arbitrary cut point. _UNCONSTRAINED_RESPONSE_FORMAT: Final[Mapping[str, str]] = MappingProxyType({"type": "text"}) _SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str]) @@ -1266,9 +1264,7 @@ class Router: ContinuationPrefillDeploymentCheck, ) - # Registered on the process-global callback list, never tracked per - # router, so discarding one router cannot drop the filter another - # still needs. It is inert unless a request carries the marker. + # 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): @@ -2559,7 +2555,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) @@ -2900,7 +2896,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: """ @@ -2966,9 +2962,6 @@ class Router: continue_after_content: Final = committed and self._mid_stream_continuation_eligible( e=e, request_kwargs=initial_kwargs ) - # Content already reached the caller and we cannot safely - # continue it (feature off, or tool/thinking/constrained output): - # surface the real error rather than restart into the same stream. if committed and not continue_after_content: self._raise_original_mid_stream_error(e) @@ -3001,8 +2994,6 @@ class Router: ) reduced_ceilings: Final = self._continuation_output_ceilings(initial_kwargs, emitted_tokens) if reduced_ceilings is None: - # The caller's output allowance is already spent; surface the - # error rather than grant a fresh allowance on this fallback hop. self._raise_original_mid_stream_error(e) initial_kwargs.update(reduced_ceilings) initial_kwargs["messages"] = self._build_completion_continuation_input( @@ -3202,8 +3193,7 @@ class Router: @staticmethod def _raise_original_mid_stream_error(e: "MidStreamFallbackError") -> NoReturn: - """Decline a continuation by surfacing the real provider error the stream - wrapper carried, rather than leaking the internal MidStreamFallbackError.""" + """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 @@ -3213,15 +3203,11 @@ class Router: e: "MidStreamFallbackError", request_kwargs: Mapping[str, object], ) -> bool: - """Whether a stream that broke after plain assistant text may be - continued via prefill. The fallback target's prefill support is checked - separately at deployment selection.""" + """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 - # Structured output cannot be resumed from an arbitrary cut point; - # `{"type": "text"}` is the unconstrained default and stays eligible. response_format: Final = request_kwargs.get("response_format") if response_format is not None and response_format != _UNCONSTRAINED_RESPONSE_FORMAT: return False @@ -3237,12 +3223,7 @@ class Router: request_kwargs: Mapping[str, object], emitted_tokens: int, ) -> Mapping[str, int] | None: - """The max_tokens / max_completion_tokens a continuation should carry, each - reduced by the tokens already emitted so the whole answer stays within the - caller's original allowance instead of getting a fresh one on every fallback - hop. Reductions compound across hops because the trimmed ceiling is what the - next hop sees. Returns None when the allowance is already exhausted, so the - stream must not be continued.""" + """The caller's output ceilings minus the tokens already emitted, or None once nothing is left.""" ceilings: Final = MappingProxyType( { key: value - emitted_tokens @@ -3259,9 +3240,7 @@ class Router: messages: Sequence[Mapping[str, str]], generated_content: str, ) -> Sequence[Mapping[str, object]]: - """Append the partial output as an assistant prefill for a prefill-capable - fallback to continue. A nested break folds the new partial into an - existing trailing prefill rather than appending a second assistant turn.""" + """Append the partial output as an assistant prefill, extending an existing trailing prefill in place.""" last: Final = messages[-1] if messages else None if last is not None and last.get("role") == "assistant" and last.get("prefix"): merged: Final = {**last, "content": str(last.get("content") or "") + generated_content} @@ -3620,7 +3599,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. @@ -3783,7 +3762,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) diff --git a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py index c640374858a..b8d6c68a353 100644 --- a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -1,10 +1,4 @@ -""" -Mid-stream fallback continuation: keep the fallback on a deployment that can -continue a prefilled assistant message. When the router marks a fallback -re-entry as a continuation, deployments whose model does not support assistant -prefill are dropped, so the partial text is continued rather than regenerated or -rejected. Requests without the marker pass through untouched. -""" +"""Keeps a mid-stream continuation on deployments whose model supports assistant prefill.""" from collections.abc import Mapping, Sequence from typing import Final @@ -15,16 +9,11 @@ from litellm.integrations.custom_logger import CustomLogger, Span from litellm.types.llms.openai import AllMessageValues from litellm.utils import supports_assistant_prefill -# The router marks a continuation re-entry by placing MID_STREAM_CONTINUATION_MARKER -# under this key. Since the proxy can forward arbitrary request-body fields into the -# router, the marker is a private object checked by type rather than a truthy value: -# a JSON request body cannot construct one, so a client cannot forge the flag to steer -# deployment selection toward prefill-capable deployments. MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation" class _ContinuationMarker: - """Unforgeable sentinel; only the router can produce an instance.""" + """Only the router constructs one, so a JSON request body cannot forge the continuation flag.""" MID_STREAM_CONTINUATION_MARKER: Final = _ContinuationMarker() @@ -37,16 +26,12 @@ def _deployment_supports_prefill(deployment: object) -> bool: deployment_map: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment) except ValidationError: return False - # A per-deployment model_info override wins, so a model that is not in the cost - # map (or is registered generically) can still opt in or out explicitly with - # `model_info: {"supports_assistant_prefill": true|false}`. try: model_info: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment_map.get("model_info")) declared: Final = model_info.get("supports_assistant_prefill") if isinstance(declared, bool): return declared except ValidationError: - # No usable model_info override; fall through to the cost-map lookup. pass try: litellm_params: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment_map.get("litellm_params")) @@ -69,4 +54,4 @@ class ContinuationPrefillDeploymentCheck(CustomLogger): if not isinstance(marker, _ContinuationMarker): return healthy_deployments eligible: Final = (deployment for deployment in healthy_deployments if _deployment_supports_prefill(deployment)) - return list(eligible) # mutable-ok: downstream deployment selection consumes a mutable list + return list(eligible)