feat(router): keep a session on one kubernetes pod

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-10-03 19:52:10 +00:00
parent b9e4f0613e
commit 277f4b4218
17 changed files with 1257 additions and 48 deletions

View file

@ -12,6 +12,7 @@ KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS: Final = float(
KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS: Final = float(
os.getenv("KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS", "300")
)
KUBERNETES_POD_ROUTING_KEY: Final = "kubernetes_pod_routing"
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

@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
import litellm
from litellm._logging import verbose_logger
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
from litellm.integrations._types.open_inference import (
OpenInferenceSpanKindValues,
SpanAttributes,
@ -2403,9 +2404,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
#############################################
############ LLM CALL METADATA ##############
#############################################
metadata: Final = standard_logging_payload["metadata"]
metadata: Final[Mapping[str, object]] = standard_logging_payload["metadata"]
for key, value in metadata.items():
self.safe_set_attribute(span=span, key=f"metadata.{key}", value=value)
self.safe_set_attribute(
span=span,
key=f"metadata.{key}",
value=safe_dumps(value) if key == KUBERNETES_POD_ROUTING_KEY else value,
)
# get hidden params
hidden_params: Final = getattr(standard_logging_payload, "hidden_params", None) or (

View file

@ -11,12 +11,13 @@ import httpx
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
from litellm.types.llms.openai import AllMessageValues, OpenAIChatCompletionFinishReason
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.types.utils import ModelResponseStream
from litellm.types.utils import ModelResponseStream, StandardLoggingKubernetesPodRouting
Span = _Span | Any
else:
@ -346,6 +347,52 @@ def proxy_stamped_used_client_oauth_token(metadata: object, litellm_params: Mapp
return metadata.get("used_client_oauth_token") if isinstance(metadata, Mapping) else None
def proxy_stamped_kubernetes_pod_routing(
metadata: object,
litellm_params: Mapping[str, object] | None,
) -> "StandardLoggingKubernetesPodRouting | None":
litellm_metadata: Final = litellm_params.get("litellm_metadata") if litellm_params is not None else None
litellm_routing: Final = _standard_logging_kubernetes_pod_routing(litellm_metadata)
if litellm_routing is not None:
return litellm_routing
return _standard_logging_kubernetes_pod_routing(metadata)
def _standard_logging_kubernetes_pod_routing(
metadata: object,
) -> "StandardLoggingKubernetesPodRouting | None":
if not isinstance(metadata, Mapping):
return None
metadata_mapping: Final[Mapping[str, object]] = metadata
routing: Final = metadata_mapping.get(KUBERNETES_POD_ROUTING_KEY)
if not isinstance(routing, Mapping):
return None
routing_mapping: Final[Mapping[str, object]] = routing
service_host: Final = routing_mapping.get("service_host")
pod_ip: Final = routing_mapping.get("pod_ip")
pod_count: Final = routing_mapping.get("pod_count")
selection: Final = routing_mapping.get("selection")
if (
not isinstance(service_host, str)
or not isinstance(pod_ip, str)
or not isinstance(pod_count, int)
or isinstance(pod_count, bool)
or pod_count < 1
):
return None
match selection:
case "round_robin" | "session_affinity" | "session_affinity_retry":
routing_record: Final[StandardLoggingKubernetesPodRouting] = {
"service_host": service_host,
"pod_ip": pod_ip,
"pod_count": pod_count,
"selection": selection,
}
return routing_record
case _:
return None
def get_litellm_metadata_from_kwargs(kwargs: dict):
"""
Helper to get litellm metadata from all litellm request kwargs

View file

@ -40,6 +40,7 @@ from litellm.constants import (
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
EMPTY_MAPPING,
KUBERNETES_POD_ROUTING_KEY,
PROVIDER_REQUEST_ID_HEADERS,
REDACTED_BY_LITELLM,
)
@ -72,6 +73,7 @@ from litellm.litellm_core_utils.classifier_logging import (
from litellm.litellm_core_utils.core_helpers import (
get_provider_response_headers_from_hidden_params,
is_expected_client_error,
proxy_stamped_kubernetes_pod_routing,
proxy_stamped_used_client_oauth_token,
reconstruct_model_name,
set_response_cost_in_hidden_params,
@ -287,7 +289,9 @@ else:
_PAGERDUTY_ALERTING_FACTORY: Final = PagerDutyAlerting
_in_memory_loggers: Final[list[CustomLogger]] = []
_STANDARD_LOGGING_METADATA_RESOLVED_KEYS: Final[frozenset[str]] = frozenset(("used_client_oauth_token",))
_STANDARD_LOGGING_METADATA_RESOLVED_KEYS: Final[frozenset[str]] = frozenset(
("used_client_oauth_token", KUBERNETES_POD_ROUTING_KEY)
)
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = (
frozenset(StandardLoggingMetadata.__annotations__.keys()) - _STANDARD_LOGGING_METADATA_RESOLVED_KEYS
)
@ -5742,7 +5746,7 @@ class StandardLoggingPayloadSetup:
@staticmethod
def get_standard_logging_metadata(
metadata: Mapping[str, object] | None,
litellm_params: dict | None = None,
litellm_params: dict[str, object] | None = None,
prompt_integration: str | None = None,
applied_guardrails: list[str] | None = None,
mcp_tool_call_metadata: StandardLoggingMCPToolCall | None = None,
@ -5782,6 +5786,8 @@ class StandardLoggingPayloadSetup:
prompt_integration=prompt_integration,
)
kubernetes_pod_routing: Final = proxy_stamped_kubernetes_pod_routing(metadata, litellm_params)
# Initialize with default values
clean_metadata = StandardLoggingMetadata(
user_api_key_hash=None,
@ -5818,6 +5824,7 @@ class StandardLoggingPayloadSetup:
user_api_key_auth_metadata=None,
team_alias=None,
team_id=None,
kubernetes_pod_routing=kubernetes_pod_routing,
used_client_oauth_token=resolve_used_client_oauth_token(
proxy_stamped_used_client_oauth_token(metadata, litellm_params),
custom_llm_provider,
@ -6834,6 +6841,7 @@ def get_standard_logging_metadata(
user_api_key_auth_metadata=None,
team_alias=None,
team_id=None,
kubernetes_pod_routing=None,
used_client_oauth_token=None,
)
if isinstance(metadata, dict):
@ -6907,6 +6915,7 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
requester_ip_address="127.0.0.1",
requester_metadata=None,
user_api_key_end_user_id="test_end_user",
kubernetes_pod_routing=None,
)
hidden_params: Final = StandardLoggingHiddenParams(

View file

@ -75,6 +75,7 @@ from litellm.types.utils import (
ProviderField,
StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingKubernetesPodRouting,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
@ -4213,6 +4214,7 @@ class SpendLogsMetadata(TypedDict):
mcp_tool_call_metadata: StandardLoggingMCPToolCall | None
vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None
routing_decision: StandardLoggingRoutingDecision | None
kubernetes_pod_routing: ReadOnly[StandardLoggingKubernetesPodRouting | None]
internal_call_origin: InternalCallOrigin | None
litellm_roi_estimator: ReadOnly[NotRequired[bool | None]]
guardrail_information: list[StandardLoggingGuardrailInformation] | None

View file

@ -16,6 +16,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
EMPTY_MAPPING,
KUBERNETES_POD_ROUTING_KEY,
LITELLM_PROXY_MASTER_KEY_ALIAS,
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
@ -33,6 +34,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, without_classifier_audit
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
proxy_stamped_kubernetes_pod_routing,
proxy_stamped_used_client_oauth_token,
reconstruct_model_name,
)
@ -59,6 +61,7 @@ from litellm.types.utils import (
CostBreakdown,
LlmProviders,
StandardLoggingGuardrailInformation,
StandardLoggingKubernetesPodRouting,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayload,
@ -158,6 +161,7 @@ _STAMPED_METADATA_KEYS: Final = frozenset(
"autorouter_savings_estimate",
"autorouter_baseline_observation",
"used_client_oauth_token",
KUBERNETES_POD_ROUTING_KEY,
"litellm_roi_estimator",
)
)
@ -184,6 +188,7 @@ def _get_spend_logs_metadata(
router_metadata: SpendLogsRouterMetadata | None = None,
azure_spillover: AzureSpillover | None = None,
used_client_oauth_token: bool | None = None,
kubernetes_pod_routing: StandardLoggingKubernetesPodRouting | None = None,
) -> SpendLogsMetadata:
if metadata is None:
return SpendLogsMetadata(
@ -230,6 +235,7 @@ def _get_spend_logs_metadata(
router_metadata=router_metadata,
azure_spillover=azure_spillover,
used_client_oauth_token=used_client_oauth_token,
kubernetes_pod_routing=None,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys()))
@ -246,6 +252,7 @@ def _get_spend_logs_metadata(
router_metadata=router_metadata,
azure_spillover=azure_spillover,
used_client_oauth_token=used_client_oauth_token,
kubernetes_pod_routing=kubernetes_pod_routing,
litellm_roi_estimator=metadata.get("litellm_roi_estimator") is True,
)
_raw_key: Final = clean_metadata.get("user_api_key")
@ -516,6 +523,16 @@ def get_logging_payload(
# standardize this function to be used across, s3, dynamoDB, langfuse logging
litellm_params: Final = kwargs.get("litellm_params", {})
metadata: Final = get_litellm_metadata_from_kwargs(kwargs)
routing_litellm_params: Final[Mapping[str, object] | None] = (
litellm_params if isinstance(litellm_params, Mapping) else None
)
routing_metadata: Final[object] = (
routing_litellm_params.get("metadata") if routing_litellm_params is not None else None
)
kubernetes_pod_routing: Final = proxy_stamped_kubernetes_pod_routing(
routing_metadata,
routing_litellm_params,
)
completion_start_time: Final = kwargs.get("completion_start_time", end_time)
call_type: Final = kwargs.get("call_type")
cache_hit: Final = kwargs.get("cache_hit", False)
@ -727,6 +744,7 @@ def get_logging_payload(
used_client_oauth_token=resolve_used_client_oauth_token(
proxy_stamped_used_client_oauth_token(litellm_params.get("metadata"), litellm_params), custom_llm_provider
),
kubernetes_pod_routing=kubernetes_pod_routing,
azure_spillover=azure_spillover(
response_headers=kwargs.get("response_headers")
if isinstance(kwargs.get("response_headers"), Mapping)

View file

@ -63,6 +63,7 @@ from litellm.constants import (
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
DEFAULT_MAX_LRU_CACHE_SIZE,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
KUBERNETES_POD_ROUTING_KEY,
OUTPUT_TOKEN_CEILING_PARAMS,
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY,
ROUTING_REQUEST_TAGS_METADATA_KEY,
@ -4027,6 +4028,7 @@ class Router:
"model_info": model_info,
"api_base": deployment_api_base,
"deployment_model_name": deployment_model_name,
KUBERNETES_POD_ROUTING_KEY: deployment.get(KUBERNETES_POD_ROUTING_KEY),
}
)

View file

@ -1,4 +1,6 @@
import asyncio
import hashlib
import inspect
import ipaddress
import socket
import threading
@ -8,7 +10,7 @@ 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
from typing import Concatenate, Final, Literal, ParamSpec, Protocol, TypeAlias, TypeVar, cast
import httpx
@ -16,14 +18,27 @@ from litellm._logging import verbose_router_logger
from litellm.constants import (
KUBERNETES_POD_DISCOVERY_IDLE_EVICTION_SECONDS,
KUBERNETES_POD_DISCOVERY_REFRESH_INTERVAL_SECONDS,
KUBERNETES_POD_ROUTING_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
)
_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]
_PodSelection: TypeAlias = tuple[str, int, Literal["round_robin", "session_affinity", "session_affinity_retry"]]
_DeploymentT = TypeVar("_DeploymentT")
def _retry_count_from_metadata(metadata: object) -> int | None:
if not isinstance(metadata, Mapping):
return None
metadata_mapping: Final[Mapping[str, object]] = metadata
retry_count: Final = metadata_mapping.get("request_retry_count")
if type(retry_count) is int and retry_count >= 0:
return retry_count
return None
@dataclass(frozen=True, slots=True)
class _PodSet:
ips: tuple[str, ...]
@ -68,49 +83,101 @@ class KubernetesPodDiscovery:
self._refreshing = self._refreshing | {key}
return None, True
def resolve_deployment(self, deployment: _DeploymentT) -> _DeploymentT:
def resolve_deployment(
self,
deployment: _DeploymentT,
request_kwargs: Mapping[str, object] | None = None,
) -> _DeploymentT:
eligible: Final = self._eligible(deployment)
if eligible is None:
return deployment
url, key, deployment_mapping = eligible
host, port = key
now: Final = self.clock()
session_id: Final = self._session_id(request_kwargs)
retry_count: Final = self._request_retry_count(request_kwargs)
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)
return self._deployment_with_cached_ip(
deployment, deployment_mapping, key, url, now, session_id, retry_count
)
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)
return self._apply_lookup(deployment, deployment_mapping, key, url, error, now, session_id, retry_count)
else:
return self._apply_lookup(deployment, deployment_mapping, key, url, records, now)
return self._apply_lookup(deployment, deployment_mapping, key, url, records, now, session_id, retry_count)
finally:
with self._lock:
self._refreshing = self._refreshing - {key}
async def async_resolve_deployment(self, deployment: _DeploymentT) -> _DeploymentT:
async def async_resolve_deployment(
self,
deployment: _DeploymentT,
request_kwargs: Mapping[str, object] | None = None,
) -> _DeploymentT:
eligible: Final = self._eligible(deployment)
if eligible is None:
return deployment
url, key, deployment_mapping = eligible
host, port = key
now: Final = self.clock()
session_id: Final = self._session_id(request_kwargs)
retry_count: Final = self._request_retry_count(request_kwargs)
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)
return self._deployment_with_cached_ip(
deployment, deployment_mapping, key, url, now, session_id, retry_count
)
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)
return self._apply_lookup(deployment, deployment_mapping, key, url, result, now, session_id, retry_count)
finally:
with self._lock:
self._refreshing = self._refreshing - {key}
@staticmethod
def _session_id(request_kwargs: Mapping[str, object] | None) -> str | None:
if request_kwargs is None:
return None
metadata_dicts: Final = tuple(
cast(Mapping[str, object], metadata) # cast-ok: runtime dict check precedes metadata lookup
for metadata in (request_kwargs.get("litellm_metadata"), request_kwargs.get("metadata"))
if isinstance(metadata, dict)
)
if any(metadata.get(SESSION_ID_GENERATED_METADATA_KEY) for metadata in metadata_dicts):
return None
metadata_session_id: Final = next(
(
session_id
for metadata in metadata_dicts
if (session_id := metadata.get("session_id")) is not None and session_id != ""
),
None,
)
if metadata_session_id is not None:
return str(metadata_session_id)
top_level_session_id: Final = request_kwargs.get("litellm_session_id")
if top_level_session_id is None or top_level_session_id == "":
return None
return str(top_level_session_id)
@staticmethod
def _request_retry_count(request_kwargs: Mapping[str, object] | None) -> int:
if request_kwargs is None:
return 0
retry_counts: Final = tuple(
_retry_count_from_metadata(request_kwargs.get(metadata_key))
for metadata_key in ("litellm_metadata", "metadata")
)
return next((retry_count for retry_count in retry_counts if retry_count is not None), 0)
def _eligible(self, deployment: object) -> tuple[httpx.URL, _CacheKey, Mapping[str, object]] | None:
if not isinstance(deployment, Mapping):
return None
@ -164,6 +231,8 @@ class KubernetesPodDiscovery:
url: httpx.URL,
result: Sequence[_AddressInfo] | OSError,
now: float,
session_id: str | None,
retry_count: int,
) -> _DeploymentT:
if isinstance(result, OSError):
if isinstance(result, socket.gaierror) and self._is_authoritative_no_pods(result):
@ -171,12 +240,14 @@ class KubernetesPodDiscovery:
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)
return self._deployment_with_cached_ip(
deployment, deployment_mapping, key, url, now, session_id, retry_count
)
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)
return self._deployment_with_cached_ip(deployment, deployment_mapping, key, url, now, session_id, retry_count)
@staticmethod
def _is_ip_literal(host: str) -> bool:
@ -218,20 +289,50 @@ class KubernetesPodDiscovery:
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:
def _next_ip(
self,
key: _CacheKey,
now: float,
session_id: str | None,
retry_count: int,
) -> _PodSelection | 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)
pod_count: Final = len(cached.ips)
ranked_ips: Final = (
tuple(
sorted(
cached.ips,
key=lambda address: (
-int.from_bytes(
hashlib.blake2b(f"{session_id}\x00{address}".encode(), digest_size=8).digest(),
"big",
),
address,
),
)
)
if session_id is not None
else ()
)
ip: Final = cached.ips[cached.cursor] if session_id is None else ranked_ips[retry_count % pod_count]
next_cursor: Final = (cached.cursor + 1) % len(cached.ips) if session_id is None else cached.cursor
selection: Final = (
"round_robin"
if session_id is None
else "session_affinity_retry"
if retry_count > 0
else "session_affinity"
)
self._cache = MappingProxyType(
{
**self._cache,
key: replace(cached, cursor=next_cursor, last_used_at=max(cached.last_used_at, now)),
}
)
return ip
return ip, pod_count, selection
def _deployment_with_cached_ip(
self,
@ -240,10 +341,13 @@ class KubernetesPodDiscovery:
key: _CacheKey,
url: httpx.URL,
now: float,
session_id: str | None,
retry_count: int,
) -> _DeploymentT:
ip: Final = self._next_ip(key, now)
if ip is None:
selection: Final = self._next_ip(key, now, session_id, retry_count)
if selection is None:
return deployment
ip, pod_count, selection_kind = selection
if self._proxy_bypasses_only_service_host(key[0], ip):
self._warn_proxy_once(key[0])
return deployment
@ -258,6 +362,12 @@ class KubernetesPodDiscovery:
{
**deployment_mapping,
"litellm_params": {**typed_params, "api_base": str(url.copy_with(host=ip))},
KUBERNETES_POD_ROUTING_KEY: {
"service_host": key[0],
"pod_ip": ip,
"pod_count": pod_count,
"selection": selection_kind,
},
},
)
@ -304,13 +414,40 @@ _R = TypeVar("_R")
_S = TypeVar("_S", bound=_HasPodDiscovery)
def _request_kwargs_from_call(
signature: inspect.Signature, *args: object, **kwargs: object
) -> Mapping[str, object] | None:
try:
request_kwargs: Final = signature.bind_partial(*args, **kwargs).arguments.get(
"request_kwargs",
kwargs.get("request_kwargs"),
)
except TypeError:
return _mapping_request_kwargs(kwargs.get("request_kwargs"))
return _mapping_request_kwargs(request_kwargs)
def _mapping_request_kwargs(request_kwargs: object) -> Mapping[str, object] | None:
if not isinstance(request_kwargs, Mapping):
return None
return cast( # cast-ok: router selector request kwargs use string keys
Mapping[str, object],
request_kwargs,
)
def resolve_pods_after(
fn: Callable[Concatenate[_S, _P], _R],
) -> Callable[Concatenate[_S, _P], _R]:
signature: Final = inspect.signature(fn)
@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 self.kubernetes_pod_discovery.resolve_deployment(
deployment,
_request_kwargs_from_call(signature, self, *args, **kwargs),
)
return wrapped
@ -318,10 +455,15 @@ def resolve_pods_after(
def async_resolve_pods_after(
fn: Callable[Concatenate[_S, _P], Awaitable[_R]],
) -> Callable[Concatenate[_S, _P], Coroutine[object, object, _R]]:
signature: Final = inspect.signature(fn)
@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 await self.kubernetes_pod_discovery.async_resolve_deployment(
deployment,
_request_kwargs_from_call(signature, self, *args, **kwargs),
)
return wrapped
@ -330,10 +472,15 @@ def resolve_pods_after_bound(
fn: Callable[_P, _R],
discovery: KubernetesPodDiscovery,
) -> Callable[_P, _R]:
signature: Final = inspect.signature(fn)
@wraps(fn)
def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R:
deployment: Final = fn(*args, **kwargs)
return discovery.resolve_deployment(deployment)
return discovery.resolve_deployment(
deployment,
_request_kwargs_from_call(signature, *args, **kwargs),
)
return wrapped
@ -342,9 +489,14 @@ def async_resolve_pods_after_bound(
fn: Callable[_P, Awaitable[_R]],
discovery: KubernetesPodDiscovery,
) -> Callable[_P, Coroutine[object, object, _R]]:
signature: Final = inspect.signature(fn)
@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 await discovery.async_resolve_deployment(
deployment,
_request_kwargs_from_call(signature, *args, **kwargs),
)
return wrapped

View file

@ -3119,6 +3119,13 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
tier_litellm_params: Mapping[str, object] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
class StandardLoggingKubernetesPodRouting(TypedDict):
service_host: ReadOnly[str]
pod_ip: ReadOnly[str]
pod_count: ReadOnly[int]
selection: ReadOnly[Literal["round_robin", "session_affinity", "session_affinity_retry"]]
# Fields whose values quote the caller's prompt. Dropped when an operator turns message
# logging off. Every other field aggregates the prompt without reproducing it and is kept,
# so a redacted row stays explainable. `test_every_routing_decision_field_is_classified`
@ -3180,6 +3187,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata):
mcp_tool_call_metadata: StandardLoggingMCPToolCall | None
vector_store_request_metadata: list[StandardLoggingVectorStoreRequest] | None
routing_decision: StandardLoggingRoutingDecision | None
kubernetes_pod_routing: ReadOnly[StandardLoggingKubernetesPodRouting | None]
applied_guardrails: list[str] | None
usage_object: dict | None
cold_storage_object_key: str | None # S3/GCS object key for cold storage retrieval

View file

@ -22,6 +22,7 @@ import litellm
from litellm._internal_context import in_post_response_phase
from litellm._logging import session_id_var, trace_id_var
from litellm.constants import REDACTED_BY_LITELLM, SENTRY_PII_DENYLIST
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
from litellm.cost_calculator import ocr_batch_cost
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
@ -4942,6 +4943,97 @@ def test_get_standard_logging_object_payload_takes_used_client_oauth_token_from_
assert payload["metadata"]["used_client_oauth_token"] is expected
@pytest.mark.parametrize(
("metadata", "litellm_metadata", "expected"),
[
(
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
}
},
{},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
},
),
(
{},
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.2",
"pod_count": 3,
"selection": "session_affinity",
}
},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.2",
"pod_count": 3,
"selection": "session_affinity",
},
),
(
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
}
},
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.3",
"pod_count": 3,
"selection": "session_affinity_retry",
}
},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.3",
"pod_count": 3,
"selection": "session_affinity_retry",
},
),
],
)
def test_standard_logging_payload_resolves_kubernetes_pod_routing_from_metadata_buckets(
logging_obj,
metadata: dict[str, object],
litellm_metadata: dict[str, object],
expected: dict[str, object],
) -> None:
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import get_standard_logging_object_payload
now: Final = datetime.now()
payload: Final = get_standard_logging_object_payload(
kwargs={
"model": "gpt-4o",
"messages": [],
"litellm_params": {"metadata": metadata, "litellm_metadata": litellm_metadata},
},
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logging_obj,
status="success",
)
assert payload is not None
assert payload["metadata"][KUBERNETES_POD_ROUTING_KEY] == expected
def test_get_standard_logging_object_payload_carries_matched_access_groups(logging_obj):
"""Access groups stamped at auth time reach the logging payload, so integrations see what a request billed."""
from datetime import datetime

View file

@ -12,6 +12,7 @@ from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm.constants import (
KUBERNETES_POD_ROUTING_KEY,
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
@ -74,6 +75,88 @@ def _get_additional_usage_values_for_usage(usage: litellm.Usage) -> dict:
return metadata["additional_usage_values"]
@pytest.mark.parametrize(
("metadata", "litellm_metadata", "expected"),
[
(
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
}
},
{},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
},
),
(
{},
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.2",
"pod_count": 3,
"selection": "session_affinity",
}
},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.2",
"pod_count": 3,
"selection": "session_affinity",
},
),
(
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "round_robin",
}
},
{
"kubernetes_pod_routing": {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.3",
"pod_count": 3,
"selection": "session_affinity_retry",
}
},
{
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.3",
"pod_count": 3,
"selection": "session_affinity_retry",
},
),
],
)
def test_get_logging_payload_includes_kubernetes_pod_routing(
metadata: dict[str, object],
litellm_metadata: dict[str, object],
expected: dict[str, object],
) -> None:
payload: Final = get_logging_payload(
kwargs={
"model": "gpt-4o-mini",
"litellm_params": {"metadata": metadata, "litellm_metadata": litellm_metadata},
},
response_obj=litellm.ModelResponse(id="chatcmpl-test", choices=[]),
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert json.loads(payload["metadata"])[KUBERNETES_POD_ROUTING_KEY] == expected
assert _get_spend_logs_metadata(None)["kubernetes_pod_routing"] is None
@pytest.mark.parametrize("store_prompts,redact", [(True, False), (False, False), (True, True)])
def test_classifier_audit_spend_storage_obeys_privacy_and_truncation(monkeypatch, store_prompts, redact):
from litellm.proxy import proxy_server

View file

@ -1,5 +1,6 @@
import asyncio
import copy
import hashlib
import json
import logging
import re
@ -16,10 +17,18 @@ import respx
from openai import AsyncOpenAI, OpenAI
import litellm
from litellm.constants import KUBERNETES_POD_ROUTING_KEY, SESSION_ID_GENERATED_METADATA_KEY
from litellm.router import CustomRoutingStrategyBase, Router
from litellm.router_utils.kubernetes_pod_discovery import KubernetesPodDiscovery
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,
)
_SERVICE_URL: Final = "http://vllm-headless.ns.svc.cluster.local:8000/v1"
_THREE_POD_CHAT_PATTERN: Final = re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2|10\.0\.0\.3):8000/v1/chat/completions")
_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)))
@ -91,6 +100,31 @@ def _empty_proxy_environment() -> Mapping[str, str]:
return {}
def _patch_router_pod_discovery(
monkeypatch: pytest.MonkeyPatch,
*,
clock: Callable[[], float] | None = None,
refresh_interval_seconds: float | None = None,
) -> None:
def create_discovery() -> KubernetesPodDiscovery:
if refresh_interval_seconds is not None and clock is not None:
return KubernetesPodDiscovery(
refresh_interval_seconds=refresh_interval_seconds,
clock=clock,
proxy_environment=_empty_proxy_environment,
)
if refresh_interval_seconds is not None:
return KubernetesPodDiscovery(
refresh_interval_seconds=refresh_interval_seconds,
proxy_environment=_empty_proxy_environment,
)
if clock is not None:
return KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment)
return KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
monkeypatch.setattr("litellm.router.KubernetesPodDiscovery", create_discovery)
def _stub_sync_dns(monkeypatch: pytest.MonkeyPatch, *ips: str) -> None:
def getaddrinfo(
host: str | None,
@ -113,12 +147,53 @@ def _api_base(deployment: Mapping[str, object]) -> str:
return api_base
def _ranked_session_ips(session_id: str) -> tuple[str, ...]:
ips: Final = ("10.0.0.1", "10.0.0.2", "10.0.0.3")
return tuple(
sorted(
ips,
key=lambda ip: (
-int.from_bytes(hashlib.blake2b(f"{session_id}\x00{ip}".encode(), digest_size=8).digest(), "big"),
ip,
),
)
)
def _request_body(content: bytes) -> Mapping[str, object]:
body: Final = json.loads(content)
assert isinstance(body, dict)
return cast(Mapping[str, object], body)
def _cache_async_client(router: Router) -> AsyncOpenAI:
client: Final = AsyncOpenAI(api_key="fake", base_url=_SERVICE_URL)
router.cache.set_cache(key="registered-id_async_client", value=client, local_only=True)
return client
def _cache_sync_client(router: Router) -> OpenAI:
client: Final = OpenAI(api_key="fake", base_url=_SERVICE_URL)
router.cache.set_cache(key="registered-id_client", value=client, local_only=True)
return client
async def _router_acompletion(router: Router, **request_kwargs: object) -> None:
await router.acompletion(
model="gpu-model",
messages=[{"role": "user", "content": "hello"}],
**request_kwargs,
)
def _router_completion(router: Router, **request_kwargs: object) -> None:
router.completion(
model="gpu-model",
messages=[{"role": "user", "content": "hello"}],
**request_kwargs,
)
def test_sync_resolution_round_robins_sorted_ips_without_mutating_deployment(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@ -700,6 +775,7 @@ async def test_router_sends_pod_hosts_without_forwarding_discovery_flag(
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
@ -709,10 +785,6 @@ async def test_router_sends_pod_hosts_without_forwarding_discovery_flag(
).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(
@ -750,6 +822,417 @@ async def test_router_sends_pod_hosts_without_forwarding_discovery_flag(
assert next(sync_lookups) == 1
@pytest.mark.asyncio
async def test_router_session_requests_stick_to_one_pod_for_sync_and_async(
monkeypatch: pytest.MonkeyPatch,
) -> None:
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", "10.0.0.3", 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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
async_router: Final = Router(model_list=[_deployment()])
sync_router: Final = Router(model_list=[_deployment()])
async_client: Final = _cache_async_client(async_router)
sync_client: Final = _cache_sync_client(sync_router)
for _ in range(6):
await _router_acompletion(async_router, metadata={"session_id": "s1"})
for _ in range(6):
_router_completion(sync_router, metadata={"session_id": "s1"})
hosts: Final = tuple(call.request.url.host for call in route.calls)
await async_client.close()
sync_client.close()
assert len(set(hosts[:6])) == 1
assert len(set(hosts[6:])) == 1
assert next(async_lookups) == 1
assert next(sync_lookups) == 1
@pytest.mark.asyncio
async def test_router_session_mapping_repeats_across_pods(
monkeypatch: pytest.MonkeyPatch,
) -> None:
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.1", "10.0.0.2", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
session_ids: Final = tuple(f"session-{index}" for index in range(30))
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for session_id in session_ids:
await _router_acompletion(router, metadata={"session_id": session_id})
await _router_acompletion(router)
for session_id in session_ids:
await _router_acompletion(router, metadata={"session_id": session_id})
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
first_pass: Final = hosts[: len(session_ids)]
second_pass: Final = hosts[len(session_ids) + 1 :]
assert len(set(first_pass)) > 1
assert hosts[len(session_ids)] == "10.0.0.1"
assert second_pass == first_pass
assert next(lookups) == 1
@pytest.mark.asyncio
async def test_generated_session_id_uses_round_robin_pod_selection(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(
router,
metadata={
"session_id": "generated-session",
SESSION_ID_GENERATED_METADATA_KEY: True,
},
)
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert hosts == ("10.0.0.1", "10.0.0.2", "10.0.0.3") * 2
@pytest.mark.asyncio
async def test_empty_session_id_uses_round_robin_pod_selection(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(router, metadata={"session_id": ""})
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert hosts == ("10.0.0.1", "10.0.0.2", "10.0.0.3") * 2
@pytest.mark.asyncio
async def test_session_requests_do_not_advance_round_robin_cursor(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
await _router_acompletion(router)
await _router_acompletion(router, metadata={"session_id": "s1"})
await _router_acompletion(router, metadata={"session_id": "s1"})
await _router_acompletion(router)
await _router_acompletion(router, metadata={"session_id": "s1"})
await _router_acompletion(router)
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert tuple(hosts[index] for index in (0, 3, 5)) == (
"10.0.0.1",
"10.0.0.2",
"10.0.0.3",
)
@pytest.mark.asyncio
async def test_session_pod_membership_change_remaps_only_affected_sessions(
monkeypatch: pytest.MonkeyPatch,
) -> None:
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]:
lookup_number: Final = next(lookups)
ips: Final = ("10.0.0.1", "10.0.0.2", "10.0.0.3") if lookup_number == 0 else ("10.0.0.2", "10.0.0.3")
return _records(*ips, port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
session_ids: Final = tuple(f"session-{index}" for index in range(60))
_patch_router_pod_discovery(
monkeypatch,
refresh_interval_seconds=10,
clock=_clock((0.0,) * 62 + (10.0,) * 10),
)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for session_id in session_ids:
await _router_acompletion(router, metadata={"session_id": session_id})
initial_hosts: Final = tuple(call.request.url.host for call in route.calls)
removed_session_ids: Final = tuple(
session_id for session_id, host in zip(session_ids, initial_hosts) if host == "10.0.0.1"
)
surviving_session_ids: Final = tuple(
session_id for session_id, host in zip(session_ids, initial_hosts) if host == "10.0.0.2"
)
assert removed_session_ids
assert len(surviving_session_ids) >= 2
await _router_acompletion(router)
await _router_acompletion(router, metadata={"session_id": surviving_session_ids[0]})
await _router_acompletion(router, metadata={"session_id": removed_session_ids[0]})
await _router_acompletion(router, metadata={"session_id": surviving_session_ids[1]})
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert hosts[60] == "10.0.0.1"
assert hosts[61] == "10.0.0.2"
assert hosts[62] in ("10.0.0.2", "10.0.0.3")
assert hosts[63] == "10.0.0.2"
assert next(lookups) == 2
@pytest.mark.asyncio
async def test_custom_routing_strategy_keeps_sync_and_async_sessions_on_one_pod(
monkeypatch: pytest.MonkeyPatch,
) -> None:
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", "10.0.0.3", 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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
async_router: Final = Router(model_list=[_deployment()])
sync_router: Final = Router(model_list=[_deployment()])
async_router.set_custom_routing_strategy(_ModelListRoutingStrategy(async_router))
sync_router.set_custom_routing_strategy(_ModelListRoutingStrategy(sync_router))
async_client: Final = _cache_async_client(async_router)
sync_client: Final = _cache_sync_client(sync_router)
for _ in range(6):
await _router_acompletion(async_router, metadata={"session_id": "custom-session"})
for _ in range(6):
_router_completion(sync_router, metadata={"session_id": "custom-session"})
hosts: Final = tuple(call.request.url.host for call in route.calls)
await async_client.close()
sync_client.close()
assert len(set(hosts[:6])) == 1
assert len(set(hosts[6:])) == 1
assert next(async_lookups) == 1
assert next(sync_lookups) == 1
@pytest.mark.asyncio
async def test_top_level_litellm_session_id_keeps_requests_on_one_pod(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(router, litellm_session_id="s1")
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert len(set(hosts)) == 1
@pytest.mark.asyncio
async def test_generated_metadata_ignores_top_level_litellm_session_id(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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", "10.0.0.3", port=int(port or 0))
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
router: Final = Router(model_list=[_deployment()])
client: Final = _cache_async_client(router)
for _ in range(6):
await _router_acompletion(
router,
litellm_session_id="s1",
metadata={
"session_id": "s1",
SESSION_ID_GENERATED_METADATA_KEY: True,
},
)
hosts: Final = tuple(call.request.url.host for call in route.calls)
await client.close()
assert hosts == ("10.0.0.1", "10.0.0.2", "10.0.0.3") * 2
@pytest.mark.asyncio
async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_requests(
monkeypatch: pytest.MonkeyPatch,
@ -784,6 +1267,7 @@ async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_req
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
@ -793,14 +1277,6 @@ async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_req
).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):
@ -852,16 +1328,13 @@ async def test_router_async_retry_uses_next_discovered_pod(monkeypatch: pytest.M
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
router.kubernetes_pod_discovery = KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
)
response: Final = await router.acompletion(
model="gpu-model",
@ -895,16 +1368,13 @@ def test_router_sync_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPa
monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
_patch_router_pod_discovery(monkeypatch, clock=clock)
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
router.kubernetes_pod_discovery = KubernetesPodDiscovery(
clock=clock,
proxy_environment=_empty_proxy_environment,
)
response: Final = router.completion(
model="gpu-model",
@ -914,3 +1384,240 @@ def test_router_sync_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPa
assert response.choices[0].message.content == "ok"
assert hosts == ("10.0.0.1", "10.0.0.2")
@pytest.mark.parametrize(
("litellm_retry_count", "metadata_retry_count", "expected_rank"),
[
(0, None, 0),
(None, None, 0),
(True, None, 0),
("1", None, 0),
(-1, None, 0),
(1, None, 1),
(True, 1, 1),
(1, 0, 1),
],
)
def test_session_retry_count_selects_ranked_pod(
monkeypatch: pytest.MonkeyPatch,
litellm_retry_count: object,
metadata_retry_count: object,
expected_rank: int,
) -> None:
_stub_sync_dns(monkeypatch, "10.0.0.1", "10.0.0.2", "10.0.0.3")
discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
request_kwargs: Final = {
"litellm_metadata": {
"session_id": "s1",
**({"request_retry_count": litellm_retry_count} if litellm_retry_count is not None else {}),
},
"metadata": ({"request_retry_count": metadata_retry_count} if metadata_retry_count is not None else {}),
}
resolved: Final = discovery.resolve_deployment(_deployment(), request_kwargs)
assert httpx.URL(_api_base(resolved)).host == _ranked_session_ips("s1")[expected_rank]
@pytest.mark.asyncio
async def test_router_session_retry_uses_next_rendezvous_pod(monkeypatch: pytest.MonkeyPatch) -> None:
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", "10.0.0.3", 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)
_patch_router_pod_discovery(monkeypatch)
session_id: Final = "retry-session"
with respx.mock(assert_all_called=True) as respx_mock:
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(side_effect=response_for)
router: Final = Router(model_list=[_deployment()], num_retries=1)
response: Final = await router.acompletion(
model="gpu-model",
messages=[{"role": "user", "content": "hello"}],
metadata={"session_id": session_id},
)
hosts: Final = tuple(call.request.url.host for call in route.calls)
assert response.choices[0].message.content == "ok"
assert hosts == _ranked_session_ips(session_id)[:2]
@pytest.mark.asyncio
async def test_positional_request_kwargs_reach_all_pod_resolver_decorators(
monkeypatch: pytest.MonkeyPatch,
) -> None:
loop: Final = asyncio.get_running_loop()
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]:
return _records("10.0.0.1", "10.0.0.2", "10.0.0.3", 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]:
return _records("10.0.0.1", "10.0.0.2", "10.0.0.3", port=int(port or 0))
class _Selector:
def __init__(self, discovery: KubernetesPodDiscovery) -> None:
self.kubernetes_pod_discovery: Final = discovery
@resolve_pods_after
def sync(self, deployment: dict[str, object], request_kwargs: Mapping[str, object]) -> dict[str, object]:
return deployment
@async_resolve_pods_after
async def async_call(
self,
deployment: dict[str, object],
request_kwargs: Mapping[str, object],
) -> dict[str, object]:
return deployment
class _BoundSelector:
def sync(self, deployment: dict[str, object], request_kwargs: Mapping[str, object]) -> dict[str, object]:
return deployment
async def async_call(
self,
deployment: dict[str, object],
request_kwargs: Mapping[str, object],
) -> dict[str, object]:
return deployment
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
session_id: Final = next(
candidate
for candidate in (f"positional-{index}" for index in range(100))
if _ranked_session_ips(candidate)[0] != "10.0.0.1"
)
request_kwargs: Final = {"metadata": {"session_id": session_id}}
sync_selector: Final = _Selector(KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment))
async_selector: Final = _Selector(KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment))
bound_selector: Final = _BoundSelector()
bound_sync: Final = resolve_pods_after_bound(
bound_selector.sync,
KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment),
)
bound_async: Final = async_resolve_pods_after_bound(
bound_selector.async_call,
KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment),
)
sync_result: Final = sync_selector.sync(_deployment(), request_kwargs)
async_result: Final = await async_selector.async_call(_deployment(), request_kwargs)
bound_sync_result: Final = bound_sync(_deployment(), request_kwargs)
bound_async_result: Final = await bound_async(_deployment(), request_kwargs)
assert (
tuple(
httpx.URL(_api_base(result)).host
for result in (sync_result, async_result, bound_sync_result, bound_async_result)
)
== (_ranked_session_ips(session_id)[0],) * 4
)
def test_direct_session_id_keyword_does_not_enable_session_affinity(monkeypatch: pytest.MonkeyPatch) -> None:
_stub_sync_dns(monkeypatch, "10.0.0.1", "10.0.0.2", "10.0.0.3")
discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
request_kwargs: Final = {"session_id": "s1", "messages": [{"role": "user", "content": "hello"}]}
selected_ips: Final = tuple(
httpx.URL(_api_base(discovery.resolve_deployment(_deployment(), request_kwargs))).host for _ in range(3)
)
assert selected_ips == ("10.0.0.1", "10.0.0.2", "10.0.0.3")
@pytest.mark.parametrize(
("request_kwargs", "selection"),
[
({}, "round_robin"),
({"metadata": {"session_id": "s1"}}, "session_affinity"),
({"metadata": {"session_id": "s1", "request_retry_count": 1}}, "session_affinity_retry"),
],
)
def test_resolved_deployment_records_pod_selection(
monkeypatch: pytest.MonkeyPatch,
request_kwargs: dict[str, object],
selection: str,
) -> None:
_stub_sync_dns(monkeypatch, "10.0.0.1", "10.0.0.2", "10.0.0.3")
discovery: Final = KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
resolved: Final = discovery.resolve_deployment(_deployment(), request_kwargs)
pod_ip: Final = httpx.URL(_api_base(resolved)).host
assert resolved[KUBERNETES_POD_ROUTING_KEY] == {
"service_host": httpx.URL(_SERVICE_URL).host,
"pod_ip": pod_ip,
"pod_count": 3,
"selection": selection,
}
@pytest.mark.parametrize("fallback", ["https", "no_proxy", "empty_ips", "dns_failure"])
def test_unresolved_deployment_does_not_record_pod_selection(
monkeypatch: pytest.MonkeyPatch,
fallback: str,
) -> None:
if fallback == "empty_ips":
_stub_sync_dns(monkeypatch)
if fallback == "dns_failure":
def failing_getaddrinfo(
host: str | None,
port: str | int | None,
family: int = 0,
type: int = 0,
proto: int = 0,
flags: int = 0,
) -> list[_AddrInfo]:
raise socket.gaierror(socket.EAI_AGAIN, "temporary failure")
monkeypatch.setattr(socket, "getaddrinfo", failing_getaddrinfo)
if fallback == "no_proxy":
_stub_sync_dns(monkeypatch, "10.0.0.1")
proxy_environment: Final = (
_proxy_environment("vllm-headless.ns.svc.cluster.local") if fallback == "no_proxy" else _empty_proxy_environment
)
api_base: Final = "https://vllm-headless.ns.svc.cluster.local:8000/v1" if fallback == "https" else _SERVICE_URL
discovery: Final = KubernetesPodDiscovery(proxy_environment=proxy_environment)
deployment: Final = _deployment(api_base=api_base)
resolved: Final = discovery.resolve_deployment(deployment)
assert resolved is deployment
assert KUBERNETES_POD_ROUTING_KEY not in resolved

View file

@ -7865,6 +7865,36 @@ def test_update_kwargs_with_deployment_model_info_in_metadata():
assert model_info["output_cost_per_token"] == 0.0015
def test_update_kwargs_with_deployment_clears_pod_routing_on_non_discovery_fallback():
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
routing_record: Final = {
"service_host": "vllm-headless.ns.svc.cluster.local",
"pod_ip": "10.0.0.1",
"pod_count": 3,
"selection": "session_affinity",
}
discovery_deployment: Final = {
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
"model_info": {"id": "discovery-id"},
KUBERNETES_POD_ROUTING_KEY: routing_record,
}
fallback_deployment: Final = {
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
"model_info": {"id": "fallback-id"},
}
router: Final = litellm.Router(model_list=[discovery_deployment])
kwargs: Final = {"metadata": {}}
router._update_kwargs_with_deployment(deployment=discovery_deployment, kwargs=kwargs)
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] == routing_record
router._update_kwargs_with_deployment(deployment=fallback_deployment, kwargs=kwargs)
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] is None
def test_combine_fallback_usage():
"""Test that _combine_fallback_usage merges partial and fallback usage."""
from litellm.router import Router

View file

@ -431,6 +431,37 @@ describe("LogDetailContent", () => {
expect(screen.queryByText("Configured key")).not.toBeInTheDocument();
});
it.each([
{ selection: "round_robin", label: "Round robin" },
{ selection: "session_affinity", label: "Session affinity" },
{ selection: "session_affinity_retry", label: "Session affinity (retry)" },
] as const)("shows the selected Kubernetes pod and routing mode", ({ selection, label }) => {
render(
<LogDetailContent
logEntry={createLogEntry({
metadata: {
status: "success",
kubernetes_pod_routing: {
service_host: "vllm-headless.ns.svc.cluster.local",
pod_ip: "10.0.0.2",
pod_count: 3,
selection,
},
},
})}
/>,
);
expect(screen.getByText("Pod")).toBeInTheDocument();
expect(screen.getByText(`10.0.0.2 · ${label} · 3 pods`)).toBeInTheDocument();
});
it("omits the Pod row when routing metadata is absent", () => {
render(<LogDetailContent logEntry={createLogEntry({ metadata: { status: "success" } })} />);
expect(screen.queryByText("Pod")).not.toBeInTheDocument();
});
it("should display guardrail label when guardrail data exists", () => {
render(
<LogDetailContent

View file

@ -23,7 +23,7 @@ import {
import { CostBreakdownViewer } from "../CostBreakdownViewer";
import { ConfigInfoMessage } from "../ConfigInfoMessage";
import { VectorStoreViewer } from "../VectorStoreViewer";
import { CREDENTIAL_LABELS } from "../constants";
import { CREDENTIAL_LABELS, KUBERNETES_POD_ROUTING_LABELS } from "../constants";
import { TruncatedValue } from "./TruncatedValue";
import { TokenFlow } from "./TokenFlow";
import { JsonViewer } from "./JsonViewer";
@ -75,6 +75,7 @@ export function LogDetailContent({
userEmail,
}: LogDetailContentProps) {
const metadata = logEntry.metadata || {};
const podRouting = metadata.kubernetes_pod_routing;
const hasError = metadata.status === "failure";
const errorInfo = hasError ? metadata.error_information : null;
const isClassifier =
@ -160,6 +161,12 @@ export function LogDetailContent({
<DescriptionItem label="API Base">
<TruncatedValue value={logEntry.api_base} maxWidth={API_BASE_MAX_WIDTH} />
</DescriptionItem>
{podRouting && typeof podRouting === "object" && (
<DescriptionItem label="Pod">
{podRouting.pod_ip} · {KUBERNETES_POD_ROUTING_LABELS[podRouting.selection]} ·{" "}
{podRouting.pod_count} pods
</DescriptionItem>
)}
{logEntry.requester_ip_address && (
<DescriptionItem label="IP Address">{logEntry.requester_ip_address}</DescriptionItem>
)}

View file

@ -10,6 +10,13 @@ export const LOGS_SORT_FIELD_MAP = {
export type LogsSortField = keyof typeof LOGS_SORT_FIELD_MAP;
export type KubernetesPodRouting = {
service_host: string;
pod_ip: string;
pod_count: number;
selection: "round_robin" | "session_affinity" | "session_affinity_retry";
};
export type LogEntry = {
request_id: string;
litellm_call_id?: string | null;
@ -29,7 +36,9 @@ export type LogEntry = {
user?: string;
end_user?: string;
custom_llm_provider?: string;
metadata?: Record<string, any>;
metadata?: Record<string, any> & {
kubernetes_pod_routing?: KubernetesPodRouting | null;
};
cache_hit: string;
cache_key?: string;
request_tags?: Record<string, any>;

View file

@ -33,6 +33,12 @@ export const CREDENTIAL_LABELS: Record<string, string> = {
false: "Configured key",
};
export const KUBERNETES_POD_ROUTING_LABELS = {
round_robin: "Round robin",
session_affinity: "Session affinity",
session_affinity_retry: "Session affinity (retry)",
} as const;
export const QUICK_SELECT_OPTIONS: { label: string; value: number; unit: string }[] = [
{ label: "Last Minute", value: 1, unit: "minutes" },
{ label: "Last 15 Minutes", value: 15, unit: "minutes" },