This commit is contained in:
TaehunOh 2026-10-03 10:33:07 -07:00 • committed by GitHub
commit 96a5369df2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 301 additions and 4 deletions

View file

@ -500,6 +500,15 @@ def _deployment_pick_attributes(model: str, request_kwargs: Mapping[str, object]
)
def _untried_fallback_target_exists(chain: Sequence[object] | None, kwargs: Mapping[str, Any]) -> bool:
"""
Whether a resolved fallback chain still names a target this request can try.
`has_unattempted_fallback_target` reports an unattempted target for an empty chain
whenever the request carries no attempt record yet, so guard the empty case here.
"""
return bool(chain) and has_unattempted_fallback_target(chain, kwargs)
def _as_retry_skipped_deployment_ids(value: object) -> tuple[str, ...]:
return tuple(item for item in value if isinstance(item, str)) if isinstance(value, tuple) else ()
@ -7793,6 +7802,14 @@ class Router:
num_retries=num_retries,
healthy_deployments=_healthy_deployments,
all_deployments=_all_deployments,
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)
@ -7863,6 +7880,14 @@ class Router:
num_retries=num_retries,
healthy_deployments=_healthy_deployments,
all_deployments=_all_deployments,
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)
@ -8020,6 +8045,7 @@ class Router:
num_retries: int,
healthy_deployments: list | None = None,
all_deployments: list | None = None,
fallback_available: bool = False,
) -> int | float:
"""
Calculate back-off, then retry
@ -8028,6 +8054,8 @@ class Router:
1. there are healthy deployments in the same model group
2. there are fallbacks for the completion call
"""
if fallback_available:
return 0
## base case - single deployment
if all_deployments is not None and len(all_deployments) == 1:
@ -8493,14 +8521,20 @@ class Router:
return self._has_content_policy_fallback(model_group, kwargs)
if self._has_default_fallbacks():
return True
fallbacks: Final = kwargs.get("fallbacks", self.fallbacks)
if fallbacks is None:
return self._regular_fallback_available(
fallbacks=kwargs.get("fallbacks", self.fallbacks), model_group=model_group, kwargs=kwargs
)
def _regular_fallback_available(
self, fallbacks: list | None, model_group: str | None, kwargs: Mapping[str, Any]
) -> bool:
if fallbacks is None or fallbacks_disabled_for_request(kwargs):
return False
resolved, _ = get_fallback_model_group_for_lookup_groups(
fallbacks=fallbacks,
lookup_groups=fallback_lookup_groups(kwargs, model_group),
)
return has_unattempted_fallback_target(resolved, kwargs)
return _untried_fallback_target_exists(resolved, kwargs)
def _anthropic_messages_order_levels(self, model_group: str, kwargs: Mapping[str, Any]) -> tuple[int, ...]:
"""
@ -8552,6 +8586,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 _untried_fallback_target_exists(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 _untried_fallback_target_exists(
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

@ -2,13 +2,241 @@
Tests for router retry backoff behavior.
"""
from unittest.mock import patch
import asyncio
from typing import Final
from unittest.mock import AsyncMock, patch
import httpx
import pytest
import litellm
from litellm import Router
from litellm.constants import MAX_RETRY_DELAY
from litellm.router import _untried_fallback_target_exists
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets, record_disable_fallbacks
from litellm.types.router import RetryPolicy
_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"
)
_CONTEXT_WINDOW_ERROR: Final = litellm.ContextWindowExceededError(
message="too long", model="gpt-5.4-mini", llm_provider="openai"
)
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": primary_error} if primary_error is not None else {}),
},
{
"model_name": "backup",
"litellm_params": {
"model": "openai/gpt-5.4-mini",
"api_key": "sk-test",
"mock_response": "answered by backup",
},
},
],
num_retries=2,
retry_after=int(MAX_RETRY_DELAY),
fallbacks=fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
retry_policy=retry_policy,
)
def _backoff_delays(sleeper: AsyncMock) -> tuple[float, ...]:
"""
The delays the router asked to wait between retries. Reading its decision keeps these
tests off the clock, which tests/unit rules out, and off the CI scheduler's timing.
"""
return tuple(call.args[0] for call in sleeper.await_args_list if call.args)
@pytest.mark.asyncio
async def test_single_deployment_group_with_fallback_does_not_back_off_before_falling_back():
router: Final = _router_with_single_failing_deployment(fallbacks=[{"primary": ["backup"]}])
sleeper: Final = AsyncMock()
with patch.object(asyncio, "sleep", sleeper):
response: Final = await router.acompletion(model="primary", messages=[{"role": "user", "content": "Hello"}])
assert response.choices[0].message.content == "answered by backup"
assert not any(delay > 0 for delay in _backoff_delays(sleeper))
@pytest.mark.asyncio
async def test_fallback_configured_for_another_group_keeps_retry_backoff():
router: Final = _router_with_single_failing_deployment(fallbacks=[{"backup": ["primary"]}])
sleeper: Final = AsyncMock()
with patch.object(asyncio, "sleep", sleeper), pytest.raises(litellm.InternalServerError):
await router.acompletion(model="primary", messages=[{"role": "user", "content": "Hello"}])
assert any(delay > 0 for delay in _backoff_delays(sleeper))
@pytest.mark.asyncio
async def test_empty_fallback_chain_keeps_retry_backoff():
"""A chain configured as an empty list names no target, so the request has nothing to
fall back to and the retries must still space themselves out."""
router: Final = _router_with_single_failing_deployment(fallbacks=[{"primary": []}])
sleeper: Final = AsyncMock()
with patch.object(asyncio, "sleep", sleeper), pytest.raises(litellm.InternalServerError):
await router.acompletion(model="primary", messages=[{"role": "user", "content": "Hello"}])
assert any(delay > 0 for delay in _backoff_delays(sleeper))
def test_untried_fallback_target_exists_is_false_for_a_chain_with_no_entries():
assert _untried_fallback_target_exists(["backup"], {}) is True
assert _untried_fallback_target_exists([], {}) is False
assert _untried_fallback_target_exists(None, {}) is False
@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=[])
sleeper: Final = AsyncMock()
with patch.object(asyncio, "sleep", sleeper):
response: Final = await router.acompletion(
model="primary",
messages=[{"role": "user", "content": "Hello"}],
fallbacks=[{"model": "backup"}],
)
assert response.choices[0].message.content == "answered by backup"
assert not any(delay > 0 for delay in _backoff_delays(sleeper))
@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),
)
sleeper: Final = AsyncMock()
with patch.object(asyncio, "sleep", sleeper), pytest.raises(litellm.ContentPolicyViolationError):
await router.acompletion(
model="primary",
messages=[{"role": "user", "content": "Hello"}],
mock_response=_CONTENT_POLICY_ERROR,
)
assert any(delay > 0 for delay in _backoff_delays(sleeper))
@pytest.mark.parametrize(
("error", "fallbacks", "context_window_fallbacks", "content_policy_fallbacks", "expected"),
[
pytest.param(_SERVER_ERROR, [{"primary": ["backup"]}], None, None, True, id="own-chain"),
pytest.param(_SERVER_ERROR, [{"backup": ["primary"]}], None, None, False, id="chain-for-another-group"),
pytest.param(_SERVER_ERROR, [{"*": ["backup"]}], None, None, True, id="generic-chain"),
pytest.param(_SERVER_ERROR, [{"model": "backup"}], None, None, True, id="client-side-list"),
pytest.param(
_CONTENT_POLICY_ERROR,
[{"primary": ["backup"]}],
None,
[{"backup": ["primary"]}],
False,
id="content-policy-list-skips-group",
),
pytest.param(
_CONTENT_POLICY_ERROR,
[{"backup": ["primary"]}],
None,
[{"primary": ["backup"]}],
True,
id="content-policy-list-covers-group",
),
pytest.param(
_CONTEXT_WINDOW_ERROR,
[{"primary": ["backup"]}],
[{"backup": ["primary"]}],
None,
False,
id="context-window-list-skips-group",
),
pytest.param(_CONTENT_POLICY_ERROR, [{"primary": ["backup"]}], None, None, True, id="no-dedicated-list"),
],
)
def test_fallback_available_for_error_follows_the_dispatch_order(
error: Exception,
fallbacks: list[dict[str, object]],
context_window_fallbacks: list[dict[str, list[str]]] | None,
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=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
model_group="primary",
kwargs={"model": "primary"},
)
assert available is expected
def test_fallback_available_for_error_is_false_without_a_model_group_or_when_disabled():
router: Final = _router_with_single_failing_deployment(fallbacks=[])
disabled_kwargs: Final = {"model": "primary", "metadata": {}}
record_disable_fallbacks(disabled_kwargs, True)
def available(model_group: str | None, kwargs: dict[str, object]) -> bool:
return router._fallback_available_for_error(
error=_SERVER_ERROR,
fallbacks=[{"model": "backup"}],
context_window_fallbacks=None,
content_policy_fallbacks=None,
model_group=model_group,
kwargs=kwargs,
)
assert (available("primary", {"model": "primary"}), available(None, {}), available("primary", disabled_kwargs)) == (
True,
False,
False,
)
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