diff --git a/litellm/constants.py b/litellm/constants.py index 69ff3cf7a5e..7b2dc698984 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -6,6 +6,12 @@ from typing import Final, Literal from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) +KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS: Final = float( + os.getenv("KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS", "5") +) +KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS: Final = float( + os.getenv("KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS", "300") +) AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) AZURE_OPENAI_AUDIO_PROVIDERS: Final = frozenset({"azure", "azure_ai"}) ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5)) diff --git a/litellm/router.py b/litellm/router.py index 0b9f12c8da3..034689d638a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -227,6 +227,13 @@ from litellm.router_utils.handle_error import ( send_llm_exception_alert, ) from litellm.router_utils.health_state_cache import DeploymentHealthCache +from litellm.router_utils.kubernetes_pod_discovery import ( + KubernetesPodDiscovery, + async_resolve_pods_after, + async_resolve_pods_after_bound, + resolve_pods_after, + resolve_pods_after_bound, +) from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( DeploymentAffinityCheck, warn_on_unknown_model_group_affinity_flags, @@ -974,6 +981,7 @@ class Router: cache_config: Final[dict[str, Any]] = {} self.client_ttl = client_ttl + self.kubernetes_pod_discovery: Final = KubernetesPodDiscovery() if redis_url is not None or (redis_host is not None and redis_port is not None): cache_type = "redis" @@ -12389,6 +12397,13 @@ class Router: client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) return client elif client_type == "async": + if ( + deployment["litellm_params"].get( # pyright: ignore[reportUnknownMemberType] # legacy mapping + "kubernetes_pod_discovery" + ) + is True + ): + return None if kwargs.get("stream") is True: cache_key = f"{model_id}_stream_async_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) @@ -12398,6 +12413,13 @@ class Router: client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) return client else: + if ( + deployment["litellm_params"].get( # pyright: ignore[reportUnknownMemberType] # legacy mapping + "kubernetes_pod_discovery" + ) + is True + ): + return None if kwargs.get("stream") is True: cache_key = f"{model_id}_stream_client" client = self.cache.get_cache(key=cache_key, parent_otel_span=parent_otel_span) @@ -13356,7 +13378,7 @@ class Router: Router._pop_effort_from_nested_carrier(request_kwargs, "output_config") Router._pop_effort_from_nested_carrier(request_kwargs, "reasoning") - async def async_get_available_deployment( + async def _async_get_available_deployment_unresolved( self, model: str, request_kwargs: dict, @@ -13393,7 +13415,7 @@ class Router: and self.routing_strategy != "latency-based-routing" and self.routing_strategy != "least-busy" ): # prevent regressions for other routing strategies, that don't have async get available deployments implemented. - return self.get_available_deployment( + return self._get_available_deployment_unresolved( model=model, messages=messages, input=input, @@ -13532,6 +13554,12 @@ class Router: ) raise e + async_get_available_deployment = ( # pyright: ignore[reportUnknownVariableType] # legacy selector return + async_resolve_pods_after( + _async_get_available_deployment_unresolved # pyright: ignore[reportUnknownArgumentType] # legacy type + ) + ) + async def async_get_available_deployment_for_pass_through( self, model: str, @@ -14272,7 +14300,7 @@ class Router: } return cast(StandardLoggingRoutingDecision, kept) # cast-ok: dropping optional keys preserves the type - def get_available_deployment( + def _get_available_deployment_unresolved( self, model: str, messages: list[dict[str, str]] | None = None, @@ -14440,6 +14468,12 @@ class Router: phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return deployment + get_available_deployment = ( # pyright: ignore[reportUnknownVariableType] # legacy selector return + resolve_pods_after( + _get_available_deployment_unresolved # pyright: ignore[reportUnknownArgumentType] # legacy annotations + ) + ) + def get_available_deployment_for_pass_through( self, model: str, @@ -14855,15 +14889,22 @@ class Router: CustomRoutingStrategy: litellm.router.CustomRoutingStrategyBase """ + strategy: Final = CustomRoutingStrategy setattr( self, "get_available_deployment", - CustomRoutingStrategy.get_available_deployment, + resolve_pods_after_bound( + strategy.get_available_deployment, # pyright: ignore[reportUnknownArgumentType] # legacy stub + self.kubernetes_pod_discovery, + ), ) setattr( self, "async_get_available_deployment", - CustomRoutingStrategy.async_get_available_deployment, + async_resolve_pods_after_bound( + strategy.async_get_available_deployment, # pyright: ignore[reportUnknownArgumentType] # legacy stub + self.kubernetes_pod_discovery, + ), ) def _reset_custom_routing_strategy(self) -> None: diff --git a/litellm/router_utils/kubernetes_pod_discovery.py b/litellm/router_utils/kubernetes_pod_discovery.py new file mode 100644 index 00000000000..00a3a6b55ef --- /dev/null +++ b/litellm/router_utils/kubernetes_pod_discovery.py @@ -0,0 +1,350 @@ +import asyncio +import ipaddress +import socket +import threading +import time +import urllib.request +from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence +from dataclasses import dataclass, replace +from functools import wraps +from types import MappingProxyType +from typing import Concatenate, Final, ParamSpec, Protocol, TypeAlias, TypeVar, cast + +import httpx + +from litellm._logging import verbose_router_logger +from litellm.constants import ( + KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS, + KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS, +) + +_SocketAddress: TypeAlias = tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes] +_AddressInfo: TypeAlias = tuple[socket.AddressFamily, socket.SocketKind, int, str, _SocketAddress] +_CacheKey: TypeAlias = tuple[str, int | None] +_DeploymentT = TypeVar("_DeploymentT") + + +@dataclass(frozen=True, slots=True) +class _PodSet: + ips: tuple[str, ...] + resolved_at: float + last_used_at: float + cursor: int = 0 + + +class KubernetesPodDiscovery: + def __init__( + self, + refresh_interval_seconds: float = KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS, + idle_eviction_seconds: float = KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS, + proxy_environment: Callable[[], Mapping[str, str]] = urllib.request.getproxies_environment, + clock: Callable[[], float] = time.monotonic, + ) -> None: + self.refresh_interval_seconds = refresh_interval_seconds + self.idle_eviction_seconds = idle_eviction_seconds + self.proxy_environment = proxy_environment + self.clock = clock + self._cache: Mapping[_CacheKey, _PodSet] = MappingProxyType({}) + self._refreshing: frozenset[_CacheKey] = frozenset() + self._https_warning_hosts: frozenset[str] = frozenset() + self._proxy_warning_hosts: frozenset[str] = frozenset() + self._lock = threading.Lock() + + def _begin_refresh(self, key: _CacheKey, now: float) -> tuple[_PodSet | None, bool]: + with self._lock: + self._cache = MappingProxyType( + { + cache_key: pod_set + for cache_key, pod_set in self._cache.items() + if cache_key in self._refreshing or now - pod_set.last_used_at <= self.idle_eviction_seconds + } + ) + cached: Final = self._cache.get(key) + refreshing: Final = key in self._refreshing + if cached is not None and (now - cached.resolved_at < self.refresh_interval_seconds or refreshing): + return cached, False + if refreshing: + return None, False + self._refreshing = self._refreshing | {key} + return None, True + + def resolve_deployment(self, deployment: _DeploymentT) -> _DeploymentT: + eligible: Final = self._eligible(deployment) + if eligible is None: + return deployment + url, key, deployment_mapping = eligible + host, port = key + now: Final = self.clock() + cached, should_refresh = self._begin_refresh(key, now) + if cached is not None: + return self._deployment_with_cached_ip(deployment, deployment_mapping, key, url, now) + if not should_refresh: + return deployment + + try: + records: Final = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM) + except OSError as error: + return self._apply_lookup(deployment, deployment_mapping, key, url, error, now) + else: + return self._apply_lookup(deployment, deployment_mapping, key, url, records, now) + finally: + with self._lock: + self._refreshing = self._refreshing - {key} + + async def async_resolve_deployment(self, deployment: _DeploymentT) -> _DeploymentT: + eligible: Final = self._eligible(deployment) + if eligible is None: + return deployment + url, key, deployment_mapping = eligible + host, port = key + now: Final = self.clock() + cached, should_refresh = self._begin_refresh(key, now) + if cached is not None: + return self._deployment_with_cached_ip(deployment, deployment_mapping, key, url, now) + if not should_refresh: + return deployment + + try: + result: Final = await self._async_getaddrinfo(host, port) + return self._apply_lookup(deployment, deployment_mapping, key, url, result, now) + finally: + with self._lock: + self._refreshing = self._refreshing - {key} + + def _eligible(self, deployment: object) -> tuple[httpx.URL, _CacheKey, Mapping[str, object]] | None: + if not isinstance(deployment, Mapping): + return None + typed_deployment: Final = cast( # cast-ok: runtime Mapping check precedes typed lookup + Mapping[str, object], deployment + ) + target: Final = self._target(typed_deployment) + if target is None: + return None + url, host, port, target_deployment = target + if url.scheme == "https": + self._warn_https_once(host) + return None + if url.scheme != "http" or self._is_ip_literal(host): + return None + return (url, (host, port), target_deployment) + + @staticmethod + def _target( + deployment_mapping: Mapping[str, object], + ) -> tuple[httpx.URL, str, int | None, Mapping[str, object]] | None: + raw_params: Final = deployment_mapping.get("litellm_params") + if not isinstance(raw_params, Mapping): + return None + params: Final = cast( # cast-ok: runtime Mapping validation precedes the string field checks + Mapping[str, object], raw_params + ) + if params.get("kubernetes_pod_discovery") is not True: + return None + raw_api_base: Final = params.get("api_base") + if not isinstance(raw_api_base, str): + return None + try: + url: Final = httpx.URL(raw_api_base) + except httpx.InvalidURL: + return None + return (url, url.host, url.port, deployment_mapping) + + @staticmethod + async def _async_getaddrinfo(host: str, port: int | None) -> Sequence[_AddressInfo] | OSError: + try: + return await asyncio.get_running_loop().getaddrinfo(host, port, type=socket.SOCK_STREAM) + except OSError as error: + return error + + def _apply_lookup( + self, + deployment: _DeploymentT, + deployment_mapping: Mapping[str, object], + key: _CacheKey, + url: httpx.URL, + result: Sequence[_AddressInfo] | OSError, + now: float, + ) -> _DeploymentT: + if isinstance(result, OSError): + if isinstance(result, socket.gaierror) and self._is_authoritative_no_pods(result): + self._store(key, (), self.clock()) + return deployment + self._stamp_failed_refresh(key, self.clock()) + verbose_router_logger.debug("Kubernetes pod discovery DNS refresh failed for %s: %s", key[0], result) + return self._deployment_with_cached_ip(deployment, deployment_mapping, key, url, now) + ips: Final = self._pod_ips(result) + self._store(key, ips, self.clock()) + if not ips: + return deployment + return self._deployment_with_cached_ip(deployment, deployment_mapping, key, url, now) + + @staticmethod + def _is_ip_literal(host: str) -> bool: + try: + ipaddress.ip_address(host) + except ValueError: + return False + return True + + @staticmethod + def _is_authoritative_no_pods(error: socket.gaierror) -> bool: + no_pod_errors: Final = (socket.EAI_NONAME, socket.EAI_NODATA) + return error.errno in no_pod_errors + + @staticmethod + def _pod_ip(record: _AddressInfo) -> str | None: + address: Final = record[4][0] + return address if isinstance(address, str) else None + + @classmethod + def _pod_ips(cls, records: Sequence[_AddressInfo]) -> tuple[str, ...]: + addresses: Final = tuple(cls._pod_ip(record) for record in records) + return tuple(sorted({address for address in addresses if address is not None})) + + def _store(self, key: _CacheKey, ips: tuple[str, ...], resolved_at: float) -> None: + with self._lock: + cached: Final = self._cache.get(key) + cursor: Final = cached.cursor % len(ips) if cached is not None and cached.ips and ips else 0 + self._cache = MappingProxyType( + { + **self._cache, + key: _PodSet(ips=ips, resolved_at=resolved_at, last_used_at=resolved_at, cursor=cursor), + } + ) + + def _stamp_failed_refresh(self, key: _CacheKey, resolved_at: float) -> None: + with self._lock: + cached: Final = self._cache.get(key) + if cached is not None: + self._cache = MappingProxyType({**self._cache, key: replace(cached, resolved_at=resolved_at)}) + + def _next_ip(self, key: _CacheKey, now: float) -> str | None: + with self._lock: + cached: Final = self._cache.get(key) + if cached is None or not cached.ips: + return None + ip: Final = cached.ips[cached.cursor] + next_cursor: Final = (cached.cursor + 1) % len(cached.ips) + self._cache = MappingProxyType( + { + **self._cache, + key: replace(cached, cursor=next_cursor, last_used_at=max(cached.last_used_at, now)), + } + ) + return ip + + def _deployment_with_cached_ip( + self, + deployment: _DeploymentT, + deployment_mapping: Mapping[str, object], + key: _CacheKey, + url: httpx.URL, + now: float, + ) -> _DeploymentT: + ip: Final = self._next_ip(key, now) + if ip is None: + return deployment + if self._proxy_bypasses_only_service_host(key[0], ip): + self._warn_proxy_once(key[0]) + return deployment + params: Final = deployment_mapping.get("litellm_params") + if not isinstance(params, Mapping): + return deployment + typed_params: Final = cast( # cast-ok: runtime Mapping validation precedes the copied provider parameters + Mapping[str, object], params + ) + return cast( # cast-ok: resolving produces a copied dict for the selector's Mapping type + _DeploymentT, + { + **deployment_mapping, + "litellm_params": {**typed_params, "api_base": str(url.copy_with(host=ip))}, + }, + ) + + def _proxy_bypasses_only_service_host(self, host: str, ip: str) -> bool: + proxies: Final = self.proxy_environment() + if not (proxies.get("http") or proxies.get("all")): + return False + no_proxy: Final = proxies.get("no", "") + entries: Final = tuple(entry.strip() for entry in no_proxy.split(",")) + service_bypassed: Final = "*" in entries or urllib.request.proxy_bypass_environment(host, dict(proxies)) + pod_bypassed: Final = urllib.request.proxy_bypass_environment(ip, dict(proxies)) + return service_bypassed and not pod_bypassed + + def _warn_https_once(self, host: str) -> None: + with self._lock: + should_warn: Final = host not in self._https_warning_hosts + if should_warn: + self._https_warning_hosts = self._https_warning_hosts | {host} + if should_warn: + verbose_router_logger.warning( + "Kubernetes pod discovery is disabled for HTTPS host %s because substituting an IP " + "breaks TLS hostname verification", + host, + ) + + def _warn_proxy_once(self, host: str) -> None: + with self._lock: + should_warn: Final = host not in self._proxy_warning_hosts + if should_warn: + self._proxy_warning_hosts = self._proxy_warning_hosts | {host} + if should_warn: + verbose_router_logger.warning( + "Kubernetes pod discovery keeps the service hostname %s: NO_PROXY exempts it but not its pod IPs", + host, + ) + + +class _HasPodDiscovery(Protocol): + kubernetes_pod_discovery: KubernetesPodDiscovery + + +_P = ParamSpec("_P") +_R = TypeVar("_R") +_S = TypeVar("_S", bound=_HasPodDiscovery) + + +def resolve_pods_after( + fn: Callable[Concatenate[_S, _P], _R], +) -> Callable[Concatenate[_S, _P], _R]: + @wraps(fn) + def wrapped(self: _S, *args: _P.args, **kwargs: _P.kwargs) -> _R: + deployment: Final = fn(self, *args, **kwargs) + return self.kubernetes_pod_discovery.resolve_deployment(deployment) + + return wrapped + + +def async_resolve_pods_after( + fn: Callable[Concatenate[_S, _P], Awaitable[_R]], +) -> Callable[Concatenate[_S, _P], Coroutine[object, object, _R]]: + @wraps(fn) + async def wrapped(self: _S, *args: _P.args, **kwargs: _P.kwargs) -> _R: + deployment: Final = await fn(self, *args, **kwargs) + return await self.kubernetes_pod_discovery.async_resolve_deployment(deployment) + + return wrapped + + +def resolve_pods_after_bound( + fn: Callable[_P, _R], + discovery: KubernetesPodDiscovery, +) -> Callable[_P, _R]: + @wraps(fn) + def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: + deployment: Final = fn(*args, **kwargs) + return discovery.resolve_deployment(deployment) + + return wrapped + + +def async_resolve_pods_after_bound( + fn: Callable[_P, Awaitable[_R]], + discovery: KubernetesPodDiscovery, +) -> Callable[_P, Coroutine[object, object, _R]]: + @wraps(fn) + async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: + deployment: Final = await fn(*args, **kwargs) + return await discovery.async_resolve_deployment(deployment) + + return wrapped diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 9119fa16b5d..c19c4a76dc7 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -175,6 +175,7 @@ class DeploymentOptions: tpm: int | None = None itpm: int | None = None otpm: int | None = None + kubernetes_pod_discovery: bool | None = None default_api_key_rpm_limit: int | None = None default_api_key_tpm_limit: int | None = None max_parallel_requests: int | None = None diff --git a/litellm/types/router.py b/litellm/types/router.py index e131a7184d4..bf830c825a0 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -653,6 +653,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): use_chat_completions_api: bool | None ## PASS-THROUGH ENDPOINTS ## use_in_pass_through: bool | None + kubernetes_pod_discovery: ReadOnly[bool | None] litellm_credential_name: str | None ## UNIFIED PROJECT/REGION ## region_name: str | None diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 574336791de..c8137d67942 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -42,9 +42,9 @@ def get_all_functions_called_in_tests(base_dir): print("dir_path: ", dir_path) for root, _, files in os.walk(dir_path): for file in files: - if file.endswith(".py") and "router" in file.lower(): + file_path = os.path.join(root, file) + if file.endswith(".py") and "router" in os.path.relpath(file_path, dir_path).lower(): print("file: ", file) - file_path = os.path.join(root, file) with open(file_path, "r", encoding="utf-8") as f: try: tree = ast.parse(f.read()) diff --git a/tests/unit/router_utils/test_kubernetes_pod_discovery.py b/tests/unit/router_utils/test_kubernetes_pod_discovery.py new file mode 100644 index 00000000000..eea2da4e388 --- /dev/null +++ b/tests/unit/router_utils/test_kubernetes_pod_discovery.py @@ -0,0 +1,916 @@ +import asyncio +import copy +import json +import logging +import re +import socket +import threading +from collections.abc import Callable, Mapping +from concurrent.futures import Future +from itertools import chain, count, repeat +from typing import Final, TypeAlias, cast + +import httpx +import pytest +import respx +from openai import AsyncOpenAI, OpenAI + +import litellm +from litellm.router import CustomRoutingStrategyBase, Router +from litellm.router_utils.kubernetes_pod_discovery import KubernetesPodDiscovery + +_SERVICE_URL: Final = "http://vllm-headless.ns.svc.cluster.local:8000/v1" +_SocketAddress: TypeAlias = tuple[str, int] | tuple[str, int, int, int] +_AddrInfo: TypeAlias = tuple[socket.AddressFamily, socket.SocketKind, int, str, _SocketAddress] +_NO_POD_ERRNOS: Final = tuple(sorted((socket.EAI_NONAME, socket.EAI_NODATA))) +_CHAT_RESPONSE: Final = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1, + "model": "my-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} + + +def _deployment(api_base: str = _SERVICE_URL, kubernetes_pod_discovery: bool | None = True) -> dict[str, object]: + params: Final = { + "model": "openai/my-model", + "api_base": api_base, + "api_key": "fake", + **({"kubernetes_pod_discovery": kubernetes_pod_discovery} if kubernetes_pod_discovery is not None else {}), + } + return { + "model_name": "gpu-model", + "litellm_params": params, + "model_info": {"id": "registered-id"}, + } + + +class _ModelListRoutingStrategy(CustomRoutingStrategyBase): + def __init__(self, router: Router) -> None: + self._router: Final = router + + def get_available_deployment(self, *args: object, **kwargs: object) -> Mapping[str, object]: + return self._router.model_list[0] + + async def async_get_available_deployment(self, *args: object, **kwargs: object) -> Mapping[str, object]: + return self._router.model_list[0] + + +def _record(ip: str, port: int) -> _AddrInfo: + address: Final = (ip, port, 0, 0) if ":" in ip else (ip, port) + family: Final = socket.AF_INET6 if ":" in ip else socket.AF_INET + return (family, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", address) + + +def _records(*ips: str, port: int = 8000) -> list[_AddrInfo]: + return [_record(ip, port) for ip in ips] + + +def _clock(ticks: tuple[float, ...]) -> Callable[[], float]: + timestamps: Final = iter(ticks) + + def monotonic() -> float: + return next(timestamps) + + return monotonic + + +def _proxy_environment(no_proxy: str) -> Callable[[], Mapping[str, str]]: + return lambda: {"http": "http://proxy.example:8080", "no": no_proxy} + + +def _empty_proxy_environment() -> Mapping[str, str]: + return {} + + +def _stub_sync_dns(monkeypatch: pytest.MonkeyPatch, *ips: str) -> None: + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + return _records(*ips, port=int(port or 8000)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + + +def _api_base(deployment: Mapping[str, object]) -> str: + params: Final = deployment["litellm_params"] + assert isinstance(params, Mapping) + api_base: Final = params["api_base"] + assert isinstance(api_base, str) + return api_base + + +def _request_body(content: bytes) -> Mapping[str, object]: + body: Final = json.loads(content) + assert isinstance(body, dict) + return cast(Mapping[str, object], body) + + +def test_sync_resolution_round_robins_sorted_ips_without_mutating_deployment( + monkeypatch: pytest.MonkeyPatch, +) -> None: + original: Final = _deployment() + before: Final = copy.deepcopy(original) + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + assert type == socket.SOCK_STREAM + return _records("10.0.0.3", "10.0.0.1", "10.0.0.2", "10.0.0.1") + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=30, + proxy_environment=_empty_proxy_environment, + ) + + results: Final = tuple(discovery.resolve_deployment(original) for _ in range(3)) + + assert tuple(_api_base(result) for result in results) == ( + "http://10.0.0.1:8000/v1", + "http://10.0.0.2:8000/v1", + "http://10.0.0.3:8000/v1", + ) + assert original == before + assert results[0] is not original + assert results[0]["model_info"] is original["model_info"] + + +@pytest.mark.asyncio +async def test_async_resolution_round_robins_sorted_ips_and_preserves_url_path( + monkeypatch: pytest.MonkeyPatch, +) -> None: + original: Final = _deployment(api_base=f"{_SERVICE_URL}/custom?version=1") + before: Final = copy.deepcopy(original) + loop: Final = asyncio.get_running_loop() + lookups: Final = count() + + async def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return _records("10.0.0.3", "10.0.0.1", "10.0.0.2", port=int(port or 0)) + + monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=30, + proxy_environment=_empty_proxy_environment, + ) + + results: Final = ( + await discovery.async_resolve_deployment(original), + await discovery.async_resolve_deployment(original), + await discovery.async_resolve_deployment(original), + ) + + assert tuple(_api_base(result) for result in results) == ( + "http://10.0.0.1:8000/v1/custom?version=1", + "http://10.0.0.2:8000/v1/custom?version=1", + "http://10.0.0.3:8000/v1/custom?version=1", + ) + assert original == before + assert results[0] is not original + assert results[0]["model_info"] is original["model_info"] + + +def test_refresh_reuses_cached_set_until_interval_then_uses_new_pods( + monkeypatch: pytest.MonkeyPatch, +) -> None: + clock: Final = _clock((10.0, 10.0, 14.99, 15.0, 15.0)) + responses: Final = iter((_records("10.0.0.1", "10.0.0.2"), _records("10.0.0.3"))) + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return list(next(responses)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + first: Final = _api_base(discovery.resolve_deployment(_deployment())) + cached: Final = _api_base(discovery.resolve_deployment(_deployment())) + refreshed: Final = _api_base(discovery.resolve_deployment(_deployment())) + + assert (first, cached, refreshed) == ( + "http://10.0.0.1:8000/v1", + "http://10.0.0.2:8000/v1", + "http://10.0.0.3:8000/v1", + ) + assert next(lookups) == 2 + + +def test_refresh_carries_round_robin_cursor_into_new_ip_set(monkeypatch: pytest.MonkeyPatch) -> None: + clock: Final = _clock((0.0, 0.0, 5.0, 5.0)) + responses: Final = iter((_records("10.0.0.1", "10.0.0.2", "10.0.0.3"), _records("10.0.0.4", "10.0.0.5"))) + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + return list(next(responses)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + first: Final = _api_base(discovery.resolve_deployment(_deployment())) + refreshed: Final = _api_base(discovery.resolve_deployment(_deployment())) + + assert (first, refreshed) == ("http://10.0.0.1:8000/v1", "http://10.0.0.5:8000/v1") + + +def test_idle_cache_eviction_restarts_round_robin_cursor(monkeypatch: pytest.MonkeyPatch) -> None: + clock: Final = _clock((0.0, 0.0, 0.0, 11.0, 11.0)) + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return _records("10.0.0.1", "10.0.0.2", "10.0.0.3", port=int(port or 0)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + idle_eviction_seconds=10, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + first: Final = _api_base(discovery.resolve_deployment(_deployment())) + second: Final = _api_base(discovery.resolve_deployment(_deployment())) + after_eviction: Final = _api_base(discovery.resolve_deployment(_deployment())) + + assert (first, second, after_eviction) == ( + "http://10.0.0.1:8000/v1", + "http://10.0.0.2:8000/v1", + "http://10.0.0.1:8000/v1", + ) + assert next(lookups) == 2 + + +def test_active_pod_set_survives_idle_eviction_when_refresh_interval_is_longer( + monkeypatch: pytest.MonkeyPatch, +) -> None: + clock_values: Final = iter(chain((0.0, 0.0, 200.0, 400.0, 550.0), repeat(550.0))) + + def clock() -> float: + return next(clock_values) + + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return _records("10.0.0.1", "10.0.0.2", "10.0.0.3", port=int(port or 0)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=600, + idle_eviction_seconds=300, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + api_bases: Final = tuple(_api_base(discovery.resolve_deployment(_deployment())) for _ in range(4)) + + assert api_bases == ( + "http://10.0.0.1:8000/v1", + "http://10.0.0.2:8000/v1", + "http://10.0.0.3:8000/v1", + "http://10.0.0.1:8000/v1", + ) + assert next(lookups) == 1 + + +def test_slow_dns_refresh_does_not_move_last_use_backward(monkeypatch: pytest.MonkeyPatch) -> None: + clock_values: Final = iter(chain((0.0, 250.0), repeat(500.0))) + + def clock() -> float: + return next(clock_values) + + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return _records("10.0.0.1", "10.0.0.2", "10.0.0.3", port=int(port or 0)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=600, + idle_eviction_seconds=300, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + api_bases: Final = tuple(_api_base(discovery.resolve_deployment(_deployment())) for _ in range(2)) + + assert api_bases == ("http://10.0.0.1:8000/v1", "http://10.0.0.2:8000/v1") + assert next(lookups) == 1 + + +def test_proxy_bypassing_only_service_hostname_keeps_api_base_and_warns_once( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + _stub_sync_dns(monkeypatch, "10.244.1.3") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_proxy_environment(".svc.cluster.local")) + original: Final = _deployment() + + with caplog.at_level(logging.WARNING): + first: Final = discovery.resolve_deployment(original) + second: Final = discovery.resolve_deployment(original) + + assert first is original + assert second is original + assert _api_base(first) == _SERVICE_URL + assert sum("NO_PROXY exempts it but not its pod IPs" in record.getMessage() for record in caplog.records) == 1 + + +def test_proxy_with_pod_cidr_in_no_proxy_keeps_service_hostname(monkeypatch: pytest.MonkeyPatch) -> None: + _stub_sync_dns(monkeypatch, "10.244.1.3") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_proxy_environment(".svc.cluster.local,10.244.0.0/16")) + + result: Final = discovery.resolve_deployment(_deployment()) + + assert _api_base(result) == _SERVICE_URL + + +def test_proxy_with_exact_pod_ip_in_no_proxy_keeps_pod_ip_substitution(monkeypatch: pytest.MonkeyPatch) -> None: + _stub_sync_dns(monkeypatch, "10.244.1.3") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_proxy_environment(".svc.cluster.local,10.244.1.3")) + + result: Final = discovery.resolve_deployment(_deployment()) + + assert _api_base(result) == "http://10.244.1.3:8000/v1" + + +def test_proxy_with_empty_no_proxy_keeps_pod_ip_substitution(monkeypatch: pytest.MonkeyPatch) -> None: + _stub_sync_dns(monkeypatch, "10.244.1.3") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_proxy_environment("")) + + result: Final = discovery.resolve_deployment(_deployment()) + + assert _api_base(result) == "http://10.244.1.3:8000/v1" + + +def test_empty_proxy_environment_keeps_pod_ip_substitution(monkeypatch: pytest.MonkeyPatch) -> None: + _stub_sync_dns(monkeypatch, "10.244.1.3") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment) + + result: Final = discovery.resolve_deployment(_deployment()) + + assert _api_base(result) == "http://10.244.1.3:8000/v1" + + +@pytest.mark.parametrize("no_pods_errno", _NO_POD_ERRNOS) +def test_authoritative_no_pods_uses_service_url_and_transient_failure_keeps_cached_pods( + monkeypatch: pytest.MonkeyPatch, no_pods_errno: int +) -> None: + clock: Final = _clock((0.0, 0.0, 0.0, 0.0, 5.0, 5.0)) + missing: Final = _deployment() + no_pod_lookups: Final = count() + + def missing_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(no_pod_lookups) + raise socket.gaierror(no_pods_errno, "missing") + + monkeypatch.setattr(socket, "getaddrinfo", missing_getaddrinfo) + discovery_without_pods: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + assert discovery_without_pods.resolve_deployment(missing) is missing + assert discovery_without_pods.resolve_deployment(missing) is missing + assert _api_base(missing) == _SERVICE_URL + assert next(no_pod_lookups) == 1 + + responses: Final = iter((_records("10.0.0.4"), socket.gaierror(socket.EAI_AGAIN, "temporary"))) + + def transient_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + response: Final = next(responses) + if isinstance(response, socket.gaierror): + raise response + return list(response) + + monkeypatch.setattr(socket, "getaddrinfo", transient_getaddrinfo) + discovery_with_pods: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + first: Final = _api_base(discovery_with_pods.resolve_deployment(_deployment())) + after_failure: Final = _api_base(discovery_with_pods.resolve_deployment(_deployment())) + + assert (first, after_failure) == ("http://10.0.0.4:8000/v1", "http://10.0.0.4:8000/v1") + + +@pytest.mark.parametrize( + ("api_base", "flag"), + ( + pytest.param(_SERVICE_URL, None, id="flag-absent"), + pytest.param(_SERVICE_URL, False, id="flag-false"), + pytest.param("https://vllm-headless.ns.svc.cluster.local:8000/v1", True, id="https"), + pytest.param("http://10.0.0.9:8000/v1", True, id="ip-literal"), + pytest.param("http://[zzzz]", True, id="malformed-url"), + ), +) +def test_unsupported_deployments_are_returned_unchanged_without_dns( + monkeypatch: pytest.MonkeyPatch, api_base: str, flag: bool | None +) -> None: + deployment: Final = _deployment(api_base=api_base, kubernetes_pod_discovery=flag) + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(lookups) + return [] + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment) + + assert discovery.resolve_deployment(deployment) is deployment + assert next(lookups) == 0 + + +@pytest.mark.asyncio +async def test_non_mapping_deployments_are_returned_unchanged() -> None: + deployment: Final = object() + discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment) + + assert discovery.resolve_deployment(deployment) is deployment + assert await discovery.async_resolve_deployment(deployment) is deployment + + +def test_https_discovery_warning_is_emitted_once_per_host(caplog: pytest.LogCaptureFixture) -> None: + caplog.set_level(logging.WARNING) + deployment: Final = _deployment(api_base="https://vllm-headless.ns.svc.cluster.local:8000/v1") + discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment) + + results: Final = ( + discovery.resolve_deployment(deployment), + discovery.resolve_deployment(deployment), + ) + warnings: Final = tuple( + record + for record in caplog.records + if "Kubernetes pod discovery is disabled for HTTPS host" in record.getMessage() + ) + + assert results == (deployment, deployment) + assert all(result is deployment for result in results) + assert len(warnings) == 1 + + +def test_ipv6_pod_address_is_bracketed_in_url(monkeypatch: pytest.MonkeyPatch) -> None: + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + return _records("2001:db8::2") + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment) + + result: Final = discovery.resolve_deployment(_deployment()) + + assert _api_base(result) == "http://[2001:db8::2]:8000/v1" + + +@pytest.mark.asyncio +async def test_concurrent_async_refresh_uses_cached_ips_while_refresh_is_in_flight( + monkeypatch: pytest.MonkeyPatch, +) -> None: + clock: Final = _clock((0.0, 0.0, 5.0, 5.0, 5.0, 5.0, 5.0)) + loop: Final = asyncio.get_running_loop() + refresh_started: Final = asyncio.Event() + finish_refresh: Final = asyncio.Event() + lookups: Final = count() + + async def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + lookup_number: Final = next(lookups) + if lookup_number == 0: + return _records("10.0.0.1", port=int(port or 0)) + if lookup_number > 1: + return _records("10.0.0.3", port=int(port or 0)) + refresh_started.set() + await finish_refresh.wait() + return _records("10.0.0.2", port=int(port or 0)) + + monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + first: Final = _api_base(await discovery.async_resolve_deployment(_deployment())) + refresh_task: Final = asyncio.create_task(discovery.async_resolve_deployment(_deployment())) + await refresh_started.wait() + concurrent: Final = _api_base(await discovery.async_resolve_deployment(_deployment())) + finish_refresh.set() + refreshed: Final = _api_base(await refresh_task) + + assert (first, concurrent, refreshed) == ( + "http://10.0.0.1:8000/v1", + "http://10.0.0.1:8000/v1", + "http://10.0.0.2:8000/v1", + ) + assert next(lookups) == 2 + + +@pytest.mark.parametrize("cache_existing", (False, True), ids=("cold-cache", "stale-cache")) +def test_concurrent_sync_refresh_uses_cached_ips_or_service_url_while_in_flight( + monkeypatch: pytest.MonkeyPatch, + cache_existing: bool, +) -> None: + ticks: Final = (0.0, 0.0, 5.0, 5.0, 5.0) if cache_existing else (0.0, 0.0, 0.0) + clock: Final = _clock(ticks) + refresh_started: Final = threading.Event() + finish_refresh: Final = threading.Event() + lookups: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + lookup_number: Final = next(lookups) + if cache_existing and lookup_number == 0: + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + refresh_started.set() + assert finish_refresh.wait(timeout=3) + return _records("10.0.0.3", port=int(port or 0)) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + discovery: Final = KubernetesPodDiscovery( + refresh_interval_seconds=5, + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + if cache_existing: + assert _api_base(discovery.resolve_deployment(_deployment())) == "http://10.0.0.1:8000/v1" + + refresh_result: Final = Future[dict[str, object]]() + + def refresh() -> None: + refresh_result.set_result(discovery.resolve_deployment(_deployment())) + + thread: Final = threading.Thread(target=refresh, daemon=True) + thread.start() + try: + assert refresh_started.wait(timeout=2) + concurrent: Final = discovery.resolve_deployment(_deployment()) + assert _api_base(concurrent) == ("http://10.0.0.2:8000/v1" if cache_existing else _SERVICE_URL) + assert next(lookups) == (2 if cache_existing else 1) + finally: + finish_refresh.set() + thread.join(timeout=2) + assert not thread.is_alive() + assert _api_base(refresh_result.result(timeout=2)) == "http://10.0.0.3:8000/v1" + + +@pytest.mark.asyncio +async def test_router_sends_pod_hosts_without_forwarding_discovery_flag( + monkeypatch: pytest.MonkeyPatch, +) -> None: + clock: Final = _clock((0.0, 0.0, 0.0, 0.0, 0.0, 5.0, 5.0)) + loop: Final = asyncio.get_running_loop() + async_lookups: Final = count() + sync_lookups: Final = count() + + async def async_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(async_lookups) + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + def sync_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(sync_lookups) + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo) + monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + 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.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( + key="registered-id_async_client", + value=cached_async_client, + local_only=True, + ) + discovery_enabled_router.cache.set_cache( + key="registered-id_client", + value=cached_sync_client, + local_only=True, + ) + for _ in range(4): + await discovery_enabled_router.acompletion( + model="gpu-model", messages=[{"role": "user", "content": "hello"}] + ) + discovery_enabled_router.completion(model="gpu-model", messages=[{"role": "user", "content": "hello"}]) + await control_router.acompletion(model="gpu-model", messages=[{"role": "user", "content": "hello"}]) + + hosts: Final = tuple(call.request.url.host for call in route.calls) + bodies: Final = tuple(_request_body(call.request.content) for call in route.calls) + await cached_async_client.close() + cached_sync_client.close() + + assert hosts == ( + "10.0.0.1", + "10.0.0.2", + "10.0.0.1", + "10.0.0.2", + "10.0.0.1", + "vllm-headless.ns.svc.cluster.local", + ) + assert all("kubernetes_pod_discovery" not in body for body in bodies) + assert next(async_lookups) == 1 + assert next(sync_lookups) == 1 + + +@pytest.mark.asyncio +async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_requests( + monkeypatch: pytest.MonkeyPatch, +) -> None: + clock: Final = _clock((0.0,) * 40) + loop: Final = asyncio.get_running_loop() + async_lookups: Final = count() + sync_lookups: Final = count() + + async def async_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(async_lookups) + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + def sync_getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + next(sync_lookups) + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo) + monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + 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.kubernetes_pod_discovery = KubernetesPodDiscovery( + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + sync_router.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): + await async_router.acompletion(model="gpu-model", messages=[{"role": "user", "content": "hello"}]) + for _ in range(4): + sync_router.completion(model="gpu-model", messages=[{"role": "user", "content": "hello"}]) + hosts: Final = tuple(call.request.url.host for call in route.calls) + + assert hosts == ("10.0.0.1", "10.0.0.2") * 4 + assert next(async_lookups) == 1 + assert next(sync_lookups) == 1 + + +@pytest.mark.asyncio +async def test_unresolved_router_selectors_keep_service_hostname() -> None: + router: Final = Router(model_list=[_deployment()]) + + sync_selected: Final = router._get_available_deployment_unresolved(model="gpu-model", request_kwargs={}) + async_selected: Final = await router._async_get_available_deployment_unresolved( + model="gpu-model", + request_kwargs={}, + ) + + assert _api_base(sync_selected) == _SERVICE_URL + assert _api_base(async_selected) == _SERVICE_URL + + +@pytest.mark.asyncio +async def test_router_async_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPatch) -> None: + clock: Final = _clock((0.0,) * 8) + loop: Final = asyncio.get_running_loop() + attempts: Final = count() + + async def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + def response_for(request: httpx.Request) -> httpx.Response: + attempt: Final = next(attempts) + status_code: Final = 500 if attempt == 0 else 200 + body: Final = {"error": {"message": "retryable", "type": "server_error"}} if attempt == 0 else _CHAT_RESPONSE + return httpx.Response(status_code, json=body, request=request) + + monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + 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.kubernetes_pod_discovery = KubernetesPodDiscovery( + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + response: Final = await router.acompletion( + model="gpu-model", + messages=[{"role": "user", "content": "hello"}], + ) + hosts: Final = tuple(call.request.url.host for call in route.calls) + + assert response.choices[0].message.content == "ok" + assert hosts == ("10.0.0.1", "10.0.0.2") + + +def test_router_sync_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPatch) -> None: + clock: Final = _clock((0.0,) * 8) + attempts: Final = count() + + def getaddrinfo( + host: str | None, + port: str | int | None, + family: int = 0, + type: int = 0, + proto: int = 0, + flags: int = 0, + ) -> list[_AddrInfo]: + return _records("10.0.0.1", "10.0.0.2", port=int(port or 0)) + + def response_for(request: httpx.Request) -> httpx.Response: + attempt: Final = next(attempts) + status_code: Final = 500 if attempt == 0 else 200 + body: Final = {"error": {"message": "retryable", "type": "server_error"}} if attempt == 0 else _CHAT_RESPONSE + return httpx.Response(status_code, json=body, request=request) + + monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + 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.kubernetes_pod_discovery = KubernetesPodDiscovery( + clock=clock, + proxy_environment=_empty_proxy_environment, + ) + + response: Final = router.completion( + model="gpu-model", + messages=[{"role": "user", "content": "hello"}], + ) + hosts: Final = tuple(call.request.url.host for call in route.calls) + + assert response.choices[0].message.content == "ok" + assert hosts == ("10.0.0.1", "10.0.0.2") diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index 8d163731a51..29026e975ec 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -124,6 +124,7 @@ OPTION_NAMES: Final = ( "use_litellm_proxy", "use_chat_completions_api", "use_in_pass_through", + "kubernetes_pod_discovery", "allowed_openai_params", "fallbacks", "context_window_fallback_dict", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx index efba26734ff..ac1154e2736 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx @@ -92,6 +92,7 @@ const alwaysMounted = { const advancedOpenExtras = { guardrails: undefined, + kubernetes_pod_discovery: undefined, tags: undefined, use_in_pass_through: undefined, vector_store_ids: undefined, @@ -181,6 +182,24 @@ describe("AddModelPanel submit payload contract", () => { }); }); + it("submits Kubernetes pod discovery in litellm_params", async () => { + const { user, openAdvanced, fillRequired, submit } = await setup(); + await fillRequired(); + await openAdvanced(); + await user.click(screen.getByRole("switch", { name: "Kubernetes pod discovery" })); + await submit(); + + expect(lastCreatedModel()).toStrictEqual({ + model_name: "gpt-4o", + litellm_params: { + ...alwaysMounted, + ...advancedOpenExtras, + kubernetes_pod_discovery: true, + }, + model_info: { ...baseModelInfo }, + }); + }); + it("drops a collapsed section's keys and the value typed into it", async () => { const { user, openAdvanced, closeAdvanced, fillRequired, submit } = await setup(); await fillRequired(); diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx index ceb9d1ca713..ca92701c913 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.test.tsx @@ -1,6 +1,8 @@ import { act, fireEvent, render, waitFor, screen } from "@testing-library/react"; +import { useFormContext } from "react-hook-form"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { MountedFormHost } from "../../../tests/mounted-form-host"; +import type { MountedFormValues } from "../common_components/MountedFormField"; import AdvancedSettings from "./advanced_settings"; const mockUsePtuCostAttributionEnabled = vi.fn(); @@ -11,7 +13,16 @@ vi.mock("@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled", () => const PTU_LABELS = ["PTU Count", "Calculated Cost per PTU / Hour (USD)", "PTU Effective From (UTC)"]; -const renderAdvancedSettings = () => +const KubernetesPodDiscoveryValue = () => { + const form = useFormContext(); + return ( + + {String(form.watch("kubernetes_pod_discovery") === true)} + + ); +}; + +const renderAdvancedSettings = (includeDiscoveryValue = false) => render( tagsList={{}} accessToken="test-token" /> + {includeDiscoveryValue && } , ); @@ -34,6 +46,23 @@ describe("AdvancedSettings", () => { renderAdvancedSettings(); }); + it("updates the Kubernetes pod discovery form value when enabled", async () => { + renderAdvancedSettings(true); + act(() => { + fireEvent.click(screen.getByText("Advanced Settings")); + }); + + const toggle = await screen.findByRole("switch", { name: "Kubernetes pod discovery" }); + const formValue = screen.getByRole("status", { name: "Kubernetes pod discovery form value" }); + expect(toggle).not.toBeChecked(); + expect(formValue).toHaveTextContent("false"); + + fireEvent.click(toggle); + + expect(toggle).toBeChecked(); + expect(formValue).toHaveTextContent("true"); + }); + it("should render tags list", async () => { renderAdvancedSettings(); fireEvent.click(screen.getByText("Advanced Settings")); diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx index eb12e49ca89..5b4a6cfbcd0 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx @@ -437,6 +437,19 @@ const AdvancedSettings: React.FC = ({ )} + + {(control) => ( + + )} + +