Merge pull request #21763 from Harshit28j/litellm_feat_sticky_sessions

feat: add session_id to have better routing
This commit is contained in:
Harshit Jain 2026-02-21 21:21:52 +05:30 committed by GitHub
commit e463fc22d9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 393 additions and 81 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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",
]

View file

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