mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
b9e4f0613e
commit
277f4b4218
17 changed files with 1257 additions and 48 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue