diff --git a/litellm/router.py b/litellm/router.py index 93400951bac..53301979b2f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -7776,9 +7776,6 @@ class Router: raise verbose_router_logger.debug("Retrying request with num_retries: %s", num_retries) - fallback_available: Final = self._regular_fallback_available( - fallbacks=fallbacks, model_group=model_group, kwargs=kwargs - ) # decides how long to sleep before retry retry_after: Final = self._time_to_sleep_before_retry( e=original_exception, @@ -7786,7 +7783,14 @@ class Router: num_retries=num_retries, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, - fallback_available=fallback_available, + fallback_available=self._fallback_available_for_error( + error=original_exception, + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + kwargs=kwargs, + ), ) await asyncio.sleep(retry_after) @@ -7857,7 +7861,14 @@ class Router: num_retries=num_retries, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, - fallback_available=fallback_available, + fallback_available=self._fallback_available_for_error( + error=e, + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + kwargs=kwargs, + ), ) await asyncio.sleep(_timeout) @@ -8555,6 +8566,41 @@ class Router: ) return has_unattempted_fallback_target(resolved, kwargs) + def _fallback_available_for_error( + self, + error: Exception, + fallbacks: list | None, + context_window_fallbacks: list | None, + content_policy_fallbacks: list | None, + model_group: str | None, + kwargs: Mapping[str, Any], + ) -> bool: + """ + Whether async_function_with_fallbacks_common_utils would hand this error to an untried + fallback, checked in the order it dispatches: client-side lists, then the dedicated + context-window or content-policy list (authoritative once set), then regular fallbacks + """ + if model_group is None or fallbacks_disabled_for_request(kwargs): + return False + if _check_non_standard_fallback_format(fallbacks=fallbacks): + return has_unattempted_fallback_target(fallbacks, kwargs) + dedicated_fallbacks: Final = ( + context_window_fallbacks + if isinstance(error, litellm.ContextWindowExceededError) + else content_policy_fallbacks + if isinstance(error, litellm.ContentPolicyViolationError) + else None + ) + if dedicated_fallbacks is not None: + return has_unattempted_fallback_target( + self._get_fallback_model_group_for_lookup_groups( + fallbacks=dedicated_fallbacks, + lookup_groups=fallback_lookup_groups(kwargs, model_group), + ), + kwargs, + ) + return self._regular_fallback_available(fallbacks=fallbacks, model_group=model_group, kwargs=kwargs) + def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool: """ Determines if a content policy error should be raised. diff --git a/tests/unit/test_router_retry_backoff_headers.py b/tests/unit/test_router_retry_backoff_headers.py index e44e7855925..7d1c416521a 100644 --- a/tests/unit/test_router_retry_backoff_headers.py +++ b/tests/unit/test_router_retry_backoff_headers.py @@ -12,20 +12,28 @@ import pytest import litellm from litellm import Router from litellm.constants import MAX_RETRY_DELAY +from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets, record_disable_fallbacks +from litellm.types.router import RetryPolicy _BACKOFF_DETECTION_TIMEOUT: Final = MAX_RETRY_DELAY / 4 +_SERVER_ERROR: Final = litellm.InternalServerError(message="provider down", model="gpt-5.4-mini", llm_provider="openai") +_CONTENT_POLICY_ERROR: Final = litellm.ContentPolicyViolationError( + message="flagged", model="gpt-5.4-mini", llm_provider="openai" +) -def _router_with_single_failing_deployment(fallbacks: list[dict[str, list[str]]]) -> Router: +def _router_with_single_failing_deployment( + fallbacks: list[dict[str, list[str]]], + primary_error: str | None = "litellm.InternalServerError", + content_policy_fallbacks: list[dict[str, list[str]]] | None = None, + retry_policy: RetryPolicy | None = None, +) -> Router: return Router( model_list=[ { "model_name": "primary", - "litellm_params": { - "model": "openai/gpt-5.4-mini", - "api_key": "sk-test", - "mock_response": "litellm.InternalServerError", - }, + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "sk-test"} + | ({"mock_response": primary_error} if primary_error is not None else {}), }, { "model_name": "backup", @@ -39,6 +47,8 @@ def _router_with_single_failing_deployment(fallbacks: list[dict[str, list[str]]] num_retries=2, retry_after=int(MAX_RETRY_DELAY), fallbacks=fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + retry_policy=retry_policy, ) @@ -65,6 +75,105 @@ async def test_fallback_configured_for_another_group_keeps_retry_backoff(): ) +@pytest.mark.asyncio +async def test_client_side_fallback_list_does_not_back_off_before_falling_back(): + router: Final = _router_with_single_failing_deployment(fallbacks=[]) + + response: Final = await asyncio.wait_for( + router.acompletion( + model="primary", + messages=[{"role": "user", "content": "Hello"}], + fallbacks=[{"model": "backup"}], + ), + timeout=_BACKOFF_DETECTION_TIMEOUT, + ) + + assert response.choices[0].message.content == "answered by backup" + + +@pytest.mark.asyncio +async def test_content_policy_error_keeps_backoff_when_its_dedicated_fallbacks_skip_the_group(): + router: Final = _router_with_single_failing_deployment( + fallbacks=[{"primary": ["backup"]}], + primary_error=None, + content_policy_fallbacks=[{"backup": ["primary"]}], + retry_policy=RetryPolicy(ContentPolicyViolationErrorRetries=2), + ) + + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for( + router.acompletion( + model="primary", + messages=[{"role": "user", "content": "Hello"}], + mock_response=_CONTENT_POLICY_ERROR, + ), + timeout=_BACKOFF_DETECTION_TIMEOUT, + ) + + +@pytest.mark.parametrize( + ("error", "fallbacks", "content_policy_fallbacks", "expected"), + [ + pytest.param(_SERVER_ERROR, [{"primary": ["backup"]}], None, True, id="own-chain"), + pytest.param(_SERVER_ERROR, [{"backup": ["primary"]}], None, False, id="chain-for-another-group"), + pytest.param(_SERVER_ERROR, [{"*": ["backup"]}], None, True, id="generic-chain"), + pytest.param(_SERVER_ERROR, [{"model": "backup"}], None, True, id="client-side-list"), + pytest.param( + _CONTENT_POLICY_ERROR, + [{"primary": ["backup"]}], + [{"backup": ["primary"]}], + False, + id="dedicated-list-skips-group", + ), + pytest.param( + _CONTENT_POLICY_ERROR, + [{"backup": ["primary"]}], + [{"primary": ["backup"]}], + True, + id="dedicated-list-covers-group", + ), + pytest.param(_CONTENT_POLICY_ERROR, [{"primary": ["backup"]}], None, True, id="no-dedicated-list"), + ], +) +def test_fallback_available_for_error_follows_the_dispatch_order( + error: Exception, + fallbacks: list[dict[str, object]], + content_policy_fallbacks: list[dict[str, list[str]]] | None, + expected: bool, +): + router: Final = _router_with_single_failing_deployment(fallbacks=[]) + + available: Final = router._fallback_available_for_error( + error=error, + fallbacks=fallbacks, + context_window_fallbacks=None, + content_policy_fallbacks=content_policy_fallbacks, + model_group="primary", + kwargs={"model": "primary"}, + ) + + assert available is expected + + +def test_regular_fallback_available_is_false_once_the_chain_is_used_up_or_disabled(): + router: Final = _router_with_single_failing_deployment(fallbacks=[]) + chain: Final = [{"primary": ["backup"]}] + disabled_kwargs: Final = {"model": "primary", "metadata": {}} + record_disable_fallbacks(disabled_kwargs, True) + + fresh: Final = router._regular_fallback_available( + fallbacks=chain, model_group="primary", kwargs={"model": "primary"} + ) + used_up: Final = router._regular_fallback_available( + fallbacks=chain, + model_group="primary", + kwargs={"model": "primary", "attempted_targets": AttemptedFallbackTargets(keys=frozenset({"backup"}))}, + ) + disabled: Final = router._regular_fallback_available(fallbacks=chain, model_group="primary", kwargs=disabled_kwargs) + + assert (fresh, used_up, disabled) == (True, False, False) + + @pytest.mark.asyncio async def test_retry_backoff_uses_current_exception_headers(): """