mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge pull request #21763 from Harshit28j/litellm_feat_sticky_sessions
feat: add session_id to have better routing
This commit is contained in:
commit
e463fc22d9
6 changed files with 393 additions and 81 deletions
|
|
@ -887,7 +887,8 @@ router = litellm.Router(
|
|||
# `responses_api_deployment_check` ensures Requests with `previous_response_id`
|
||||
# are routed to the same deployment. `deployment_affinity` adds sticky sessions
|
||||
# for requests without `previous_response_id` (useful for implicit caching).
|
||||
optional_pre_call_checks=["responses_api_deployment_check", "deployment_affinity"],
|
||||
# `session_affinity` adds sticky sessions based on `session_id` metadata.
|
||||
optional_pre_call_checks=["responses_api_deployment_check", "deployment_affinity", "session_affinity"],
|
||||
# Optional (default is 3600 seconds / 1 hour)
|
||||
deployment_affinity_ttl_seconds=3600,
|
||||
)
|
||||
|
|
@ -919,10 +920,12 @@ follow_up = await router.aresponses(
|
|||
To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks` in your proxy config.yaml.
|
||||
|
||||
- `responses_api_deployment_check`: high priority routing when `previous_response_id` is provided
|
||||
- `session_affinity`: sticky sessions based on session id (takes priority over `deployment_affinity`)
|
||||
- `deployment_affinity`: sticky sessions based on user key (applies even without `previous_response_id`)
|
||||
|
||||
Notes:
|
||||
- User-key affinity is keyed on `metadata.user_api_key_hash` (the API key hash). The OpenAI `user` request parameter is an end-user identifier and is intentionally not used for deployment affinity.
|
||||
- Session-ID affinity is keyed on `metadata.session_id`. For proxy requests, this can be passed via the `x-litellm-session-id` HTTP header. For Python SDK requests, you can pass it via `litellm_metadata={"session_id": "value"}` in request args.
|
||||
- `user_api_key_hash` is already SHA-256, and is used as-is (no double hashing).
|
||||
- Affinity is scoped by a stable model identifier (the model-map key, e.g. `model_map_information.model_map_key`) so model aliases map to the same stickiness bucket.
|
||||
- The mapping TTL is controlled by `deployment_affinity_ttl_seconds` (configured on Router init / proxy startup).
|
||||
|
|
@ -945,6 +948,7 @@ model_list:
|
|||
router_settings:
|
||||
optional_pre_call_checks:
|
||||
- responses_api_deployment_check
|
||||
- session_affinity
|
||||
- deployment_affinity
|
||||
# Optional (default is 3600 seconds / 1 hour)
|
||||
deployment_affinity_ttl_seconds: 3600
|
||||
|
|
|
|||
|
|
@ -569,13 +569,25 @@ class LiteLLMProxyRequestSetup:
|
|||
#########################################################################################
|
||||
agent_id_from_header = headers.get("x-litellm-agent-id")
|
||||
trace_id_from_header = headers.get("x-litellm-trace-id")
|
||||
session_id_from_header = headers.get("x-litellm-session-id")
|
||||
|
||||
if agent_id_from_header:
|
||||
metadata_from_headers["agent_id"] = agent_id_from_header
|
||||
verbose_proxy_logger.debug(f"Extracted agent_id from header: {agent_id_from_header}")
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracted agent_id from header: {agent_id_from_header}"
|
||||
)
|
||||
|
||||
if trace_id_from_header:
|
||||
metadata_from_headers["trace_id"] = trace_id_from_header
|
||||
verbose_proxy_logger.debug(f"Extracted trace_id from header: {trace_id_from_header}")
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracted trace_id from header: {trace_id_from_header}"
|
||||
)
|
||||
|
||||
if session_id_from_header:
|
||||
metadata_from_headers["session_id"] = session_id_from_header
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracted session_id from header: {session_id_from_header}"
|
||||
)
|
||||
|
||||
if isinstance(data[_metadata_variable_name], dict):
|
||||
data[_metadata_variable_name].update(metadata_from_headers)
|
||||
|
|
@ -1589,9 +1601,7 @@ def _match_and_track_policies(
|
|||
for name in applied_policy_names
|
||||
if name in policy_reasons
|
||||
}
|
||||
add_policy_sources_to_metadata(
|
||||
request_data=data, policy_sources=applied_reasons
|
||||
)
|
||||
add_policy_sources_to_metadata(request_data=data, policy_sources=applied_reasons)
|
||||
|
||||
return applied_policy_names, policy_reasons
|
||||
|
||||
|
|
@ -1626,9 +1636,9 @@ def _apply_resolved_guardrails_to_metadata(
|
|||
pipelines
|
||||
)
|
||||
data[metadata_variable_name]["_guardrail_pipelines"] = pipelines
|
||||
data[metadata_variable_name]["_pipeline_managed_guardrails"] = (
|
||||
pipeline_managed_guardrails
|
||||
)
|
||||
data[metadata_variable_name][
|
||||
"_pipeline_managed_guardrails"
|
||||
] = pipeline_managed_guardrails
|
||||
verbose_proxy_logger.debug(
|
||||
f"Policy engine: resolved {len(pipelines)} pipeline(s), "
|
||||
f"managed guardrails: {pipeline_managed_guardrails}"
|
||||
|
|
|
|||
|
|
@ -1197,7 +1197,12 @@ class Router:
|
|||
enable_responses_api_affinity = (
|
||||
"responses_api_deployment_check" in optional_pre_call_checks
|
||||
)
|
||||
if enable_user_key_affinity or enable_responses_api_affinity:
|
||||
enable_session_id_affinity = "session_affinity" in optional_pre_call_checks
|
||||
if (
|
||||
enable_user_key_affinity
|
||||
or enable_responses_api_affinity
|
||||
or enable_session_id_affinity
|
||||
):
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
|
||||
|
|
@ -1216,6 +1221,10 @@ class Router:
|
|||
existing_affinity_callback.enable_responses_api_affinity
|
||||
or enable_responses_api_affinity
|
||||
)
|
||||
existing_affinity_callback.enable_session_id_affinity = (
|
||||
existing_affinity_callback.enable_session_id_affinity
|
||||
or enable_session_id_affinity
|
||||
)
|
||||
existing_affinity_callback.ttl_seconds = (
|
||||
self.deployment_affinity_ttl_seconds
|
||||
)
|
||||
|
|
@ -1225,11 +1234,10 @@ class Router:
|
|||
ttl_seconds=self.deployment_affinity_ttl_seconds,
|
||||
enable_user_key_affinity=enable_user_key_affinity,
|
||||
enable_responses_api_affinity=enable_responses_api_affinity,
|
||||
enable_session_id_affinity=enable_session_id_affinity,
|
||||
)
|
||||
self.optional_callbacks.append(affinity_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(
|
||||
affinity_callback
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# Remaining optional pre-call checks
|
||||
|
|
@ -1239,6 +1247,7 @@ class Router:
|
|||
if pre_call_check in (
|
||||
"deployment_affinity",
|
||||
"responses_api_deployment_check",
|
||||
"session_affinity",
|
||||
):
|
||||
continue
|
||||
if pre_call_check == "prompt_caching":
|
||||
|
|
@ -1948,9 +1957,7 @@ class Router:
|
|||
return deployment_pydantic_obj
|
||||
|
||||
@staticmethod
|
||||
def _merge_tools_from_deployment(
|
||||
deployment: dict, kwargs: dict
|
||||
) -> None:
|
||||
def _merge_tools_from_deployment(deployment: dict, kwargs: dict) -> None:
|
||||
"""
|
||||
Merge tools from deployment litellm_params with request kwargs.
|
||||
When both have tools, concatenate them (deployment tools first, then request tools).
|
||||
|
|
@ -3779,7 +3786,7 @@ class Router:
|
|||
)
|
||||
raise e
|
||||
|
||||
async def _acreate_file( # noqa: PLR0915
|
||||
async def _acreate_file( # noqa: PLR0915
|
||||
self,
|
||||
model: str,
|
||||
**kwargs,
|
||||
|
|
@ -3840,8 +3847,12 @@ class Router:
|
|||
)
|
||||
|
||||
kwargs_copy["file"] = file
|
||||
if "gcs_bucket_name" in data: # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there
|
||||
kwargs_copy.setdefault("litellm_metadata", {})["gcs_bucket_name"] = data["gcs_bucket_name"]
|
||||
if (
|
||||
"gcs_bucket_name" in data
|
||||
): # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there
|
||||
kwargs_copy.setdefault("litellm_metadata", {})[
|
||||
"gcs_bucket_name"
|
||||
] = data["gcs_bucket_name"]
|
||||
response = litellm.acreate_file(
|
||||
**{
|
||||
**data,
|
||||
|
|
@ -3920,15 +3931,15 @@ class Router:
|
|||
):
|
||||
"""
|
||||
Create a vector store for a specific model.
|
||||
|
||||
|
||||
Args:
|
||||
model: Model name from router config
|
||||
**kwargs: Vector store creation parameters
|
||||
|
||||
|
||||
Returns:
|
||||
VectorStoreCreateResponse
|
||||
"""
|
||||
try:
|
||||
try:
|
||||
# If model is None, use the factory function approach (direct SDK call)
|
||||
if model is None:
|
||||
from litellm.vector_stores.main import acreate
|
||||
|
|
@ -3938,8 +3949,9 @@ class Router:
|
|||
acreate, call_type="avector_store_create"
|
||||
)
|
||||
return await factory_fn(**kwargs)
|
||||
|
||||
|
||||
from litellm.vector_stores import acreate as avector_store_create_sdk
|
||||
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
deployment = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
|
|
@ -3950,7 +3962,9 @@ class Router:
|
|||
data = deployment["litellm_params"].copy()
|
||||
model_name = data["model"]
|
||||
self._update_kwargs_with_deployment(
|
||||
deployment=deployment, kwargs=kwargs, function_name="avector_store_create"
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
function_name="avector_store_create",
|
||||
)
|
||||
|
||||
model_client = self._get_async_openai_model_client(
|
||||
|
|
@ -3961,7 +3975,7 @@ class Router:
|
|||
|
||||
# Get custom provider
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=data["model"])
|
||||
|
||||
|
||||
response = avector_store_create_sdk(
|
||||
**{
|
||||
**data,
|
||||
|
|
@ -4006,7 +4020,6 @@ class Router:
|
|||
self.fail_calls[model] += 1
|
||||
raise e
|
||||
|
||||
|
||||
def _override_vector_store_methods_for_router(self):
|
||||
"""
|
||||
Override factory-generated vector store methods with router-aware implementations.
|
||||
|
|
@ -4724,20 +4737,20 @@ class Router:
|
|||
):
|
||||
"""
|
||||
Initialize the Vector Store API endpoints on the router.
|
||||
|
||||
|
||||
If a model is provided in kwargs, use model-based routing to get
|
||||
the deployment credentials. Otherwise, call the original function directly.
|
||||
"""
|
||||
if custom_llm_provider and "custom_llm_provider" not in kwargs:
|
||||
kwargs["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
|
||||
# If model is provided, use generic API call with fallbacks for proper routing
|
||||
if kwargs.get("model"):
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
original_function=original_function,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
# Otherwise, call the original function directly
|
||||
return await original_function(**kwargs)
|
||||
|
||||
|
|
@ -5141,7 +5154,9 @@ class Router:
|
|||
)
|
||||
## ADD RETRY TRACKING TO METADATA - used for spend logs retry tracking
|
||||
_metadata["attempted_retries"] = 0
|
||||
_metadata["max_retries"] = num_retries # Updated after overrides in exception handler
|
||||
_metadata[
|
||||
"max_retries"
|
||||
] = num_retries # Updated after overrides in exception handler
|
||||
try:
|
||||
self._handle_mock_testing_rate_limit_error(
|
||||
model_group=model_group, kwargs=kwargs
|
||||
|
|
@ -5175,10 +5190,7 @@ class Router:
|
|||
# Check retry policy FIRST, before should_retry_this_error
|
||||
# This allows retry policies to override the healthy deployments check
|
||||
_retry_policy_applies = False
|
||||
if (
|
||||
self.retry_policy is not None
|
||||
or model_group_retry_policy is not None
|
||||
):
|
||||
if self.retry_policy is not None or model_group_retry_policy is not None:
|
||||
# get num_retries from retry policy
|
||||
# Use the model_group captured at the start of the function, or get it from metadata
|
||||
# kwargs.get("model") at this point is the deployment model, not the model_group
|
||||
|
|
@ -6173,9 +6185,7 @@ class Router:
|
|||
# unique model_id above.
|
||||
_custom_pricing_fields = CustomPricingLiteLLMParams.model_fields.keys()
|
||||
_shared_model_info = {
|
||||
k: v
|
||||
for k, v in _model_info.items()
|
||||
if k not in _custom_pricing_fields
|
||||
k: v for k, v in _model_info.items() if k not in _custom_pricing_fields
|
||||
}
|
||||
litellm.register_model(
|
||||
model_cost={
|
||||
|
|
@ -8365,9 +8375,7 @@ class Router:
|
|||
litellm_router_instance=self, parent_otel_span=parent_otel_span
|
||||
)
|
||||
if verbose_router_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_router_logger.debug(
|
||||
f"cooldown deployments: {cooldown_deployments}"
|
||||
)
|
||||
verbose_router_logger.debug(f"cooldown deployments: {cooldown_deployments}")
|
||||
healthy_deployments = self._filter_cooldown_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
cooldown_deployments=cooldown_deployments,
|
||||
|
|
|
|||
|
|
@ -45,12 +45,14 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
ttl_seconds: int,
|
||||
enable_user_key_affinity: bool,
|
||||
enable_responses_api_affinity: bool,
|
||||
enable_session_id_affinity: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.cache = cache
|
||||
self.ttl_seconds = ttl_seconds
|
||||
self.enable_user_key_affinity = enable_user_key_affinity
|
||||
self.enable_responses_api_affinity = enable_responses_api_affinity
|
||||
self.enable_session_id_affinity = enable_session_id_affinity
|
||||
|
||||
@staticmethod
|
||||
def _looks_like_sha256_hex(value: str) -> bool:
|
||||
|
|
@ -78,7 +80,9 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return hashlib.sha256(user_key.encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _get_model_map_key_from_litellm_model_name(litellm_model_name: str) -> Optional[str]:
|
||||
def _get_model_map_key_from_litellm_model_name(
|
||||
litellm_model_name: str,
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Best-effort derivation of a stable "model map key" for affinity scoping.
|
||||
|
||||
|
|
@ -133,8 +137,10 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return base_model
|
||||
litellm_model_name = litellm_params.get("model")
|
||||
if isinstance(litellm_model_name, str) and litellm_model_name:
|
||||
return DeploymentAffinityCheck._get_model_map_key_from_litellm_model_name(
|
||||
litellm_model_name
|
||||
return (
|
||||
DeploymentAffinityCheck._get_model_map_key_from_litellm_model_name(
|
||||
litellm_model_name
|
||||
)
|
||||
)
|
||||
|
||||
return None
|
||||
|
|
@ -175,6 +181,10 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
hashed_user_key = cls._hash_user_key(user_key=user_key)
|
||||
return f"{cls.CACHE_KEY_PREFIX}:{model_group}:{hashed_user_key}"
|
||||
|
||||
@classmethod
|
||||
def get_session_affinity_cache_key(cls, model_group: str, session_id: str) -> str:
|
||||
return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{session_id}"
|
||||
|
||||
@staticmethod
|
||||
def _get_user_key_from_metadata_dict(metadata: dict) -> Optional[str]:
|
||||
# NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the
|
||||
|
|
@ -184,6 +194,13 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return None
|
||||
return str(user_key)
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_metadata_dict(metadata: dict) -> Optional[str]:
|
||||
session_id = metadata.get("session_id")
|
||||
if session_id is None:
|
||||
return None
|
||||
return str(session_id)
|
||||
|
||||
@staticmethod
|
||||
def _iter_metadata_dicts(request_kwargs: dict) -> List[dict]:
|
||||
"""
|
||||
|
|
@ -220,13 +237,27 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _find_deployment_by_model_id(healthy_deployments: List[dict], model_id: str) -> Optional[dict]:
|
||||
def _get_session_id_from_request_kwargs(request_kwargs: dict) -> Optional[str]:
|
||||
for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs):
|
||||
session_id = DeploymentAffinityCheck._get_session_id_from_metadata_dict(
|
||||
metadata=metadata
|
||||
)
|
||||
if session_id is not None:
|
||||
return session_id
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _find_deployment_by_model_id(
|
||||
healthy_deployments: List[dict], model_id: str
|
||||
) -> Optional[dict]:
|
||||
for deployment in healthy_deployments:
|
||||
model_info = deployment.get("model_info")
|
||||
if not isinstance(model_info, dict):
|
||||
continue
|
||||
deployment_model_id = model_info.get("id")
|
||||
if deployment_model_id is not None and str(deployment_model_id) == str(model_id):
|
||||
if deployment_model_id is not None and str(deployment_model_id) == str(
|
||||
model_id
|
||||
):
|
||||
return deployment
|
||||
return None
|
||||
|
||||
|
|
@ -250,7 +281,11 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
if self.enable_responses_api_affinity:
|
||||
previous_response_id = request_kwargs.get("previous_response_id")
|
||||
if previous_response_id is not None:
|
||||
responses_model_id = ResponsesAPIRequestUtils.get_model_id_from_response_id(str(previous_response_id))
|
||||
responses_model_id = (
|
||||
ResponsesAPIRequestUtils.get_model_id_from_response_id(
|
||||
str(previous_response_id)
|
||||
)
|
||||
)
|
||||
if responses_model_id is not None:
|
||||
deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments=typed_healthy_deployments,
|
||||
|
|
@ -263,7 +298,52 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
)
|
||||
return [deployment]
|
||||
|
||||
# 2) User key -> deployment affinity
|
||||
stable_model_map_key = self._get_stable_model_map_key_from_deployments(
|
||||
healthy_deployments=typed_healthy_deployments
|
||||
)
|
||||
if stable_model_map_key is None:
|
||||
return typed_healthy_deployments
|
||||
|
||||
# 2) Session-id -> deployment affinity
|
||||
if self.enable_session_id_affinity:
|
||||
session_id = self._get_session_id_from_request_kwargs(
|
||||
request_kwargs=request_kwargs
|
||||
)
|
||||
if session_id is not None:
|
||||
session_cache_key = self.get_session_affinity_cache_key(
|
||||
model_group=stable_model_map_key, session_id=session_id
|
||||
)
|
||||
session_cache_result = await self.cache.async_get_cache(
|
||||
key=session_cache_key
|
||||
)
|
||||
|
||||
session_model_id: Optional[str] = None
|
||||
if isinstance(session_cache_result, dict):
|
||||
session_model_id = cast(
|
||||
Optional[str], session_cache_result.get("model_id")
|
||||
)
|
||||
elif isinstance(session_cache_result, str):
|
||||
session_model_id = session_cache_result
|
||||
|
||||
if session_model_id:
|
||||
session_deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments=typed_healthy_deployments,
|
||||
model_id=session_model_id,
|
||||
)
|
||||
if session_deployment is not None:
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: session-id affinity hit -> deployment=%s session_id=%s",
|
||||
session_model_id,
|
||||
session_id,
|
||||
)
|
||||
return [session_deployment]
|
||||
else:
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: session-id pinned deployment=%s not found in healthy_deployments",
|
||||
session_model_id,
|
||||
)
|
||||
|
||||
# 3) User key -> deployment affinity
|
||||
if not self.enable_user_key_affinity:
|
||||
return typed_healthy_deployments
|
||||
|
||||
|
|
@ -271,12 +351,6 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
if user_key is None:
|
||||
return typed_healthy_deployments
|
||||
|
||||
stable_model_map_key = self._get_stable_model_map_key_from_deployments(
|
||||
healthy_deployments=typed_healthy_deployments
|
||||
)
|
||||
if stable_model_map_key is None:
|
||||
return typed_healthy_deployments
|
||||
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
model_group=stable_model_map_key, user_key=user_key
|
||||
)
|
||||
|
|
@ -320,11 +394,18 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
- LiteLLM runs async success callbacks via a background logging worker for performance.
|
||||
- We want affinity to be immediately available for subsequent requests.
|
||||
"""
|
||||
if not self.enable_user_key_affinity:
|
||||
if not self.enable_user_key_affinity and not self.enable_session_id_affinity:
|
||||
return None
|
||||
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
if user_key is None:
|
||||
user_key = None
|
||||
if self.enable_user_key_affinity:
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
|
||||
session_id = None
|
||||
if self.enable_session_id_affinity:
|
||||
session_id = self._get_session_id_from_request_kwargs(request_kwargs=kwargs)
|
||||
|
||||
if user_key is None and session_id is None:
|
||||
return None
|
||||
|
||||
metadata_dicts = self._iter_metadata_dicts(kwargs)
|
||||
|
|
@ -357,7 +438,10 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
deployment_model_name: Optional[str] = None
|
||||
for metadata in metadata_dicts:
|
||||
maybe_deployment_model_name = metadata.get("deployment_model_name")
|
||||
if isinstance(maybe_deployment_model_name, str) and maybe_deployment_model_name:
|
||||
if (
|
||||
isinstance(maybe_deployment_model_name, str)
|
||||
and maybe_deployment_model_name
|
||||
):
|
||||
deployment_model_name = maybe_deployment_model_name
|
||||
break
|
||||
|
||||
|
|
@ -368,29 +452,55 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
model_group=deployment_model_name, user_key=user_key
|
||||
)
|
||||
await self.cache.async_set_cache(
|
||||
cache_key,
|
||||
DeploymentAffinityCacheValue(model_id=str(model_id)),
|
||||
ttl=self.ttl_seconds,
|
||||
)
|
||||
if user_key is not None:
|
||||
try:
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
model_group=deployment_model_name, user_key=user_key
|
||||
)
|
||||
await self.cache.async_set_cache(
|
||||
cache_key,
|
||||
DeploymentAffinityCacheValue(model_id=str(model_id)),
|
||||
ttl=self.ttl_seconds,
|
||||
)
|
||||
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: set affinity mapping model_map_key=%s deployment=%s ttl=%s user_key=%s",
|
||||
deployment_model_name,
|
||||
model_id,
|
||||
self.ttl_seconds,
|
||||
self._shorten_for_logs(user_key),
|
||||
)
|
||||
except Exception as e:
|
||||
# Non-blocking: affinity is a best-effort optimization.
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: failed to set affinity cache. model_map_key=%s error=%s",
|
||||
deployment_model_name,
|
||||
e,
|
||||
)
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: set affinity mapping model_map_key=%s deployment=%s ttl=%s user_key=%s",
|
||||
deployment_model_name,
|
||||
model_id,
|
||||
self.ttl_seconds,
|
||||
self._shorten_for_logs(user_key),
|
||||
)
|
||||
except Exception as e:
|
||||
# Non-blocking: affinity is a best-effort optimization.
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: failed to set user key affinity cache. model_map_key=%s error=%s",
|
||||
deployment_model_name,
|
||||
e,
|
||||
)
|
||||
|
||||
# Also persist Session-ID affinity if enabled and session-id is provided
|
||||
if session_id is not None:
|
||||
try:
|
||||
session_cache_key = self.get_session_affinity_cache_key(
|
||||
model_group=deployment_model_name, session_id=session_id
|
||||
)
|
||||
await self.cache.async_set_cache(
|
||||
session_cache_key,
|
||||
DeploymentAffinityCacheValue(model_id=str(model_id)),
|
||||
ttl=self.ttl_seconds,
|
||||
)
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: set session affinity mapping model_map_key=%s deployment=%s ttl=%s session_id=%s",
|
||||
deployment_model_name,
|
||||
model_id,
|
||||
self.ttl_seconds,
|
||||
session_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: failed to set session affinity cache. model_map_key=%s error=%s",
|
||||
deployment_model_name,
|
||||
e,
|
||||
)
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -803,6 +803,7 @@ OptionalPreCallChecks = List[
|
|||
"router_budget_limiting",
|
||||
"responses_api_deployment_check",
|
||||
"deployment_affinity",
|
||||
"session_affinity",
|
||||
"forward_client_headers_by_model_group",
|
||||
"enforce_model_rate_limits",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,179 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import json
|
||||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
|
||||
DeploymentAffinityCheck,
|
||||
)
|
||||
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, json_data, status_code):
|
||||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_session_id_affinity_routes_to_same_deployment():
|
||||
"""
|
||||
When session_affinity is enabled, subsequent requests from the same session id
|
||||
should route to the same deployment.
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-resp-123",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "azure/computer-use-preview",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_123",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "output_text", "text": "Hello there!", "annotations": []}
|
||||
],
|
||||
}
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 10,
|
||||
"total_tokens": 15,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
},
|
||||
"text": {"format": {"type": "text"}},
|
||||
"error": None,
|
||||
"previous_response_id": None,
|
||||
}
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "azure-computer-use-preview",
|
||||
"litellm_params": {
|
||||
"model": "azure/computer-use-preview-1",
|
||||
"api_key": "mock-api-key-1",
|
||||
"api_version": "mock-api-version",
|
||||
"api_base": "https://mock-endpoint-1.openai.azure.com",
|
||||
},
|
||||
"model_info": {"base_model": "computer-use-preview"},
|
||||
},
|
||||
{
|
||||
"model_name": "azure-computer-use-preview",
|
||||
"litellm_params": {
|
||||
"model": "azure/computer-use-preview-2",
|
||||
"api_key": "mock-api-key-2",
|
||||
"api_version": "mock-api-version-2",
|
||||
"api_base": "https://mock-endpoint-2.openai.azure.com",
|
||||
},
|
||||
"model_info": {"base_model": "computer-use-preview"},
|
||||
},
|
||||
],
|
||||
optional_pre_call_checks=["session_affinity"],
|
||||
)
|
||||
|
||||
model_group = "azure-computer-use-preview"
|
||||
session_id = "test-session-id-1"
|
||||
|
||||
choice_calls = {"count": 0}
|
||||
|
||||
def deterministic_choice(seq):
|
||||
choice_calls["count"] += 1
|
||||
if choice_calls["count"] == 1:
|
||||
return seq[0]
|
||||
return seq[1] if len(seq) > 1 else seq[0]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post, patch(
|
||||
"litellm.router_strategy.simple_shuffle.random.choice",
|
||||
side_effect=deterministic_choice,
|
||||
):
|
||||
mock_post.return_value = MockResponse(mock_response_data, 200)
|
||||
|
||||
first_response = await router.aresponses(
|
||||
model=model_group,
|
||||
input="Hello, how are you?",
|
||||
truncation="auto",
|
||||
litellm_metadata={"session_id": session_id},
|
||||
)
|
||||
first_model_id = first_response._hidden_params["model_id"]
|
||||
|
||||
second_response = await router.aresponses(
|
||||
model=model_group,
|
||||
input="Follow-up question",
|
||||
truncation="auto",
|
||||
litellm_metadata={"session_id": session_id},
|
||||
)
|
||||
assert second_response._hidden_params["model_id"] == first_model_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_session_id_affinity_priority_over_user_key():
|
||||
"""
|
||||
If both session_affinity and deployment_affinity are enabled,
|
||||
session_affinity should have priority. We test this by sending different
|
||||
session ids for the same user.
|
||||
"""
|
||||
cache = DualCache()
|
||||
callback = DeploymentAffinityCheck(
|
||||
cache=cache,
|
||||
ttl_seconds=123,
|
||||
enable_user_key_affinity=True,
|
||||
enable_responses_api_affinity=False,
|
||||
enable_session_id_affinity=True,
|
||||
)
|
||||
|
||||
healthy_deployments = [
|
||||
{
|
||||
"model_name": "model_group",
|
||||
"litellm_params": {"model": "model_1"},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
},
|
||||
{
|
||||
"model_name": "model_group",
|
||||
"litellm_params": {"model": "model_2"},
|
||||
"model_info": {"id": "deployment-2"},
|
||||
},
|
||||
]
|
||||
|
||||
await callback.cache.async_set_cache(
|
||||
DeploymentAffinityCheck.get_affinity_cache_key("model_group", "user1"),
|
||||
{"model_id": "deployment-1"},
|
||||
)
|
||||
|
||||
await callback.cache.async_set_cache(
|
||||
DeploymentAffinityCheck.get_session_affinity_cache_key(
|
||||
"model_group", "session1"
|
||||
),
|
||||
{"model_id": "deployment-2"},
|
||||
)
|
||||
|
||||
# Should use session mapping
|
||||
filtered = await callback.async_filter_deployments(
|
||||
model="model_group",
|
||||
healthy_deployments=healthy_deployments,
|
||||
messages=[],
|
||||
request_kwargs={
|
||||
"metadata": {"user_api_key_hash": "user1", "session_id": "session1"}
|
||||
},
|
||||
)
|
||||
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["model_info"]["id"] == "deployment-2"
|
||||
Loading…
Add table
Reference in a new issue