mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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
This commit is contained in:
parent
c75e3a69e4
commit
ba3374732e
5 changed files with 44 additions and 62 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue