diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index f3274151e5a..486c95e227e 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index d0f4e5ff44d..dec88c0238b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index b2002f9a0f9..692c207003b 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -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, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index cd235d8de67..223206ab5b3 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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.