fix(router): keep pod routing records accurate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-10-03 22:13:56 +00:00
parent afb68249ba
commit a0966ceede
12 changed files with 268 additions and 105 deletions

View file

@ -49,6 +49,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
"fallback_access_check",
"fallback_budget_check",
"auto_router_capability_limit",
"kubernetes_pod_discovery",
}
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))

View file

@ -6999,7 +6999,7 @@ class ProxyConfig:
"Key '%s' is not a valid argument for Router.__init__(). Ignoring this key.", k
)
router = litellm.Router(
**router_params,
**router_params, # pyright: ignore[reportUnknownArgumentType] # router settings are filtered at runtime
assistants_config=assistants_config,
search_tools=search_tools,
router_general_settings=RouterGeneralSettings(

View file

@ -45,9 +45,13 @@ def _route_user_config_request(data: dict, route_type: str):
# Filter router_config to only include valid Router.__init__ arguments
# This prevents TypeError when invalid parameters are stored in the database
valid_args: Final = litellm.Router.get_valid_args()
filtered_config: Final = {k: v for k, v in router_config.items() if k in valid_args}
filtered_config: Final = {
k: v for k, v in router_config.items() if k in valid_args and k != "kubernetes_pod_discovery"
}
user_router: Final = litellm.Router(**filtered_config)
user_router: Final = litellm.Router(
**filtered_config, # pyright: ignore[reportUnknownArgumentType] # config is filtered by get_valid_args
)
ret_val: Final = getattr(user_router, f"{route_type}")(**data)
user_router.discard()
return ret_val

View file

@ -843,6 +843,7 @@ class Router:
fallback_access_check: FallbackAccessCheck | None = None,
fallback_budget_check: FallbackBudgetCheck | None = None,
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
kubernetes_pod_discovery: KubernetesPodDiscovery | None = None,
) -> None:
"""
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
@ -952,7 +953,9 @@ class Router:
cache_config: Final[dict[str, Any]] = {}
self.client_ttl = client_ttl
self.kubernetes_pod_discovery: Final = KubernetesPodDiscovery()
self.kubernetes_pod_discovery: Final = (
kubernetes_pod_discovery if kubernetes_pod_discovery is not None else KubernetesPodDiscovery()
)
if redis_url is not None or (redis_host is not None and redis_port is not None):
cache_type = "redis"
@ -1250,7 +1253,8 @@ class Router:
"""
Returns a list of valid arguments for the Router.__init__ method.
"""
arg_spec: Final = inspect.getfullargspec(Router.__init__)
router_init: Final[Callable[..., None]] = Router.__init__
arg_spec: Final = inspect.getfullargspec(router_init)
valid_args: Final = arg_spec.args + arg_spec.kwonlyargs
if "self" in valid_args:
valid_args.remove("self")
@ -4021,6 +4025,7 @@ class Router:
metadata_variable_name: Final = _get_router_metadata_variable_name(
function_name=function_name,
)
routing: Final = deployment.get(KUBERNETES_POD_ROUTING_KEY)
kwargs.setdefault(metadata_variable_name, {}).update(
{
@ -4028,9 +4033,16 @@ class Router:
"model_info": model_info,
"api_base": deployment_api_base,
"deployment_model_name": deployment_model_name,
KUBERNETES_POD_ROUTING_KEY: deployment.get(KUBERNETES_POD_ROUTING_KEY),
KUBERNETES_POD_ROUTING_KEY: routing,
}
)
other_bucket: Final = "metadata" if metadata_variable_name == "litellm_metadata" else "litellm_metadata"
other_bucket_metadata: Final = kwargs.get(other_bucket)
if isinstance(other_bucket_metadata, Mapping) and KUBERNETES_POD_ROUTING_KEY in other_bucket_metadata:
kwargs[other_bucket] = {
**other_bucket_metadata,
KUBERNETES_POD_ROUTING_KEY: routing,
}
# A retry/fallback reuses this same kwargs dict for the next deployment.
# Refund and clear any reservation the previous deployment attempt left

View file

@ -297,7 +297,7 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
# -------- every fallback entry stays reachable across hops --------
def _make_three_tier_router(**router_kwargs) -> Router:
def _make_three_tier_router(fallbacks: list[dict[str, list[str]]] | None = None) -> Router:
return Router(
model_list=[
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "sk-test"}},
@ -305,7 +305,7 @@ def _make_three_tier_router(**router_kwargs) -> Router:
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "sk-test"}},
],
num_retries=0,
**router_kwargs,
fallbacks=[] if fallbacks is None else fallbacks,
)

View file

@ -25,8 +25,16 @@ from litellm.router_utils.fallback_event_handlers import _trigger_cooldown_for_f
from litellm.types.router import AllowedFailsPolicy
def _make_router(model_list: list, **kwargs) -> Router:
return Router(model_list=model_list, **kwargs)
def _make_router(
model_list: list,
allowed_fails: int | None = None,
allowed_fails_policy: AllowedFailsPolicy | None = None,
) -> Router:
return Router(
model_list=model_list,
allowed_fails=allowed_fails,
allowed_fails_policy=allowed_fails_policy,
)
class TestDeploymentLevelAllowedFails:

View file

@ -12,6 +12,7 @@ from litellm.proxy.common_utils.discoverable_model_filter import (
discoverable_rows,
undiscoverable_model_names,
)
from litellm.types.router import RouterModelGroupAliasItem
def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
@ -22,8 +23,11 @@ def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
}
def _router(*deployments, **router_kwargs) -> Router:
return Router(model_list=list(deployments), **router_kwargs)
def _router(
*deployments,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
) -> Router:
return Router(model_list=list(deployments), model_group_alias=model_group_alias)
def _non_admin() -> UserAPIKeyAuth:

View file

@ -17,6 +17,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
initialize_guardrail,
)
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.router import RouterModelGroupAliasItem
from litellm.types.utils import LLM_AS_A_JUDGE_GUARDRAIL_CALL_ORIGIN
# ---------------------------------------------------------------------------
@ -598,19 +599,22 @@ def _judge_response_mock() -> MagicMock:
return MagicMock(choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(90.0))))])
def _real_router(model_list, **router_kwargs):
def _real_router(
model_list,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
):
"""Build a real Router so the router-membership decision is exercised for
real (wildcards, model_group_alias, exact names), stubbing only the outbound
completion so no network call is made."""
from litellm import Router
router = Router(model_list=model_list, **router_kwargs)
router = Router(model_list=model_list, model_group_alias=model_group_alias)
router.acompletion = AsyncMock(return_value=_judge_response_mock())
return router
@pytest.mark.parametrize(
"model_list, router_kwargs, judge_model",
"model_list, model_group_alias, judge_model",
[
(
[
@ -619,12 +623,12 @@ def _real_router(model_list, **router_kwargs):
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
}
],
{},
None,
"my-judge-alias",
),
(
[{"model_name": "anthropic/*", "litellm_params": {"model": "anthropic/*", "api_key": "sk-ant-test"}}],
{},
None,
"anthropic/claude-sonnet-4-6",
),
(
@ -634,7 +638,7 @@ def _real_router(model_list, **router_kwargs):
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
}
],
{"model_group_alias": {"my-judge-alias": "backing-group"}},
{"my-judge-alias": "backing-group"},
"my-judge-alias",
),
(
@ -644,7 +648,7 @@ def _real_router(model_list, **router_kwargs):
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
}
],
{"model_group_alias": {"my-judge-alias": {"model": "backing-group", "hidden": True}}},
{"my-judge-alias": {"model": "backing-group", "hidden": True}},
"my-judge-alias",
),
],
@ -653,13 +657,16 @@ def _real_router(model_list, **router_kwargs):
@pytest.mark.asyncio
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
async def test_judge_routes_through_router_for_router_served_model(
mock_sdk_completion, model_list, router_kwargs, judge_model
mock_sdk_completion,
model_list,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None,
judge_model,
):
"""Any judge_model the Router can serve must resolve its credentials via the
Router. Wildcard and alias shapes regress the naive `judge_model in
get_model_names()` check, which reports patterns/aliases literally and so
routes a servable model to the SDK, where deployment creds do not resolve."""
router = _real_router(model_list, **router_kwargs)
router = _real_router(model_list, model_group_alias)
guardrail = _make_guardrail(judge_model=judge_model, router_provider=lambda: router)
inputs = {"texts": ["good response"]}
request_data: dict = {"messages": [{"role": "user", "content": "hi"}], "metadata": {}}

View file

@ -20,6 +20,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy import proxy_server
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.router import RouterModelGroupAliasItem
class _Gate(CustomLogger):
@ -70,8 +71,12 @@ def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
}
def _install_router(monkeypatch, *deployments, **router_kwargs) -> Router:
router = Router(model_list=list(deployments), **router_kwargs)
def _install_router(
monkeypatch,
*deployments,
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
) -> Router:
router = Router(model_list=list(deployments), model_group_alias=model_group_alias)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list)
monkeypatch.setattr(proxy_server, "prisma_client", None)

View file

@ -100,31 +100,6 @@ def _empty_proxy_environment() -> Mapping[str, str]:
return {}
def _patch_router_pod_discovery(
monkeypatch: pytest.MonkeyPatch,
*,
clock: Callable[[], float] | None = None,
refresh_interval_seconds: float | None = None,
) -> None:
def create_discovery() -> KubernetesPodDiscovery:
if refresh_interval_seconds is not None and clock is not None:
return KubernetesPodDiscovery(
refresh_interval_seconds=refresh_interval_seconds,
clock=clock,
proxy_environment=_empty_proxy_environment,
)
if refresh_interval_seconds is not None:
return KubernetesPodDiscovery(
refresh_interval_seconds=refresh_interval_seconds,
proxy_environment=_empty_proxy_environment,
)
if clock is not None:
return KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment)
return KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
monkeypatch.setattr("litellm.router.KubernetesPodDiscovery", create_discovery)
def _stub_sync_dns(monkeypatch: pytest.MonkeyPatch, *ips: str) -> None:
def getaddrinfo(
host: str | None,
@ -775,16 +750,20 @@ async def test_router_sends_pod_hosts_without_forwarding_discovery_flag(
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(
r"http://(?:10\.0\.0\.1|10\.0\.0\.2|vllm-headless\.ns\.svc\.cluster\.local):8000/v1/chat/completions"
)
).mock(return_value=httpx.Response(200, json=_CHAT_RESPONSE))
discovery_enabled_router: Final = Router(model_list=[_deployment()])
control_router: Final = Router(model_list=[_deployment(kubernetes_pod_discovery=None)])
discovery_enabled_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment),
)
control_router: Final = Router(
model_list=[_deployment(kubernetes_pod_discovery=None)],
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment),
)
cached_async_client: Final = AsyncOpenAI(api_key="fake", base_url=_SERVICE_URL)
cached_sync_client: Final = OpenAI(api_key="fake", base_url=_SERVICE_URL)
discovery_enabled_router.cache.set_cache(
@ -855,14 +834,18 @@ async def test_router_session_requests_stick_to_one_pod_for_sync_and_async(
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
async_router: Final = Router(model_list=[_deployment()])
sync_router: Final = Router(model_list=[_deployment()])
async_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=lambda: 0.0, proxy_environment=_empty_proxy_environment),
)
sync_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=lambda: 0.0, proxy_environment=_empty_proxy_environment),
)
async_client: Final = _cache_async_client(async_router)
sync_client: Final = _cache_sync_client(sync_router)
for _ in range(6):
@ -899,14 +882,18 @@ async def test_router_session_mapping_repeats_across_pods(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
session_ids: Final = tuple(f"session-{index}" for index in range(30))
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for session_id in session_ids:
await _router_acompletion(router, metadata={"session_id": session_id})
@ -942,13 +929,17 @@ async def test_generated_session_id_uses_round_robin_pod_selection(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(
@ -982,13 +973,17 @@ async def test_empty_session_id_uses_round_robin_pod_selection(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(router, metadata={"session_id": ""})
@ -1016,13 +1011,17 @@ async def test_session_requests_do_not_advance_round_robin_cursor(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
await _router_acompletion(router)
await _router_acompletion(router, metadata={"session_id": "s1"})
@ -1063,16 +1062,18 @@ async def test_session_pod_membership_change_remaps_only_affected_sessions(
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
session_ids: Final = tuple(f"session-{index}" for index in range(60))
_patch_router_pod_discovery(
monkeypatch,
refresh_interval_seconds=10,
clock=_clock((0.0,) * 62 + (10.0,) * 10),
)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
refresh_interval_seconds=10,
clock=_clock((0.0,) * 62 + (10.0,) * 10),
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for session_id in session_ids:
await _router_acompletion(router, metadata={"session_id": session_id})
@ -1132,14 +1133,24 @@ async def test_custom_routing_strategy_keeps_sync_and_async_sessions_on_one_pod(
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
async_router: Final = Router(model_list=[_deployment()])
sync_router: Final = Router(model_list=[_deployment()])
async_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
sync_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
async_router.set_custom_routing_strategy(_ModelListRoutingStrategy(async_router))
sync_router.set_custom_routing_strategy(_ModelListRoutingStrategy(sync_router))
async_client: Final = _cache_async_client(async_router)
@ -1176,13 +1187,17 @@ async def test_top_level_litellm_session_id_keeps_requests_on_one_pod(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(router, litellm_session_id="s1")
@ -1210,13 +1225,17 @@ async def test_generated_metadata_ignores_top_level_litellm_session_id(
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=lambda: 0.0,
proxy_environment=_empty_proxy_environment,
),
)
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(
@ -1267,16 +1286,26 @@ async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_req
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(
r"http://(?:10\.0\.0\.1|10\.0\.0\.2|vllm-headless\.ns\.svc\.cluster\.local):8000/v1/chat/completions"
)
).mock(return_value=httpx.Response(200, json=_CHAT_RESPONSE))
async_router: Final = Router(model_list=[_deployment()])
sync_router: Final = Router(model_list=[_deployment()])
async_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
),
)
sync_router: Final = Router(
model_list=[_deployment()],
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
),
)
async_router.set_custom_routing_strategy(_ModelListRoutingStrategy(async_router))
sync_router.set_custom_routing_strategy(_ModelListRoutingStrategy(sync_router))
for _ in range(4):
@ -1328,13 +1357,18 @@ async def test_router_async_retry_uses_next_discovered_pod(monkeypatch: pytest.M
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
router: Final = Router(
model_list=[_deployment()],
num_retries=1,
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
),
)
response: Final = await router.acompletion(
model="gpu-model",
@ -1368,13 +1402,18 @@ def test_router_sync_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPa
monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
router: Final = Router(
model_list=[_deployment()],
num_retries=1,
kubernetes_pod_discovery=KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
),
)
response: Final = router.completion(
model="gpu-model",
@ -1443,12 +1482,14 @@ async def test_router_session_retry_uses_next_rendezvous_pod(monkeypatch: pytest
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch)
session_id: Final = "retry-session"
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
router: Final = Router(
model_list=[_deployment()],
num_retries=1,
kubernetes_pod_discovery=KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment),
)
response: Final = await router.acompletion(
model="gpu-model",

View file

@ -11,7 +11,7 @@ import warnings
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final, Literal
from typing import Final, Literal, TypedDict
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -19,6 +19,7 @@ import openai
import pytest
import respx
from fastapi import HTTPException
from typing_extensions import ReadOnly
import litellm
from litellm import Router
@ -7895,6 +7896,64 @@ def test_update_kwargs_with_deployment_clears_pod_routing_on_non_discovery_fallb
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] is None
def test_update_kwargs_with_deployment_synchronizes_existing_litellm_metadata_pod_routing():
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
routing_record: Final = {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "session_affinity",
}
spoofed_record: Final = {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.9",
"pod_count": 3,
"selection": "round_robin",
}
deployment: Final = {
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
"model_info": {"id": "discovery-id"},
KUBERNETES_POD_ROUTING_KEY: routing_record,
}
original_litellm_metadata: Final = {KUBERNETES_POD_ROUTING_KEY: spoofed_record}
kwargs: Final = {"metadata": {}, "litellm_metadata": original_litellm_metadata}
router: Final = litellm.Router(model_list=[deployment])
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="completion")
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] == routing_record
assert kwargs["litellm_metadata"][KUBERNETES_POD_ROUTING_KEY] == routing_record
assert kwargs["litellm_metadata"] is not original_litellm_metadata
assert original_litellm_metadata[KUBERNETES_POD_ROUTING_KEY] == spoofed_record
def test_update_kwargs_with_non_discovery_deployment_clears_other_pod_routing_bucket():
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
previous_record: Final = {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "session_affinity",
}
deployment: Final = {
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
"model_info": {"id": "fallback-id"},
}
original_litellm_metadata: Final = {KUBERNETES_POD_ROUTING_KEY: previous_record}
kwargs: Final = {"metadata": {}, "litellm_metadata": original_litellm_metadata}
router: Final = litellm.Router(model_list=[deployment])
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="completion")
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] is None
assert kwargs["litellm_metadata"][KUBERNETES_POD_ROUTING_KEY] is None
assert original_litellm_metadata[KUBERNETES_POD_ROUTING_KEY] == previous_record
def test_combine_fallback_usage():
"""Test that _combine_fallback_usage merges partial and fallback usage."""
from litellm.router import Router
@ -13728,8 +13787,14 @@ def _anthropic_messages_make_wrapper() -> FallbackAwareAnthropicMessagesStream:
return FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), object())
def _anthropic_messages_make_router(**router_kwargs) -> Router:
router_kwargs.setdefault("fallbacks", [{"primary": ["fallback"]}])
class _AnthropicMessagesRouterKwargs(TypedDict, total=False):
fallbacks: ReadOnly[list[dict[str, list[str]]] | None]
content_policy_fallbacks: ReadOnly[list[dict[str, list[str]]] | None]
enable_weighted_failover: ReadOnly[bool]
def _anthropic_messages_make_router(router_kwargs: _AnthropicMessagesRouterKwargs) -> Router:
fallbacks: Final = router_kwargs.get("fallbacks", [{"primary": ["fallback"]}])
return Router(
model_list=[
{
@ -13746,7 +13811,9 @@ def _anthropic_messages_make_router(**router_kwargs) -> Router:
},
},
],
**router_kwargs,
fallbacks=fallbacks,
content_policy_fallbacks=router_kwargs.get("content_policy_fallbacks"),
enable_weighted_failover=router_kwargs.get("enable_weighted_failover", False),
)
@ -14115,8 +14182,12 @@ def _anthropic_messages_two_order_primary_model_list() -> list:
pytest.param({"fallbacks": None, "enable_weighted_failover": True}, {"model": "primary"}, True, id="weighted-failover"),
],
)
def test_anthropic_messages_stream_can_fall_back_direct_call(router_kwargs, request_kwargs, expected):
router = _anthropic_messages_make_router(**router_kwargs)
def test_anthropic_messages_stream_can_fall_back_direct_call(
router_kwargs: _AnthropicMessagesRouterKwargs,
request_kwargs: dict[str, object],
expected: bool,
) -> None:
router = _anthropic_messages_make_router(router_kwargs)
assert router._anthropic_messages_stream_can_fall_back("primary", request_kwargs) is expected

View file

@ -359,10 +359,20 @@ class TestNoProviderRetryAmplification:
await session.aclose()
@staticmethod
def _router(api_base: str, litellm_params: dict, **router_kwargs) -> Router:
def _router(
api_base: str,
litellm_params: dict,
*,
num_retries: int,
retry_policy: RetryPolicy | None = None,
) -> Router:
params = {"model": "openai/gpt-4o-mini", "api_base": api_base, "api_key": "sk-fake"}
params.update(litellm_params)
return Router(model_list=[{"model_name": "mock", "litellm_params": params}], **router_kwargs)
return Router(
model_list=[{"model_name": "mock", "litellm_params": params}],
num_retries=num_retries,
retry_policy=retry_policy,
)
async def _call_and_count(self, router: Router, **call_kwargs) -> int:
counter = self._install_counting_upstream()