fix(router): check the fallback the dispatcher will use before skipping backoff

This commit is contained in:
RosieOh 2026-09-22 18:12:40 +09:00
parent 077fc62e92
commit e0c1783975
2 changed files with 166 additions and 11 deletions

View file

@ -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.

View file

@ -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():
"""