From ba3374732e3e00b2e78afb86a8b8886788a175eb Mon Sep 17 00:00:00 2001 From: Ayush Date: Tue, 15 Sep 2026 02:59:51 +0530 Subject: [PATCH] fix(router): address review of mid-stream continuation - register the continuation deployment filter as a process-global singleton instead of tracking it per router, so discarding one router no longer removes the type-deduplicated filter that other live routers still depend on - disqualify reasoning_content: it reaches the caller as visible reasoning that a text-only prefill cannot carry, so a reasoning-capable fallback would re-derive it and produce an incoherent response - move MID_STREAM_CONTINUATION_KWARG to constants and lazy-import the filter class in the router, breaking the module-level import cycle CodeQL flagged - trim the added docstrings to the repository comment policy --- litellm/constants.py | 4 ++ .../litellm_core_utils/streaming_handler.py | 16 ++---- litellm/router.py | 50 ++++++------------- .../continuation_prefill_check.py | 18 ++----- .../test_streaming_handler.py | 18 +++++-- 5 files changed, 44 insertions(+), 62 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 5751e6e46af..487c1a1e020 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2039,3 +2039,7 @@ BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" # Shared read-only empty mapping, for defaulting optional Mapping parameters without # constructing a fresh mutable dict at each call site. EMPTY_MAPPING: Final = MappingProxyType({}) + +# Marks a fallback re-entry as a mid-stream continuation, read by the deployment +# pre-call filter. Lives here so router and the filter share it without an import cycle. +MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation" diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index f62a23d2c42..b6301929ce3 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2390,9 +2390,12 @@ 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", @@ -2402,23 +2405,14 @@ class CustomStreamWrapper: @classmethod def _delta_disqualifies_continuation(cls, delta: object) -> bool: - """ - True when a streamed delta carries output a text-only prefill - continuation cannot represent: tool/function calls, signed Anthropic - thinking blocks, structured reasoning items, audio or image parts, or - annotations. Plain ``reasoning_content`` is deliberately not here - it - is out-of-band, never reaches the caller as answer text, and so does - not block a continuation (parity with the Responses-API path). - """ get: Final = getattr(delta, "get", None) if not callable(get): return False return any(get(field) for field in cls._CONTINUATION_DISQUALIFYING_DELTA_FIELDS) def _accumulate_streamed_delta(self, delta: object) -> None: - """Grow the running answer text and latch whether anything a - continuation cannot carry has been streamed. One home for both so the - three iteration sites (sync, async, non-aiohttp) stay in step.""" + # 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 "" self.response_uptil_now += content or "" diff --git a/litellm/router.py b/litellm/router.py index 3de4a70d5c4..9f770b23608 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -62,6 +62,7 @@ from litellm.constants import ( DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER, DEFAULT_MAX_LRU_CACHE_SIZE, INTERNAL_CALL_ORIGIN_METADATA_KEY, + MID_STREAM_CONTINUATION_KWARG, OUTPUT_TOKEN_CEILING_PARAMS, ROUTING_REQUEST_TAGS_METADATA_KEY, RUNTIME_UPDATABLE_ROUTER_SETTINGS, @@ -189,10 +190,6 @@ from litellm.router_utils.handle_error import ( send_llm_exception_alert, ) from litellm.router_utils.health_state_cache import DeploymentHealthCache -from litellm.router_utils.pre_call_checks.continuation_prefill_check import ( - MID_STREAM_CONTINUATION_KWARG, - ContinuationPrefillDeploymentCheck, -) from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( DeploymentAffinityCheck, warn_on_unknown_model_group_affinity_flags, @@ -1186,12 +1183,14 @@ class Router: self.add_optional_pre_call_checks(default_pre_call_checks) if self.enable_mid_stream_fallback_continuation: - if self.optional_callbacks is None: - self.optional_callbacks = [] - if not any(isinstance(cb, ContinuationPrefillDeploymentCheck) for cb in self.optional_callbacks): - continuation_check: Final = ContinuationPrefillDeploymentCheck() - self.optional_callbacks.append(continuation_check) - litellm.logging_callback_manager.add_litellm_callback(continuation_check) + from litellm.router_utils.pre_call_checks.continuation_prefill_check import ( + 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. + litellm.logging_callback_manager.add_litellm_callback(ContinuationPrefillDeploymentCheck()) def discard(self): """ @@ -3051,22 +3050,14 @@ class Router: e: "MidStreamFallbackError", request_kwargs: Mapping[str, object], ) -> bool: - """ - Whether a chat-completions stream that broke after content may be - continued on a fallback deployment via assistant prefill, instead of - re-raising. Only plain assistant text is safe: a continuation built from - ``generated_content`` (text-only) cannot carry tool calls, signed - thinking blocks, audio or images, and a constrained (JSON / forced - tool_choice) or merged-reasoning response cannot be resumed from an - arbitrary cut point. The fallback target's prefill support is enforced - separately at deployment selection. - """ + """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.""" if not self.enable_mid_stream_fallback_continuation: return False if not e.generated_content or e.emitted_disqualifying_content: return False - # Any structured-output request (response_format, or a forced tool call) - # produces a partial that cannot be resumed from an arbitrary cut point. + # 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 != {"type": "text"}: @@ -3083,18 +3074,9 @@ class Router: messages: list[dict[str, str]], generated_content: str, ) -> Sequence[Mapping[str, object]]: - """ - Append the partial assistant output as a prefill so a prefill-capable - fallback continues where the broken stream stopped instead of - regenerating text already delivered to the caller. The deployment filter - guarantees the target supports ``prefix: True`` (parity with - ``_build_responses_continuation_input`` for the Responses-API path). - - A nested mid-stream break can re-enter here with a prefill already - appended; the new partial is folded into that trailing assistant turn so - the request keeps a single prefill rather than two consecutive assistant - messages a non-merging provider would reject. - """ + """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.""" if messages and messages[-1].get("role") == "assistant" and messages[-1].get("prefix"): last: Final = messages[-1] merged: dict[str, object] = {**last, "content": str(last.get("content") or "") + generated_content} 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 14e637280d8..2a527c8b249 100644 --- a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -1,28 +1,20 @@ """ Mid-stream fallback continuation: keep the fallback on a deployment that can -actually continue a prefilled assistant message. - -When a chat-completions stream breaks after content and the Router re-enters the -fallback chain to continue it (the request is marked with -``MID_STREAM_CONTINUATION_KWARG``), only a deployment whose model supports -assistant prefill can pick up the partial text without regenerating it. -Deployments that cannot are dropped, so selection lands on a -continuation-capable one. If a group has none it empties and the fallback chain -moves on, surfacing the original error rather than sending a request the target -would reject or duplicate. A request without the marker is passed through -untouched. +continue a prefilled assistant message. When a request carries +``MID_STREAM_CONTINUATION_KWARG``, 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. """ from typing import Final from pydantic import TypeAdapter, ValidationError +from litellm.constants import MID_STREAM_CONTINUATION_KWARG from litellm.integrations.custom_logger import CustomLogger, Span from litellm.types.llms.openai import AllMessageValues from litellm.utils import supports_assistant_prefill -MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation" - _STR_KEYED_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index f233d6e9bb2..7aa244a8433 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -4925,14 +4925,24 @@ class TestContinuationDisqualifiers: @pytest.mark.parametrize( "field", - ["tool_calls", "function_call", "thinking_blocks", "reasoning_items", "audio", "images", "annotations"], + [ + "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"}, {"reasoning_content": "thinking"}, {}, {"role": "assistant"}]) - def test_plain_text_and_reasoning_content_not_flagged(self, delta): - # plain reasoning_content is out-of-band and must NOT block a continuation + @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):