mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge d7d5d8afdb into e55dbaf347
This commit is contained in:
commit
e4dda9d80f
5 changed files with 575 additions and 4 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
388
tests/test_litellm/router_strategy/test_session_key_fallback.py
Normal file
388
tests/test_litellm/router_strategy/test_session_key_fallback.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue