fix(router): honor content_policy_fallbacks on streaming requests

Streaming requests bypassed content_policy_fallbacks entirely. The 4xx
filter in CustomStreamWrapper._handle_stream_fallback_error raised any
client error except 429 directly, so a ContentPolicyViolationError
(status 400) never became a MidStreamFallbackError and the router's
dedicated content-policy fallback branch never engaged. The same
payload with stream=false fell back correctly.

Exempt ContentPolicyViolationError and ContextWindowExceededError from
the direct-raise in both 4xx guards; both types have dedicated fallback
configs and the router already special-cases the pair when ordering
fallbacks. Additionally, _acompletion_streaming_iterator passed the
wrapper MidStreamFallbackError to
async_function_with_fallbacks_common_utils, whose content-policy and
context-window branches match on isinstance of the original exception
types, so a wrapped error only ever reached the generic fallbacks.
Unwrap original_exception for those two types before dispatching.

Fixes #28599.
This commit is contained in:
Filippo Mattia Menghi 2026-06-10 09:49:32 +02:00
parent e15b37a18e
commit ed3b86686d
4 changed files with 120 additions and 2 deletions

View file

@ -2318,8 +2318,15 @@ class CustomStreamWrapper:
429 (rate-limit) is explicitly exempted from the 4xx filter because
it is transient and the Router should switch to another model group.
Context-window and content-policy errors are also exempted because
the Router has dedicated fallback configs for them
(context_window_fallbacks / content_policy_fallbacks).
"""
from litellm.exceptions import MidStreamFallbackError
from litellm.exceptions import (
ContentPolicyViolationError,
ContextWindowExceededError,
MidStreamFallbackError,
)
# Map to OpenAI exception format
if isinstance(e, OpenAIError):
@ -2358,19 +2365,27 @@ class CustomStreamWrapper:
mapped_status_code = _normalize_status_code(mapped_exception)
original_status_code = _normalize_status_code(e)
has_dedicated_fallback_config = isinstance(
mapped_exception,
(ContentPolicyViolationError, ContextWindowExceededError),
)
# Raise non-retriable client errors directly (skip fallback).
# Exception: 429 (rate-limit) IS retriable/transient — allow it
# through so the Router can switch to a different model group.
# Same for errors with dedicated fallback configs.
if (
mapped_status_code is not None
and 400 <= mapped_status_code < 500
and mapped_status_code != 429
and not has_dedicated_fallback_config
):
raise mapped_exception
if (
original_status_code is not None
and 400 <= original_status_code < 500
and original_status_code != 429
and not has_dedicated_fallback_config
):
raise mapped_exception

View file

@ -2196,9 +2196,21 @@ class Router:
self._update_kwargs_before_fallbacks(
model=model_group, kwargs=initial_kwargs
)
# Unwrap errors with dedicated fallback configs so the
# isinstance checks in the common utils can route them to
# context_window_fallbacks / content_policy_fallbacks.
fallback_trigger_exception: Exception = e
if isinstance(
e.original_exception,
(
litellm.ContextWindowExceededError,
litellm.ContentPolicyViolationError,
),
):
fallback_trigger_exception = e.original_exception
fallback_response = (
await self.async_function_with_fallbacks_common_utils(
e=e,
e=fallback_trigger_exception,
disable_fallbacks=False,
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,

View file

@ -876,6 +876,41 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging):
assert "invalid maxOutputTokens" in str(excinfo.value)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"exception_class",
[litellm.ContentPolicyViolationError, litellm.ContextWindowExceededError],
)
async def test_streaming_dedicated_fallback_error_triggers_midstream_fallback(
logging_obj: Logging, exception_class
):
"""400s with dedicated fallback configs (content_policy_fallbacks /
context_window_fallbacks) must wrap into MidStreamFallbackError instead
of raising directly, so the Router's fallback chain can engage.
Regression test for https://github.com/BerriAI/litellm/issues/28599
"""
from litellm.exceptions import MidStreamFallbackError
async def _raise_dedicated_fallback_error(**kwargs):
raise exception_class(
message="content blocked", model="gpt-4", llm_provider="openai"
)
response = CustomStreamWrapper(
completion_stream=None,
model="gpt-4",
logging_obj=logging_obj,
custom_llm_provider="openai",
make_call=_raise_dedicated_fallback_error,
)
with pytest.raises(MidStreamFallbackError) as excinfo:
await response.__anext__()
assert isinstance(excinfo.value.original_exception, exception_class)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_streaming_read_timeout_triggers_midstream_fallback(
logging_obj: Logging,

View file

@ -1623,6 +1623,62 @@ async def test_acompletion_streaming_iterator_preserves_hidden_params():
assert result._hidden_params.get("_response_ms") == 500.0
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_unwraps_content_policy_violation():
"""A MidStreamFallbackError wrapping a ContentPolicyViolationError must
pass the original exception to the fallback utils, so the dedicated
content_policy_fallbacks branch matches on its isinstance check.
Regression test for https://github.com/BerriAI/litellm/issues/28599
"""
from litellm.exceptions import ContentPolicyViolationError, MidStreamFallbackError
cpv_error = ContentPolicyViolationError(
message="content blocked", model="gpt-4", llm_provider="openai"
)
mid_stream_error = MidStreamFallbackError(
message=str(cpv_error),
model="gpt-4",
llm_provider="openai",
original_exception=cpv_error,
is_pre_first_chunk=True,
)
class FailingStream:
model = "gpt-4"
custom_llm_provider = "openai"
logging_obj = MagicMock()
chunks = []
def __aiter__(self):
return self
async def __anext__(self):
raise mid_stream_error
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
}
],
content_policy_fallbacks=[{"gpt-4": ["gpt-3.5-turbo"]}],
)
with patch.object(
router, "async_function_with_fallbacks_common_utils", return_value=object()
) as mock_fallback_utils:
result = await router._acompletion_streaming_iterator(
model_response=FailingStream(),
messages=[{"role": "user", "content": "Hello"}],
initial_kwargs={"model": "gpt-4", "stream": True},
)
async for _ in result:
pass
assert mock_fallback_utils.call_args.kwargs["e"] is cpv_error
def test_completion_streaming_iterator_fallback_on_429():
"""Sync streaming: MidStreamFallbackError (429 pre-first-chunk) triggers fallback.