Merge pull request #38969 from emerzon/litellm_strict_order_fallback

fix(router): keep order fallback on the requested order level
This commit is contained in:
ryan-crabbe-berri 2026-09-01 15:21:59 -07:00 committed by GitHub
commit 9c417ba08b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 308 additions and 28 deletions

View file

@ -412,6 +412,7 @@ async def run_async_fallback(
# LOGGING
kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception)
verbose_router_logger.info("Falling back to model_group = %s", mask_sensitive_structure(mg))
kwargs.pop("_target_order", None) # rebind-ok: next hop must not inherit the previous order target
if isinstance(mg, str):
kwargs["model"] = mg
elif isinstance(mg, dict):

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

@ -4876,11 +4876,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,19 @@ should be tried first, and higher order deployments should be used as fallbacks
when lower order deployments fail.
"""
from typing import Optional
import json
from typing import Final, Optional
import httpx
import pytest
from openai import AsyncOpenAI
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.types.router import RouterRateLimitError
from litellm.utils import _get_deployment_order, _get_order_filtered_deployments
# ---------------------------------------------------------------------------
# Unit tests for _get_order_filtered_deployments
@ -49,13 +56,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,35 +422,239 @@ 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(RouterRateLimitError, match="No deployments available") 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)
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
@pytest.mark.asyncio
async def test_generic_api_call_strips_target_order_from_provider_kwargs():
captured: Final = {}
async def _fake_provider(**provider_kwargs):
captured.update(provider_kwargs)
return "ok"
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {"model": "gpt-4o", "api_key": "key", "order": 2},
"model_info": {"id": "2"},
},
],
)
response = await router._ageneric_api_call_with_fallbacks_helper(
model="test-model",
original_generic_function=_fake_provider,
_target_order=2,
messages=[{"role": "user", "content": "hi"}],
)
assert response == "ok"
assert captured["model"] == "gpt-4o"
assert "_target_order" not in captured
@pytest.mark.asyncio
async def test_text_completion_order_fallback_hop_does_not_send_target_order_upstream():
upstream_bodies: Final[list[dict]] = []
def _upstream(request: httpx.Request) -> httpx.Response:
upstream_bodies.append(json.loads(request.content))
return httpx.Response(
200,
json={
"id": "cmpl-1",
"object": "text_completion",
"created": 0,
"model": "gpt-3.5-turbo-instruct",
"choices": [{"text": "ok from order 2", "index": 0, "logprobs": None, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
)
upstream_client: Final = AsyncOpenAI(
api_key="key",
base_url="http://upstream.test",
http_client=httpx.AsyncClient(transport=httpx.MockTransport(_upstream)),
)
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "text-completion-openai/gpt-3.5-turbo-instruct",
"api_key": "key",
"mock_response": Exception("fail order 1"),
"order": 1,
},
"model_info": {"id": "1"},
},
{
"model_name": "test-model",
"litellm_params": {
"model": "text-completion-openai/gpt-3.5-turbo-instruct",
"api_key": "key",
"api_base": "http://upstream.test",
"order": 2,
},
"model_info": {"id": "2"},
},
],
num_retries=0,
)
try:
response = await router.atext_completion(model="test-model", prompt="hi", client=upstream_client)
finally:
await upstream_client.close()
assert response._hidden_params["model_id"] == "2"
assert upstream_bodies
assert all("_target_order" not in body for body in upstream_bodies)
def test_check_non_standard_fallback_format():
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,
)
# Standard formats
assert (
_check_non_standard_fallback_format([{"gpt-3.5-turbo": ["claude-3-haiku"]}])
== False
)
assert _check_non_standard_fallback_format([{"gpt-3.5-turbo": ["claude-3-haiku"]}]) == False
assert _check_non_standard_fallback_format([{"model": ["qwen-backup"]}]) == False
assert (
_check_non_standard_fallback_format(
[{"model": ["qwen-backup"], "region": ["us-east-1"]}]
)
== False
)
assert _check_non_standard_fallback_format([{"model": ["qwen-backup"], "region": ["us-east-1"]}]) == False
# Non-standard formats
assert _check_non_standard_fallback_format([{"model": "qwen-backup"}]) == True
assert (
_check_non_standard_fallback_format(
[{"model": "qwen-backup", "messages": [{"role": "user", "content": "hi"}]}]
)
== True
)
assert (
_check_non_standard_fallback_format(
[{"model": ["qwen-backup"], "api_key": "some-key"}]
)
_check_non_standard_fallback_format([{"model": "qwen-backup", "messages": [{"role": "user", "content": "hi"}]}])
== True
)
assert _check_non_standard_fallback_format([{"model": ["qwen-backup"], "api_key": "some-key"}]) == True