mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): add per-model-group deployment affinity configuration
Enable deployment_affinity, responses_api_deployment_check, and session_affinity to be configured per model group via router_settings.model_group_affinity_config, falling back to global settings for unconfigured groups. - Add model_group_affinity_config parameter to Router and DeploymentAffinityCheck - Add _get_effective_flags helper to resolve flags per model group - Update async_filter_deployments and async_pre_call_deployment_hook to use per-group config - Add 4 comprehensive tests covering per-group config, fallback, and override scenarios This allows fine-grained control of affinity behavior across model groups, e.g., enabling stickiness only for cross-provider deployments while leaving other groups free to load-balance. Co-Authored-By: Claude Haiku 4.5 <noreply@anthropic.com>
This commit is contained in:
parent
e5baa2232f
commit
528daa8cf4
4 changed files with 371 additions and 28 deletions
|
|
@ -301,6 +301,7 @@ class Router:
|
|||
RouterGeneralSettings
|
||||
] = RouterGeneralSettings(),
|
||||
deployment_affinity_ttl_seconds: int = 3600,
|
||||
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
|
||||
ignore_invalid_deployments: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -641,6 +642,9 @@ class Router:
|
|||
self.model_group_retry_policy: Optional[
|
||||
Dict[str, RetryPolicy]
|
||||
] = model_group_retry_policy
|
||||
self.model_group_affinity_config: Optional[
|
||||
Dict[str, List[str]]
|
||||
] = model_group_affinity_config
|
||||
|
||||
self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None
|
||||
if allowed_fails_policy is not None:
|
||||
|
|
@ -661,6 +665,26 @@ class Router:
|
|||
if optional_pre_call_checks is not None:
|
||||
self.add_optional_pre_call_checks(optional_pre_call_checks)
|
||||
|
||||
# If model_group_affinity_config is set but no global affinity checks were
|
||||
# enabled, we still need the DeploymentAffinityCheck callback (with global
|
||||
# flags all False) so per-group config can activate affinity per model group.
|
||||
if self.model_group_affinity_config and not any(
|
||||
isinstance(cb, DeploymentAffinityCheck)
|
||||
for cb in (self.optional_callbacks or [])
|
||||
):
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
affinity_callback = DeploymentAffinityCheck(
|
||||
cache=self.cache,
|
||||
ttl_seconds=self.deployment_affinity_ttl_seconds,
|
||||
enable_user_key_affinity=False,
|
||||
enable_responses_api_affinity=False,
|
||||
enable_session_id_affinity=False,
|
||||
model_group_affinity_config=self.model_group_affinity_config,
|
||||
)
|
||||
self.optional_callbacks.append(affinity_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)
|
||||
|
||||
if self.alerting_config is not None:
|
||||
self._initialize_alerting()
|
||||
|
||||
|
|
@ -1311,6 +1335,10 @@ class Router:
|
|||
existing_affinity_callback.ttl_seconds = (
|
||||
self.deployment_affinity_ttl_seconds
|
||||
)
|
||||
if self.model_group_affinity_config:
|
||||
existing_affinity_callback.model_group_affinity_config = (
|
||||
self.model_group_affinity_config
|
||||
)
|
||||
else:
|
||||
affinity_callback = DeploymentAffinityCheck(
|
||||
cache=self.cache,
|
||||
|
|
@ -1318,6 +1346,7 @@ class Router:
|
|||
enable_user_key_affinity=enable_user_key_affinity,
|
||||
enable_responses_api_affinity=enable_responses_api_affinity,
|
||||
enable_session_id_affinity=enable_session_id_affinity,
|
||||
model_group_affinity_config=self.model_group_affinity_config,
|
||||
)
|
||||
self.optional_callbacks.append(affinity_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ where routing to a consistent deployment is still beneficial.
|
|||
"""
|
||||
|
||||
import hashlib
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
from typing import Any, Dict, List, Optional, Tuple, cast
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
|
@ -46,6 +46,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
enable_user_key_affinity: bool,
|
||||
enable_responses_api_affinity: bool,
|
||||
enable_session_id_affinity: bool = False,
|
||||
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.cache = cache
|
||||
|
|
@ -53,6 +54,32 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
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
|
||||
self.model_group_affinity_config: Dict[str, List[str]] = (
|
||||
model_group_affinity_config or {}
|
||||
)
|
||||
|
||||
def _get_effective_flags(
|
||||
self, model_group: str
|
||||
) -> Tuple[bool, bool, bool]:
|
||||
"""
|
||||
Return (enable_user_key_affinity, enable_responses_api_affinity, enable_session_id_affinity)
|
||||
for the given model group.
|
||||
|
||||
If the model group has an explicit entry in model_group_affinity_config, use it.
|
||||
Otherwise fall back to the global instance flags.
|
||||
"""
|
||||
group_checks = self.model_group_affinity_config.get(model_group)
|
||||
if group_checks is not None:
|
||||
return (
|
||||
"deployment_affinity" in group_checks,
|
||||
"responses_api_deployment_check" in group_checks,
|
||||
"session_affinity" in group_checks,
|
||||
)
|
||||
return (
|
||||
self.enable_user_key_affinity,
|
||||
self.enable_responses_api_affinity,
|
||||
self.enable_session_id_affinity,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _looks_like_sha256_hex(value: str) -> bool:
|
||||
|
|
@ -277,8 +304,12 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
request_kwargs = request_kwargs or {}
|
||||
typed_healthy_deployments = cast(List[dict], healthy_deployments)
|
||||
|
||||
enable_user_key, enable_responses_api, enable_session_id = (
|
||||
self._get_effective_flags(model)
|
||||
)
|
||||
|
||||
# 1) Responses API continuity (high priority)
|
||||
if self.enable_responses_api_affinity:
|
||||
if enable_responses_api:
|
||||
previous_response_id = request_kwargs.get("previous_response_id")
|
||||
if previous_response_id is not None:
|
||||
responses_model_id = (
|
||||
|
|
@ -305,7 +336,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return typed_healthy_deployments
|
||||
|
||||
# 2) Session-id -> deployment affinity
|
||||
if self.enable_session_id_affinity:
|
||||
if enable_session_id:
|
||||
session_id = self._get_session_id_from_request_kwargs(
|
||||
request_kwargs=request_kwargs
|
||||
)
|
||||
|
|
@ -344,7 +375,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
)
|
||||
|
||||
# 3) User key -> deployment affinity
|
||||
if not self.enable_user_key_affinity:
|
||||
if not enable_user_key:
|
||||
return typed_healthy_deployments
|
||||
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs)
|
||||
|
|
@ -394,22 +425,42 @@ 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 and not self.enable_session_id_affinity:
|
||||
metadata_dicts = self._iter_metadata_dicts(kwargs)
|
||||
|
||||
# Extract deployment_model_name first — needed for both per-group flag resolution
|
||||
# and cache key scoping.
|
||||
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
|
||||
):
|
||||
deployment_model_name = maybe_deployment_model_name
|
||||
break
|
||||
|
||||
if not deployment_model_name:
|
||||
return None
|
||||
|
||||
# Resolve effective flags for this model group
|
||||
enable_user_key, _enable_responses_api, enable_session_id = (
|
||||
self._get_effective_flags(deployment_model_name)
|
||||
)
|
||||
|
||||
if not enable_user_key and not enable_session_id:
|
||||
return None
|
||||
|
||||
user_key = None
|
||||
if self.enable_user_key_affinity:
|
||||
if enable_user_key:
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
|
||||
session_id = None
|
||||
if self.enable_session_id_affinity:
|
||||
if enable_session_id:
|
||||
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)
|
||||
|
||||
model_info = kwargs.get("model_info")
|
||||
if not isinstance(model_info, dict):
|
||||
model_info = None
|
||||
|
|
@ -433,25 +484,6 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
)
|
||||
return None
|
||||
|
||||
# Scope affinity by the Router deployment model name (alias-safe, consistent across
|
||||
# heterogeneous providers, and matches standard logging's `model_map_key`).
|
||||
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
|
||||
):
|
||||
deployment_model_name = maybe_deployment_model_name
|
||||
break
|
||||
|
||||
if not deployment_model_name:
|
||||
verbose_router_logger.warning(
|
||||
"DeploymentAffinityCheck: deployment_model_name missing; skipping affinity cache update. model_id=%s",
|
||||
model_id,
|
||||
)
|
||||
return None
|
||||
|
||||
if user_key is not None:
|
||||
try:
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ class UpdateRouterConfig(BaseModel):
|
|||
routing_strategy_args: Optional[dict] = None
|
||||
routing_strategy: Optional[str] = None
|
||||
model_group_retry_policy: Optional[dict] = None
|
||||
model_group_affinity_config: Optional[Dict[str, List[str]]] = None
|
||||
allowed_fails: Optional[int] = None
|
||||
cooldown_time: Optional[float] = None
|
||||
num_retries: Optional[int] = None
|
||||
|
|
|
|||
|
|
@ -657,3 +657,284 @@ def test_cache_key_does_not_double_hash_user_api_key_hash():
|
|||
user_key=user_api_key_hash,
|
||||
)
|
||||
assert key.endswith(user_api_key_hash)
|
||||
|
||||
|
||||
def test_get_effective_flags_returns_per_group_config():
|
||||
"""
|
||||
_get_effective_flags should return per-group flags when the model group has an entry
|
||||
in model_group_affinity_config, and global flags otherwise.
|
||||
"""
|
||||
callback = DeploymentAffinityCheck(
|
||||
cache=AsyncMock(),
|
||||
ttl_seconds=60,
|
||||
enable_user_key_affinity=True,
|
||||
enable_responses_api_affinity=True,
|
||||
enable_session_id_affinity=False,
|
||||
model_group_affinity_config={
|
||||
"gpt-4": ["deployment_affinity"],
|
||||
"claude-3": ["session_affinity", "responses_api_deployment_check"],
|
||||
},
|
||||
)
|
||||
|
||||
# gpt-4: only deployment_affinity
|
||||
user_key, responses_api, session_id = callback._get_effective_flags("gpt-4")
|
||||
assert user_key is True
|
||||
assert responses_api is False
|
||||
assert session_id is False
|
||||
|
||||
# claude-3: session_affinity + responses_api_deployment_check
|
||||
user_key, responses_api, session_id = callback._get_effective_flags("claude-3")
|
||||
assert user_key is False
|
||||
assert responses_api is True
|
||||
assert session_id is True
|
||||
|
||||
# unconfigured-model: falls back to global flags
|
||||
user_key, responses_api, session_id = callback._get_effective_flags(
|
||||
"unconfigured-model"
|
||||
)
|
||||
assert user_key is True
|
||||
assert responses_api is True
|
||||
assert session_id is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_affinity_config_only_applies_to_configured_group():
|
||||
"""
|
||||
When model_group_affinity_config is set without global optional_pre_call_checks,
|
||||
only configured model groups should get affinity behavior.
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-resp-per-group",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "openai/gpt-4",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_pg",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Per-group response"}],
|
||||
}
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 5, "total_tokens": 10},
|
||||
"text": {"format": {"type": "text"}},
|
||||
"error": None,
|
||||
"previous_response_id": None,
|
||||
}
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4-deploy-1",
|
||||
"api_key": "mock-key-1",
|
||||
"api_base": "https://mock-gpt4-1.openai.azure.com",
|
||||
"api_version": "2024-02-01",
|
||||
},
|
||||
"model_info": {"base_model": "gpt-4"},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4-deploy-2",
|
||||
"api_key": "mock-key-2",
|
||||
"api_base": "https://mock-gpt4-2.openai.azure.com",
|
||||
"api_version": "2024-02-01",
|
||||
},
|
||||
"model_info": {"base_model": "gpt-4"},
|
||||
},
|
||||
{
|
||||
"model_name": "claude-3",
|
||||
"litellm_params": {
|
||||
"model": "azure/claude-3-deploy-1",
|
||||
"api_key": "mock-key-3",
|
||||
"api_base": "https://mock-claude-1.openai.azure.com",
|
||||
"api_version": "2024-02-01",
|
||||
},
|
||||
"model_info": {"base_model": "claude-3"},
|
||||
},
|
||||
{
|
||||
"model_name": "claude-3",
|
||||
"litellm_params": {
|
||||
"model": "azure/claude-3-deploy-2",
|
||||
"api_key": "mock-key-4",
|
||||
"api_base": "https://mock-claude-2.openai.azure.com",
|
||||
"api_version": "2024-02-01",
|
||||
},
|
||||
"model_info": {"base_model": "claude-3"},
|
||||
},
|
||||
],
|
||||
# No global optional_pre_call_checks — only per-group
|
||||
model_group_affinity_config={
|
||||
"gpt-4": ["deployment_affinity"],
|
||||
},
|
||||
)
|
||||
|
||||
user_api_key_hash = "test-per-group-key"
|
||||
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)
|
||||
|
||||
# gpt-4: affinity should work — second request pinned to same deployment
|
||||
first = await router.aresponses(
|
||||
model="gpt-4",
|
||||
input="Hello",
|
||||
truncation="auto",
|
||||
litellm_metadata={"user_api_key_hash": user_api_key_hash},
|
||||
)
|
||||
first_model_id = first._hidden_params["model_id"]
|
||||
|
||||
second = await router.aresponses(
|
||||
model="gpt-4",
|
||||
input="Follow-up",
|
||||
truncation="auto",
|
||||
litellm_metadata={"user_api_key_hash": user_api_key_hash},
|
||||
)
|
||||
assert second._hidden_params["model_id"] == first_model_id
|
||||
|
||||
# claude-3: no affinity configured — should NOT be pinned
|
||||
choice_calls["count"] = 0
|
||||
first_claude = await router.aresponses(
|
||||
model="claude-3",
|
||||
input="Hello",
|
||||
truncation="auto",
|
||||
litellm_metadata={"user_api_key_hash": user_api_key_hash},
|
||||
)
|
||||
first_claude_id = first_claude._hidden_params["model_id"]
|
||||
|
||||
second_claude = await router.aresponses(
|
||||
model="claude-3",
|
||||
input="Follow-up",
|
||||
truncation="auto",
|
||||
litellm_metadata={"user_api_key_hash": user_api_key_hash},
|
||||
)
|
||||
# With deterministic choice and len>1, second call picks seq[1]
|
||||
assert second_claude._hidden_params["model_id"] != first_claude_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_affinity_config_falls_back_to_global():
|
||||
"""
|
||||
When both global optional_pre_call_checks and model_group_affinity_config are set,
|
||||
unconfigured model groups should use the global settings.
|
||||
"""
|
||||
callback = DeploymentAffinityCheck(
|
||||
cache=DualCache(),
|
||||
ttl_seconds=60,
|
||||
enable_user_key_affinity=True,
|
||||
enable_responses_api_affinity=False,
|
||||
enable_session_id_affinity=False,
|
||||
model_group_affinity_config={
|
||||
"claude-3": ["session_affinity"],
|
||||
},
|
||||
)
|
||||
|
||||
stable_model_map_key = "gpt-4"
|
||||
user_key = "test-fallback-key"
|
||||
|
||||
healthy_deployments = [
|
||||
{
|
||||
"model_name": stable_model_map_key,
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
},
|
||||
{
|
||||
"model_name": stable_model_map_key,
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "deployment-2"},
|
||||
},
|
||||
]
|
||||
|
||||
# Set up affinity cache for gpt-4 (should work since global has deployment_affinity)
|
||||
await callback.async_pre_call_deployment_hook(
|
||||
kwargs={
|
||||
"model_info": {"id": "deployment-1"},
|
||||
"metadata": {
|
||||
"user_api_key_hash": user_key,
|
||||
"deployment_model_name": stable_model_map_key,
|
||||
},
|
||||
},
|
||||
call_type=None,
|
||||
)
|
||||
|
||||
# gpt-4 not in model_group_affinity_config, so global flags apply (user_key affinity ON)
|
||||
filtered = await callback.async_filter_deployments(
|
||||
model="gpt-4",
|
||||
healthy_deployments=healthy_deployments,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"user_api_key_hash": user_key}},
|
||||
parent_otel_span=None,
|
||||
)
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["model_info"]["id"] == "deployment-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_affinity_config_overrides_global():
|
||||
"""
|
||||
When model_group_affinity_config specifies session_affinity for a model group,
|
||||
user-key affinity (from global config) should NOT apply to that group.
|
||||
"""
|
||||
callback = DeploymentAffinityCheck(
|
||||
cache=DualCache(),
|
||||
ttl_seconds=60,
|
||||
enable_user_key_affinity=True,
|
||||
enable_responses_api_affinity=False,
|
||||
enable_session_id_affinity=False,
|
||||
model_group_affinity_config={
|
||||
"claude-3": ["session_affinity"],
|
||||
},
|
||||
)
|
||||
|
||||
stable_model_map_key = "claude-3"
|
||||
user_key = "test-override-key"
|
||||
|
||||
healthy_deployments = [
|
||||
{
|
||||
"model_name": stable_model_map_key,
|
||||
"litellm_params": {"model": "anthropic/claude-3-opus"},
|
||||
"model_info": {"id": "deployment-1"},
|
||||
},
|
||||
{
|
||||
"model_name": stable_model_map_key,
|
||||
"litellm_params": {"model": "anthropic/claude-3-opus"},
|
||||
"model_info": {"id": "deployment-2"},
|
||||
},
|
||||
]
|
||||
|
||||
# Set up user-key affinity cache for claude-3
|
||||
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(
|
||||
model_group=stable_model_map_key, user_key=user_key
|
||||
)
|
||||
await callback.cache.async_set_cache(
|
||||
cache_key, {"model_id": "deployment-1"}, ttl=60
|
||||
)
|
||||
|
||||
# claude-3 has per-group config (session_affinity only), so user-key affinity
|
||||
# should NOT apply even though it's globally enabled
|
||||
filtered = await callback.async_filter_deployments(
|
||||
model="claude-3",
|
||||
healthy_deployments=healthy_deployments,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"user_api_key_hash": user_key}},
|
||||
parent_otel_span=None,
|
||||
)
|
||||
# All deployments returned (user-key affinity disabled for this group)
|
||||
assert len(filtered) == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue