diff --git a/litellm/constants.py b/litellm/constants.py
index 61ac302849c..540ac1fb6ae 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -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))
diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py
index 8d588896b2f..59009109cc0 100644
--- a/litellm/integrations/opentelemetry.py
+++ b/litellm/integrations/opentelemetry.py
@@ -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 (
diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py
index 39e95fbf687..420ae7bb372 100644
--- a/litellm/litellm_core_utils/core_helpers.py
+++ b/litellm/litellm_core_utils/core_helpers.py
@@ -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
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index 05393049f58..f41b8e8ac93 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -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(
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 0abec51cc49..762f31b0232 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -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
diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py
index b71b834a31c..60c0d06ddab 100644
--- a/litellm/proxy/spend_tracking/spend_tracking_utils.py
+++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py
@@ -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)
diff --git a/litellm/router.py b/litellm/router.py
index 4b16bc42fd6..21e91bf763d 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -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),
}
)
diff --git a/litellm/router_utils/kubernetes_pod_discovery.py b/litellm/router_utils/kubernetes_pod_discovery.py
index 00a3a6b55ef..8f539a2f99d 100644
--- a/litellm/router_utils/kubernetes_pod_discovery.py
+++ b/litellm/router_utils/kubernetes_pod_discovery.py
@@ -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
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 6919fd6fd27..3e500fed576 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -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
diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py
index 4b25da2ff79..df61b724838 100644
--- a/tests/unit/litellm_core_utils/test_litellm_logging.py
+++ b/tests/unit/litellm_core_utils/test_litellm_logging.py
@@ -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
diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py
index a3de9328437..f18344a8cea 100644
--- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py
+++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py
@@ -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
diff --git a/tests/unit/router_utils/test_kubernetes_pod_discovery.py b/tests/unit/router_utils/test_kubernetes_pod_discovery.py
index eea2da4e388..47979dbf049 100644
--- a/tests/unit/router_utils/test_kubernetes_pod_discovery.py
+++ b/tests/unit/router_utils/test_kubernetes_pod_discovery.py
@@ -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
diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py
index 96dddf15869..22f76370ae9 100644
--- a/tests/unit/test_router/test_router.py
+++ b/tests/unit/test_router/test_router.py
@@ -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
diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx
index e0551c61062..bbdd0323ffb 100644
--- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx
@@ -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(
+ ,
+ );
+
+ 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();
+
+ expect(screen.queryByText("Pod")).not.toBeInTheDocument();
+ });
+
it("should display guardrail label when guardrail data exists", () => {
render(
+ {podRouting && typeof podRouting === "object" && (
+
+ {podRouting.pod_ip} · {KUBERNETES_POD_ROUTING_LABELS[podRouting.selection]} ·{" "}
+ {podRouting.pod_count} pods
+
+ )}
{logEntry.requester_ip_address && (
{logEntry.requester_ip_address}
)}
diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.tsx b/ui/litellm-dashboard/src/components/view_logs/columns.tsx
index ce4a72a6d38..dc5fa4aaf5b 100644
--- a/ui/litellm-dashboard/src/components/view_logs/columns.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/columns.tsx
@@ -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;
+ metadata?: Record & {
+ kubernetes_pod_routing?: KubernetesPodRouting | null;
+ };
cache_hit: string;
cache_key?: string;
request_tags?: Record;
diff --git a/ui/litellm-dashboard/src/components/view_logs/constants.ts b/ui/litellm-dashboard/src/components/view_logs/constants.ts
index 9474c54d6b6..44231370521 100644
--- a/ui/litellm-dashboard/src/components/view_logs/constants.ts
+++ b/ui/litellm-dashboard/src/components/view_logs/constants.ts
@@ -33,6 +33,12 @@ export const CREDENTIAL_LABELS: Record = {
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" },