fix(router): keep order fallback on the requested order level

When a pre-call filter left no order-2 deployments, target_order matching
fell through to the remaining healthy list and reselected the failed
primary. Prompt-cache and deployment affinity also pinned that hop back
to order 1. Match the requested order strictly, skip those pins while
target_order is set, and keep target_order across retries of that hop.
This commit is contained in:
Emerson Gomes 2026-08-31 10:21:22 -05:00
parent 4ba8517134
commit bd0b9c78bd
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
8 changed files with 220 additions and 11 deletions

View file

@ -2174,6 +2174,7 @@ class Router:
"client": model_client,
**kwargs,
}
input_kwargs.pop("_target_order", None)
response: Final = litellm.completion(**input_kwargs)
verbose_router_logger.info("litellm.completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
@ -3197,6 +3198,7 @@ class Router:
}
input_kwargs.pop("silent_model", None)
input_kwargs.pop("include_fallback_errors", None)
input_kwargs.pop("_target_order", None)
_response: Final = litellm.acompletion(**input_kwargs)
@ -11928,7 +11930,7 @@ class Router:
)
## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
_target_order: Final = (request_kwargs or {}).pop("_target_order", None)
_target_order: Final = (request_kwargs or {}).get("_target_order")
healthy_deployments = litellm.utils._get_order_filtered_deployments(
cast(list[dict], healthy_deployments), target_order=_target_order
)
@ -12693,7 +12695,7 @@ class Router:
)
## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
_target_order: Final = (request_kwargs or {}).pop("_target_order", None)
_target_order: Final = (request_kwargs or {}).get("_target_order")
healthy_deployments = litellm.utils._get_order_filtered_deployments(
healthy_deployments, target_order=_target_order
)

View file

