From 6e52b45dbff214226447a348b95dc7155f2908c2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:14:48 -0700 Subject: [PATCH] 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. --- litellm/router.py | 16 ++++-- .../continuation_prefill_check.py | 26 ++++++++++ tests/unit/test_router/test_router.py | 49 +++++++++++++++++-- 3 files changed, 83 insertions(+), 8 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index d479519c18e..849edfc66a7 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 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 b8d6c68a353..0d6d462d417 100644 --- a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -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) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 753161a3bbc..5e1b2f3b8ef 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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",