feat(router): discover kubernetes pods behind a headless service api_base

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-10-03 08:52:51 +00:00
parent 09313bcf1b
commit b9e4f0613e
11 changed files with 1385 additions and 8 deletions

View file

@ -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))

View file

@ -225,6 +225,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,
@ -944,6 +951,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"
@ -12294,6 +12302,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)
@ -12303,6 +12318,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)
@ -13261,7 +13283,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,
@ -13281,7 +13303,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,
@ -13412,6 +13434,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,
@ -14132,7 +14160,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,
@ -14293,6 +14321,12 @@ class Router:
)
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,
@ -14708,15 +14742,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:

View file

@ -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

View file

@ -144,6 +144,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

View file

@ -586,6 +586,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

View file

@ -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())

View file

@ -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")

View file

@ -103,6 +103,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",

View file

@ -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();

View file

@ -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<MountedFormValues>();
return (
<output aria-label="Kubernetes pod discovery form value" role="status">
{String(form.watch("kubernetes_pod_discovery") === true)}
</output>
);
};
const renderAdvancedSettings = (includeDiscoveryValue = false) =>
render(
<MountedFormHost>
<AdvancedSettings
@ -21,6 +32,7 @@ const renderAdvancedSettings = () =>
tagsList={{}}
accessToken="test-token"
/>
{includeDiscoveryValue && <KubernetesPodDiscoveryValue />}
</MountedFormHost>,
);
@ -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"));

View file

@ -437,6 +437,19 @@ const AdvancedSettings: React.FC<AdvancedSettingsProps> = ({
)}
</MountedFormField>
<MountedFormField
name="kubernetes_pod_discovery"
label={labelWithHint(
"Kubernetes pod discovery",
"Treat the API Base hostname as a Kubernetes headless service: resolve it to its ready pod IPs and spread requests across the pods. LiteLLM must run inside the same cluster.",
)}
className="mb-4"
>
{(control) => (
<Switch id={control.id} checked={control.value === true} onCheckedChange={control.onChange} />
)}
</MountedFormField>
<MountedFormField
name="cache_control"
label={labelWithHint(CACHE_CONTROL_LABEL, CACHE_CONTROL_TOOLTIP)}