@ -414,7 +414,9 @@ async def run_async_fallback(
verbose_router_logger.info("Falling back to model_group = %s", mask_sensitive_structure(mg))
if isinstance(mg, str):
kwargs["model"] = mg
kwargs.pop("_target_order", None)
elif isinstance(mg, dict):
kwargs.pop("_target_order", None)
kwargs.update(mg)
fallback_depth = fallback_depth + 1
_hop_metadata = dict(kwargs.get(metadata_variable_name) or {})

View file

@ -427,6 +427,8 @@ class DeploymentAffinityCheck(CustomLogger):
"""
request_kwargs = request_kwargs or {}
typed_healthy_deployments: Final = cast(list[dict], healthy_deployments)
if request_kwargs.get("_target_order") is not None:
return typed_healthy_deployments
(
enable_user_key,

View file

@ -58,6 +58,9 @@ class PromptCachingDeploymentCheck(CustomLogger):
request_kwargs: dict | None = None,
parent_otel_span: Span | None = None,
) -> list[dict]:
if request_kwargs is not None and request_kwargs.get("_target_order") is not None:
return healthy_deployments
if messages is not None and is_prompt_caching_valid_prompt(
messages=messages,
model=model,

View file

@ -4859,11 +4859,7 @@ def _get_deployment_order(deployment: dict | Any) -> int | None:
def _get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list:
if target_order is not None:
filtered: Final = [d for d in healthy_deployments if _get_deployment_order(d) == target_order]
if filtered:
return filtered
# target_order doesn't match any deployment (e.g., external fallback model) — return all
return healthy_deployments
return [d for d in healthy_deployments if _get_deployment_order(d) == target_order]
# Default: pick min order group
_valid_orders: Final[list[int]] = [

View file

@ -598,6 +598,45 @@ async def test_async_filter_deployments_falls_back_when_cached_deployment_is_unh
assert filtered == healthy_deployments
@pytest.mark.asyncio
async def test_async_filter_deployments_does_not_pin_when_target_order_is_set():
user_key = "user-key-order-fallback"
stable_model_map_key = "claude-sonnet-4-5@20250929"
cache = AsyncMock()
cache.async_get_cache = AsyncMock(return_value={"model_id": "deployment-1"})
callback = DeploymentAffinityCheck(
cache=cache,
ttl_seconds=123,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
)
healthy_deployments = [
{
"model_name": stable_model_map_key,
"litellm_params": {"model": f"vertex_ai/{stable_model_map_key}"},
"model_info": {"id": "deployment-1"},
},
{
"model_name": stable_model_map_key,
"litellm_params": {
"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"
},
"model_info": {"id": "deployment-2"},
},
]
filtered = await callback.async_filter_deployments(
model="some-router-model-group",
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"_target_order": 2, "metadata": {"user_api_key_hash": user_key}},
parent_otel_span=None,
)
assert filtered == healthy_deployments
cache.async_get_cache.assert_not_called()
@pytest.mark.asyncio
async def test_async_user_key_affinity_ttl_expiry_allows_reroute():
"""

View file

@ -150,6 +150,25 @@ async def test_async_filter_deployments_narrows_prompt_above_model_minimum():
assert filtered == [deployments[1]]
@pytest.mark.asyncio
async def test_async_filter_deployments_does_not_pin_when_target_order_is_set():
cache = DualCache()
check = PromptCachingDeploymentCheck(cache=cache)
deployments = _deployments("anthropic/claude-opus-4-6", "anthropic/claude-opus-4-6")
messages = _messages(word_count=5000)
await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-2", messages=messages, tools=None)
filtered = await check.async_filter_deployments(
model=MODEL_GROUP_ALIAS,
healthy_deployments=deployments,
messages=messages,
request_kwargs={"_target_order": 2},
)
assert filtered == deployments
@pytest.mark.asyncio
async def test_async_filter_deployments_narrows_for_group_whose_model_minimum_is_lower():
"""

View file

@ -6,12 +6,15 @@ should be tried first, and higher order deployments should be used as fallbacks
when lower order deployments fail.
"""
from typing import Optional
from typing import Final, Optional
import pytest
import litellm
from litellm import Router
from litellm.utils import _get_order_filtered_deployments
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.prompt_caching_cache import PromptCachingCache
from litellm.utils import _get_deployment_order, _get_order_filtered_deployments
# ---------------------------------------------------------------------------
# Unit tests for _get_order_filtered_deployments
@ -49,13 +52,22 @@ class TestGetOrderFilteredDeployments:
assert len(result) == 1
assert result[0]["model_info"]["id"] == "b"
def test_target_order_no_match_returns_all(self):
def test_target_order_no_match_returns_empty(self):
deps = [
self._make_deployment(1, "a"),
self._make_deployment(2, "b"),
]
result = _get_order_filtered_deployments(deps, target_order=99)
assert len(result) == 2
assert result == []
def test_target_order_no_match_does_not_reselect_lower_order(self):
deps = [
self._make_deployment(1, "a"),
self._make_deployment(2, "b"),
]
remaining_after_pre_call = [deps[0]]
result = _get_order_filtered_deployments(remaining_after_pre_call, target_order=2)
assert result == []
def test_no_order_set_returns_all(self):
deps = [
@ -406,6 +418,140 @@ async def test_router_order_fallback_with_hidden_model_group_alias():
assert response._hidden_params["model_id"] == "2"
@pytest.mark.asyncio
async def test_router_order_fallback_does_not_reselect_order_1_when_order_2_is_filtered_out():
class _DropOrder2(CustomLogger):
async def async_filter_deployments(
self, model, healthy_deployments, messages, request_kwargs=None, parent_otel_span=None
):
return [d for d in healthy_deployments if _get_deployment_order(d) != 2]
drop_order_2: Final = _DropOrder2()
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "key",
"mock_response": "litellm.RateLimitError",
"order": 1,
},
"model_info": {"id": "1"},
},
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "key",
"mock_response": "success from order 2",
"order": 2,
},
"model_info": {"id": "2"},
},
],
num_retries=0,
)
litellm.callbacks.append(drop_order_2)
try:
with pytest.raises(Exception) as exc_info:
await router.acompletion(
model="test-model",
messages=[{"role": "user", "content": "hi"}],
)
assert "success from order 2" not in str(exc_info.value)
assert getattr(exc_info.value, "_hidden_params", {}).get("model_id") != "1"
finally:
litellm.callbacks.remove(drop_order_2)
@pytest.mark.asyncio
async def test_router_order_fallback_ignores_prompt_cache_pin_on_target_order():
messages = [{"role": "user", "content": "word " * 5000}]
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "bad",
"mock_response": Exception("azure peak load"),
"order": 1,
},
"model_info": {"id": "1"},
},
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "good",
"mock_response": "success from order 2",
"order": 2,
},
"model_info": {"id": "2"},
},
],
num_retries=0,
optional_pre_call_checks=["prompt_caching"],
)
await PromptCachingCache(cache=router.cache).async_add_model_id(
model_id="1",
messages=messages,
tools=None,
)
response = await router.acompletion(model="test-model", messages=messages)
assert response._hidden_params["model_id"] == "2"
@pytest.mark.asyncio
async def test_router_order_fallback_retries_keep_target_order():
seen_target_orders: Final = []
class _RecordTargetOrder(CustomLogger):
async def async_filter_deployments(
self, model, healthy_deployments, messages, request_kwargs=None, parent_otel_span=None
):
seen_target_orders.append((request_kwargs or {}).get("_target_order"))
return healthy_deployments
recorder: Final = _RecordTargetOrder()
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "bad",
"mock_response": Exception("fail order 1"),
"order": 1,
},
"model_info": {"id": "1"},
},
{
"model_name": "test-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "bad",
"mock_response": Exception("fail order 2"),
"order": 2,
},
"model_info": {"id": "2"},
},
],
num_retries=1,
)
litellm.callbacks.append(recorder)
try:
with pytest.raises(Exception, match="fail order 2"):
await router.acompletion(
model="test-model",
messages=[{"role": "user", "content": "hi"}],
)
finally:
litellm.callbacks.remove(recorder)
assert seen_target_orders.count(2) >= 2
def test_check_non_standard_fallback_format():
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,