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:
Sameer Kankute 2026-03-19 14:44:01 +05:30
parent e5baa2232f
commit 528daa8cf4
4 changed files with 371 additions and 28 deletions

View file

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

View file

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

View file

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

View file

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