refactor(router): type the streaming fallback kwargs and drop narration comments

Removes an unused mutable-ok marker (LIT013) and the explanatory comments
Greptile flagged, and types initial_kwargs on both streaming fallback
iterators as dict[str, Any] like their Responses siblings so the new
continuation helpers take a typed mapping and the basedpyright budget
stays at or below main.
This commit is contained in:
mateo-berri 2026-10-01 17:04:07 -07:00
parent ba472dfa71
commit eab6d76e9a
2 changed files with 12 additions and 48 deletions

View file

@ -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)

View file

@ -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)