fix(router): run post_call_rules over the joined text of a continued stream

Each provider stream wrapper validates only the fragment it produced, so a
rule that trips on the primary's prefix plus the continuation's text never
fired. The router now feeds litellm.post_call_rules the emitted text plus
every continuation chunk before yielding it. Also drops a narrating
test-helper docstring.
This commit is contained in:
mateo-berri 2026-10-01 17:14:48 -07:00
parent eab6d76e9a
commit 6e52b45dbf
3 changed files with 83 additions and 8 deletions

View file

@ -2983,12 +2983,16 @@ class Router:
"content_policy_fallbacks", self.content_policy_fallbacks
)
initial_kwargs["original_function"] = self._acompletion
if continue_after_content:
from litellm.router_utils.pre_call_checks.continuation_prefill_check import (
MID_STREAM_CONTINUATION_KWARG,
MID_STREAM_CONTINUATION_MARKER,
)
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
)
@ -3035,6 +3039,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

View file

@ -5,8 +5,11 @@ 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"
@ -55,3 +58,26 @@ class ContinuationPrefillDeploymentCheck(CustomLogger):
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)

View file

@ -2663,9 +2663,6 @@ async def test_acompletion_streaming_iterator():
def _make_midstream_source(error, chunks=None):
"""A minimal stand-in for the CustomStreamWrapper: yields one content chunk,
then raises ``error`` on the next pull. Carries the attributes
_acompletion_streaming_iterator reads."""
from unittest.mock import MagicMock
first_chunk = MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))])
@ -2756,6 +2753,52 @@ async def test_acompletion_streaming_iterator_continues_after_content_when_eligi
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",