This commit is contained in:
Roshan Sah 2026-08-28 01:12:14 -04:00 committed by GitHub
commit e4dda9d80f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 575 additions and 4 deletions

View file

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

View file

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

View file

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

View file

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

View file

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