diff --git a/litellm/router_strategy/adaptive_router/hooks.py b/litellm/router_strategy/adaptive_router/hooks.py index b59ce6e3621..779eae48533 100644 --- a/litellm/router_strategy/adaptive_router/hooks.py +++ b/litellm/router_strategy/adaptive_router/hooks.py @@ -55,9 +55,13 @@ def _resolve_session_key(kwargs: dict[str, Any]) -> str | None: sid = litellm_params.get("litellm_session_id") if sid: return str(sid) - metadata: Final = litellm_params.get("metadata") or {} + metadata: Final = ( + litellm_params.get("metadata") + if isinstance(litellm_params.get("metadata"), dict) + else (kwargs.get("metadata") if isinstance(kwargs.get("metadata"), dict) else kwargs.get("litellm_metadata")) + ) or {} if isinstance(metadata, dict): - sid = metadata.get("session_id") or metadata.get("litellm_session_id") + sid = metadata.get("session_id") or metadata.get("litellm_session_id") or metadata.get("prompt_cache_key") if sid: return str(sid) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index f1f791ba72e..851513d4c98 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -16,6 +16,7 @@ Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter from __future__ import annotations import asyncio +import hashlib import random import re from collections.abc import Iterator, Mapping, Sequence @@ -2151,6 +2152,162 @@ class ComplexityRouter(CustomLogger): return str(session_id) return None + @staticmethod + def _get_prompt_cache_key_from_request_kwargs(request_kwargs: dict) -> str | None: + """Resolve a client-supplied prompt_cache_key from request_kwargs, metadata, litellm_params, extra_body, or headers.""" + val = request_kwargs.get("prompt_cache_key") + if val is not None and str(val).strip(): + return str(val).strip() + + for metadata in ComplexityRouter._iter_metadata_dicts(request_kwargs): + val = metadata.get("prompt_cache_key") + if val is not None and str(val).strip(): + return str(val).strip() + + litellm_params = request_kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + val = litellm_params.get("prompt_cache_key") + if val is not None and str(val).strip(): + return str(val).strip() + lp_meta = litellm_params.get("metadata") + if isinstance(lp_meta, dict): + val = lp_meta.get("prompt_cache_key") + if val is not None and str(val).strip(): + return str(val).strip() + + extra_body = request_kwargs.get("extra_body") + if isinstance(extra_body, dict): + val = extra_body.get("prompt_cache_key") + if val is not None and str(val).strip(): + return str(val).strip() + + headers = request_kwargs.get("headers") + if isinstance(headers, dict): + for hkey in ("prompt_cache_key", "prompt-cache-key", "x-prompt-cache-key", "X-Prompt-Cache-Key"): + val = headers.get(hkey) + if val is not None and str(val).strip(): + return str(val).strip() + + return None + + @staticmethod + def _get_user_identifier_from_request_kwargs(request_kwargs: dict) -> str: + """Extract user identifier (user_api_key_hash, user_api_key, user_id, or user).""" + for metadata in ComplexityRouter._iter_metadata_dicts(request_kwargs): + for key in ("user_api_key_hash", "user_api_key", "user_api_key_user_id", "user_id", "user"): + val = metadata.get(key) + if val is not None and str(val).strip(): + return str(val).strip() + user = request_kwargs.get("user") + if user is not None and str(user).strip(): + return str(user).strip() + return "" + + @classmethod + def _extract_message_text(cls, content: Any) -> str: + """Extract text from message content (str or structured parts).""" + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for part in content: + if isinstance(part, str): + parts.append(part) + elif isinstance(part, dict): + if part.get("type") == "text" and "text" in part: + text_val = part["text"] + if isinstance(text_val, str): + parts.append(text_val) + elif "content" in part and isinstance(part["content"], str): + parts.append(part["content"]) + return "".join(parts) + if isinstance(content, dict): + if content.get("type") == "text" and "text" in content: + text_val = content["text"] + if isinstance(text_val, str): + return text_val + elif "content" in content and isinstance(content["content"], str): + return content["content"] + return "" + if content is None: + return "" + return str(content) + + def _derive_prefix_hash( + self, + request_kwargs: dict, + resolved_messages: list[dict[str, Any]] | None = None, + model: str | None = None, + ) -> str | None: + """ + Derive a session key hash via SHA256 of: + (user_api_key / user_id) + (model / model_group) + (normalized first system message) + (normalized first user message) + """ + messages = resolved_messages or request_kwargs.get("messages") or [] + first_system_msg = "" + first_user_msg = "" + + if isinstance(messages, list): + for msg in messages: + if not isinstance(msg, dict): + continue + role = msg.get("role") + if not first_system_msg and role in ("system", "developer"): + first_system_msg = self._extract_message_text(msg.get("content"))[:1024].strip() + elif not first_user_msg and role == "user": + first_user_msg = self._extract_message_text(msg.get("content"))[:1024].strip() + if first_system_msg and first_user_msg: + break + + if not first_system_msg and not first_user_msg: + return None + + user_identifier = self._get_user_identifier_from_request_kwargs(request_kwargs) + model_group = model or self.model_name or request_kwargs.get("model") or "" + + payload = f"litellm-session-key:{user_identifier}:{model_group}:{first_system_msg}:{first_user_msg}" + return hashlib.sha256(payload.encode("utf-8", errors="replace")).hexdigest() + + def _resolve_session_id( + self, + request_kwargs: dict, + resolved_messages: list[dict[str, Any]] | None = None, + model: str | None = None, + ) -> str | None: + """Resolve a session_id: client-supplied first, then session_key_fallback if configured.""" + session_id = self._get_session_id_from_request_kwargs(request_kwargs) + if session_id is not None: + return session_id + + fallback = self.config.session_key_fallback + if fallback == "none" or not fallback: + return None + + derived_key: str | None = None + if fallback == "prompt_cache_key": + derived_key = self._get_prompt_cache_key_from_request_kwargs(request_kwargs) + elif fallback == "prefix_hash": + derived_key = self._derive_prefix_hash( + request_kwargs=request_kwargs, + resolved_messages=resolved_messages, + model=model, + ) + + if derived_key is not None: + sanitized_key = derived_key.replace("\r", "").replace("\n", "")[:32] + verbose_router_logger.info( + "ComplexityRouter: resolved fallback session key '%s' using strategy '%s'", + sanitized_key, + fallback, + ) + metadata_key = "litellm_metadata" if "litellm_metadata" in request_kwargs else "metadata" + metadata = request_kwargs.setdefault(metadata_key, {}) + if isinstance(metadata, dict) and "session_id" not in metadata: + metadata["session_id"] = derived_key + return derived_key + + return None + @staticmethod def _get_user_api_key_hash_from_request_kwargs(request_kwargs: dict) -> str | None: """Resolve the proxy-derived API key hash, the same trust boundary @@ -2227,8 +2384,21 @@ class ComplexityRouter(CustomLogger): conversation_continuing: Final = _conversation_is_continuing(resolved_messages) use_session_affinity: Final = self._uses_tier_pin - session_id: Final = self._get_session_id_from_request_kwargs(request_kwargs) if use_session_affinity else None - cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None + use_deployment_affinity: Final = self._uses_deployment_pin + session_id: Final = ( + self._resolve_session_id( + request_kwargs=request_kwargs, + resolved_messages=resolved_messages, + model=model, + ) + if (use_session_affinity or use_deployment_affinity or self.config.adaptive) + else None + ) + cache_key = ( + self._get_session_affinity_cache_key(session_id, request_kwargs) + if (use_session_affinity and session_id is not None) + else None + ) if cache_key is not None: pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key) diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 9b2a25f5d28..adea84999a5 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -829,6 +829,14 @@ class ComplexityRouterConfig(BaseModel): "idle time for the session's routing decisions rather than total session length" ), ) + session_key_fallback: Literal["none", "prompt_cache_key", "prefix_hash"] = Field( + default="none", + description=( + "Fallback method to derive a session key when session_id is absent from request metadata. " + "Supported values: 'none' (default, no fallback), 'prompt_cache_key' (use prompt_cache_key from request), " + "or 'prefix_hash' (SHA256 hash of user_api_key/user_id + model/model_group + normalized first system message + normalized first user message)." + ), + ) plugins: list[RoutingPlugin] | None = Field( default=None, diff --git a/litellm/types/router.py b/litellm/types/router.py index 97bd93f3f47..405e64e9fb5 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -1029,6 +1029,7 @@ class AdaptiveRouterWeights(BaseModel): class AdaptiveRouterConfig(BaseModel): available_models: list[str] weights: AdaptiveRouterWeights = Field(default_factory=AdaptiveRouterWeights) + session_key_fallback: Literal["none", "prompt_cache_key", "prefix_hash"] = "none" class AdaptiveRouterPreferences(BaseModel): diff --git a/tests/test_litellm/router_strategy/test_session_key_fallback.py b/tests/test_litellm/router_strategy/test_session_key_fallback.py new file mode 100644 index 00000000000..5e336e86a2a --- /dev/null +++ b/tests/test_litellm/router_strategy/test_session_key_fallback.py @@ -0,0 +1,388 @@ +""" +Unit tests for Auto Router v2 session_key_fallback derivation (Issue #34766). + +Tests cover: +- Session resolution with explicit session_id. +- Fallback to prompt_cache_key when configured. +- Fallback to prefix_hash when configured. +- Normal operation with "none" (default behavior unchanged). +- Deployment affinity with session_key_fallback. +- Adaptive router session key resolution with fallback metadata. +""" + +import hashlib +from unittest.mock import AsyncMock + +import pytest + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.router_strategy.adaptive_router.hooks import _resolve_session_key +from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter +from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig + + +@pytest.fixture +def mock_router_instance(): + class _MockRouter: + def __init__(self): + self.cache = DualCache() + self.model_list = [ + { + "model_name": "gpt-4o-mini", + "litellm_params": {"model": "openai/gpt-4o-mini", "input_cost_per_token": 0.0}, + "model_info": {}, + }, + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "input_cost_per_token": 0.0}, + "model_info": {}, + }, + { + "model_name": "claude-sonnet-4-20250514", + "litellm_params": {"model": "anthropic/claude-sonnet-4-20250514", "input_cost_per_token": 0.0}, + "model_info": {}, + }, + { + "model_name": "o1-preview", + "litellm_params": {"model": "openai/o1-preview", "input_cost_per_token": 0.0}, + "model_info": {}, + }, + ] + self.model_name_to_deployment_indices = { + "gpt-4o-mini": [0], + "gpt-4o": [1], + "claude-sonnet-4-20250514": [2], + "o1-preview": [3], + } + + return _MockRouter() + + +@pytest.fixture +def base_config(): + return { + "tiers": { + "SIMPLE": "gpt-4o-mini", + "MEDIUM": "gpt-4o", + "COMPLEX": "claude-sonnet-4-20250514", + "REASONING": "o1-preview", + }, + "tier_boundaries": { + "simple_medium": 0.25, + "medium_complex": 0.50, + "complex_reasoning": 0.75, + }, + "session_affinity": True, + "session_affinity_ttl_seconds": 3600, + } + + +class TestSessionKeyFallback: + SIMPLE_MESSAGE = [{"role": "user", "content": "Hello!"}] + REASONING_MESSAGE = [ + {"role": "system", "content": "You are a mathematics tutor."}, + { + "role": "user", + "content": "Let's think step by step and reason through this problem carefully.", + }, + ] + + def test_config_default_fallback_is_none(self): + cfg = ComplexityRouterConfig(tiers={"SIMPLE": "gpt-4o-mini"}) + assert cfg.session_key_fallback == "none" + + def test_config_supports_valid_fallback_modes(self): + cfg_cache = ComplexityRouterConfig(tiers={"SIMPLE": "gpt-4o-mini"}, session_key_fallback="prompt_cache_key") + assert cfg_cache.session_key_fallback == "prompt_cache_key" + + cfg_prefix = ComplexityRouterConfig(tiers={"SIMPLE": "gpt-4o-mini"}, session_key_fallback="prefix_hash") + assert cfg_prefix.session_key_fallback == "prefix_hash" + + cfg_none = ComplexityRouterConfig(tiers={"SIMPLE": "gpt-4o-mini"}, session_key_fallback="none") + assert cfg_none.session_key_fallback == "none" + + @pytest.mark.asyncio + async def test_explicit_session_id_takes_precedence_over_fallback(self, mock_router_instance, base_config): + """When an explicit session_id is provided, it must be used directly, ignoring session_key_fallback.""" + config = {**base_config, "session_key_fallback": "prefix_hash"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + request_kwargs = { + "metadata": {"session_id": "my-explicit-session", "user_api_key_hash": "key123"}, + "prompt_cache_key": "cache-key-should-be-ignored", + } + + resolved_id = router._resolve_session_id( + request_kwargs=request_kwargs, + resolved_messages=self.REASONING_MESSAGE, + model="test-router", + ) + assert resolved_id == "my-explicit-session" + + # Pre routing hook pins model under explicit session id + res1 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=self.REASONING_MESSAGE, + ) + assert res1 is not None + assert res1.model == "o1-preview" + + # Turn 2: simple message under same explicit session_id should hit cache pin + req2_kwargs = {"metadata": {"session_id": "my-explicit-session", "user_api_key_hash": "key123"}} + res2 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=req2_kwargs, + messages=self.SIMPLE_MESSAGE, + ) + assert res2 is not None + assert res2.model == "o1-preview" + assert res2.routing_decision["cause"] == "session_affinity_pin" + + @pytest.mark.asyncio + async def test_fallback_to_prompt_cache_key(self, mock_router_instance, base_config): + """When session_id is absent and session_key_fallback='prompt_cache_key', use prompt_cache_key.""" + config = {**base_config, "session_key_fallback": "prompt_cache_key"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + request_kwargs = { + "prompt_cache_key": "custom-prompt-cache-key-999", + "metadata": {"user_api_key_hash": "key123"}, + } + + resolved_id = router._resolve_session_id( + request_kwargs=request_kwargs, + resolved_messages=self.REASONING_MESSAGE, + model="test-router", + ) + assert resolved_id == "custom-prompt-cache-key-999" + assert request_kwargs["metadata"]["session_id"] == "custom-prompt-cache-key-999" + + # Pre routing hook turn 1 + res1 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=self.REASONING_MESSAGE, + ) + assert res1 is not None + assert res1.model == "o1-preview" + + # Pre routing hook turn 2 with same prompt_cache_key + req2_kwargs = { + "prompt_cache_key": "custom-prompt-cache-key-999", + "metadata": {"user_api_key_hash": "key123"}, + } + res2 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=req2_kwargs, + messages=self.SIMPLE_MESSAGE, + ) + assert res2 is not None + assert res2.model == "o1-preview" + assert res2.routing_decision["cause"] == "session_affinity_pin" + + @pytest.mark.asyncio + async def test_fallback_prompt_cache_key_from_extra_body_or_headers(self, mock_router_instance, base_config): + """prompt_cache_key in extra_body or headers is also resolved.""" + config = {**base_config, "session_key_fallback": "prompt_cache_key"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + req_extra_body = {"extra_body": {"prompt_cache_key": "extra-body-key"}} + resolved = router._resolve_session_id(req_extra_body) + assert resolved == "extra-body-key" + + req_headers = {"headers": {"x-prompt-cache-key": "header-key"}} + resolved_hdr = router._resolve_session_id(req_headers) + assert resolved_hdr == "header-key" + + @pytest.mark.asyncio + async def test_fallback_prompt_cache_key_missing_returns_none(self, mock_router_instance, base_config): + """When prompt_cache_key is absent and fallback='prompt_cache_key', returns None (reclassifies).""" + config = {**base_config, "session_key_fallback": "prompt_cache_key"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + request_kwargs = {"metadata": {}} + resolved = router._resolve_session_id(request_kwargs, resolved_messages=self.SIMPLE_MESSAGE) + assert resolved is None + + @pytest.mark.asyncio + async def test_fallback_to_prefix_hash(self, mock_router_instance, base_config): + """When session_id is absent and session_key_fallback='prefix_hash', derive SHA256 prefix hash.""" + config = {**base_config, "session_key_fallback": "prefix_hash"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + request_kwargs = { + "metadata": {"user_api_key_hash": "user_hash_abc"}, + } + + # Expected SHA256 computation: + # user_api_key_hash : model_group : first_system_msg : first_user_msg + system_text = "You are a mathematics tutor." + user_text = "Let's think step by step and reason through this problem carefully." + expected_payload = f"litellm-session-key:user_hash_abc:test-router:{system_text}:{user_text}" + expected_hash = hashlib.sha256(expected_payload.encode("utf-8")).hexdigest() + + derived_id = router._resolve_session_id( + request_kwargs=request_kwargs, + resolved_messages=self.REASONING_MESSAGE, + model="test-router", + ) + assert derived_id == expected_hash + assert request_kwargs["metadata"]["session_id"] == expected_hash + + # Turn 1: Classifies as REASONING (o1-preview) and pins it under derived prefix hash + res1 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=self.REASONING_MESSAGE, + ) + assert res1 is not None + assert res1.model == "o1-preview" + + # Turn 2: Follow-up message in the same multi-turn conversation + turn2_messages = [ + {"role": "system", "content": "You are a mathematics tutor."}, + { + "role": "user", + "content": "Let's think step by step and reason through this problem carefully.", + }, + {"role": "assistant", "content": "Here is step 1..."}, + {"role": "user", "content": "Thanks! Can you summarize step 1 in one line?"}, + ] + turn2_kwargs = { + "metadata": {"user_api_key_hash": "user_hash_abc"}, + } + + turn2_derived_id = router._resolve_session_id( + request_kwargs=turn2_kwargs, + resolved_messages=turn2_messages, + model="test-router", + ) + # Prefix hash MUST match Turn 1 because initial system & user message are identical + assert turn2_derived_id == expected_hash + + res2 = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=turn2_kwargs, + messages=turn2_messages, + ) + assert res2 is not None + assert res2.model == "o1-preview" + assert res2.routing_decision["cause"] == "session_affinity_pin" + + @pytest.mark.asyncio + async def test_fallback_prefix_hash_different_conversations_segregate(self, mock_router_instance, base_config): + """Different initial prompts or different users produce different prefix hashes and route independently.""" + config = {**base_config, "session_key_fallback": "prefix_hash"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + conv_a_kwargs = {"metadata": {"user_api_key_hash": "user_a"}} + conv_b_kwargs = {"metadata": {"user_api_key_hash": "user_b"}} + + res_a = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=conv_a_kwargs, + messages=self.REASONING_MESSAGE, + ) + res_b = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=conv_b_kwargs, + messages=self.SIMPLE_MESSAGE, + ) + + assert res_a.model == "o1-preview" + assert res_b.model == "gpt-4o-mini" + assert conv_a_kwargs["metadata"]["session_id"] != conv_b_kwargs["metadata"]["session_id"] + + @pytest.mark.asyncio + async def test_fallback_prefix_hash_empty_messages_returns_none(self, mock_router_instance, base_config): + """When no message content is present, prefix_hash cannot be derived and returns None.""" + config = {**base_config, "session_key_fallback": "prefix_hash"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + resolved = router._resolve_session_id(request_kwargs={}, resolved_messages=[]) + assert resolved is None + + @pytest.mark.asyncio + async def test_default_none_behavior_unchanged(self, mock_router_instance, base_config): + """When session_key_fallback is 'none' (default), no fallback key is derived.""" + config = {**base_config, "session_key_fallback": "none"} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + cache = AsyncMock() + mock_router_instance.cache = cache + + # Without session_id, no cache lookup or cache set happens + res = await router.async_pre_routing_hook( + model="test-router", + request_kwargs={}, + messages=self.SIMPLE_MESSAGE, + ) + assert res.model == "gpt-4o-mini" + cache.async_get_cache.assert_not_called() + cache.async_set_cache.assert_not_called() + + @pytest.mark.asyncio + async def test_deployment_affinity_uses_fallback_session_id(self, mock_router_instance, base_config): + """When deployment_affinity is active and fallback derives session_id, deployment pin TTL is added.""" + config = { + **base_config, + "session_affinity": False, + "deployment_affinity": True, + "session_key_fallback": "prompt_cache_key", + } + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + + request_kwargs = {"prompt_cache_key": "cache-dep-123"} + res = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=self.SIMPLE_MESSAGE, + ) + assert res is not None + assert res.session_affinity_ttl_seconds == 3600 + assert request_kwargs["metadata"]["session_id"] == "cache-dep-123" + + def test_adaptive_router_hooks_resolve_session_key_with_fallback_metadata(self): + """Adaptive router post call hook picks up session_id populated in metadata by fallback.""" + kwargs = { + "metadata": {"session_id": "fallback-derived-session-id-123"}, + } + key = _resolve_session_key(kwargs) + assert key == "fallback-derived-session-id-123"