mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
e15b37a18e
commit
ed3b86686d
4 changed files with 120 additions and 2 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue