mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
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:
parent
4ba8517134
commit
bd0b9c78bd
8 changed files with 220 additions and 11 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 {})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]] = [
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue