mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
eab6d76e9a
commit
6e52b45dbf
3 changed files with 83 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue