From 456d8f5524a469486dea770904c3baf8a61f3970 Mon Sep 17 00:00:00 2001 From: Harshit Jain Date: Sat, 21 Feb 2026 18:45:50 +0530 Subject: [PATCH] feat: add session_id to have better routing --- docs/my-website/docs/response_api.md | 6 +- litellm/proxy/litellm_pre_call_utils.py | 28 ++- litellm/router.py | 70 ++++--- .../deployment_affinity_check.py | 190 ++++++++++++++---- litellm/types/router.py | 1 + .../test_session_id_affinity.py | 179 +++++++++++++++++ 6 files changed, 393 insertions(+), 81 deletions(-) create mode 100644 tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py diff --git a/docs/my-website/docs/response_api.md b/docs/my-website/docs/response_api.md index 90b1beefa0f..b37be2b5bc2 100644 --- a/docs/my-website/docs/response_api.md +++ b/docs/my-website/docs/response_api.md @@ -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 diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index dc9710c2ce2..235bd85c5be 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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}" diff --git a/litellm/router.py b/litellm/router.py index 855f21387aa..01089f98d5d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index d34607732b8..8044f71d904 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -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 diff --git a/litellm/types/router.py b/litellm/types/router.py index 3abe9f202aa..00b7853cc4c 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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", ] diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py new file mode 100644 index 00000000000..f33f332a2dd --- /dev/null +++ b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py @@ -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"