diff --git a/litellm/constants.py b/litellm/constants.py index 540ac1fb6ae..df579873fd2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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)) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9af080f1b50..5e16e792c00 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 7da09ddcb68..0fdd3195d29 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 21e91bf763d..bd4cde32edb 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py index 5370089eef5..1b1e157257b 100644 --- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py +++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py @@ -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, ) diff --git a/tests/router_unit_tests/test_router_cooldown_per_deployment.py b/tests/router_unit_tests/test_router_cooldown_per_deployment.py index 964782c348b..062c0182b88 100644 --- a/tests/router_unit_tests/test_router_cooldown_per_deployment.py +++ b/tests/router_unit_tests/test_router_cooldown_per_deployment.py @@ -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: diff --git a/tests/unit/proxy/common_utils/test_discoverable_model_filter.py b/tests/unit/proxy/common_utils/test_discoverable_model_filter.py index 17619afbb07..2f583e60266 100644 --- a/tests/unit/proxy/common_utils/test_discoverable_model_filter.py +++ b/tests/unit/proxy/common_utils/test_discoverable_model_filter.py @@ -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: diff --git a/tests/unit/proxy/guardrails/test_llm_as_a_judge.py b/tests/unit/proxy/guardrails/test_llm_as_a_judge.py index 6e00958eba4..c9ae83237a3 100644 --- a/tests/unit/proxy/guardrails/test_llm_as_a_judge.py +++ b/tests/unit/proxy/guardrails/test_llm_as_a_judge.py @@ -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": {}} diff --git a/tests/unit/proxy/test_model_list_callback_filter.py b/tests/unit/proxy/test_model_list_callback_filter.py index 00fbfee24ed..c632786418b 100644 --- a/tests/unit/proxy/test_model_list_callback_filter.py +++ b/tests/unit/proxy/test_model_list_callback_filter.py @@ -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) diff --git a/tests/unit/router_utils/test_kubernetes_pod_discovery.py b/tests/unit/router_utils/test_kubernetes_pod_discovery.py index 47979dbf049..e2344f63595 100644 --- a/tests/unit/router_utils/test_kubernetes_pod_discovery.py +++ b/tests/unit/router_utils/test_kubernetes_pod_discovery.py @@ -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", diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 22f76370ae9..ca678e97d39 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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 diff --git a/tests/unit/test_router_per_deployment_num_retries.py b/tests/unit/test_router_per_deployment_num_retries.py index 99ad7c224f8..fb304269e69 100644 --- a/tests/unit/test_router_per_deployment_num_retries.py +++ b/tests/unit/test_router_per_deployment_num_retries.py @@ -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()