Merge pull request #39594 from BerriAI/litellm_lit6867_jwt_deployment_affinity

fix(router): pin JWT-authenticated callers by user id in deployment_affinity
This commit is contained in:
Mateo Wang 2026-09-03 13:14:17 -07:00 committed by GitHub
commit c137828a92
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 237 additions and 28 deletions

View file

@ -82,6 +82,7 @@ class DeploymentAffinityCheck(CustomLogger):
"""
CACHE_KEY_PREFIX = "deployment_affinity:v1"
USER_ID_AFFINITY_PREFIX: Final = "user_id:"
def __init__(
self,
@ -253,15 +254,6 @@ class DeploymentAffinityCheck(CustomLogger):
hashed_user_key: Final = cls._hash_user_key(user_key) if user_key is not None else "unscoped"
return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{hashed_user_key}:{session_id}"
@staticmethod
def _get_user_key_from_metadata_dict(metadata: dict) -> str | None:
# NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the
# OpenAI `user` parameter, which is an end-user identifier).
user_key: Final = metadata.get("user_api_key_hash")
if user_key is None:
return None
return str(user_key)
@staticmethod
def _get_session_id_from_metadata_dict(metadata: dict) -> str | None:
session_id: Final = metadata.get("session_id")
@ -285,22 +277,30 @@ class DeploymentAffinityCheck(CustomLogger):
return metadata_dicts
@staticmethod
def _get_user_key_from_request_kwargs(request_kwargs: dict) -> str | None:
def _first_metadata_value(metadata_dicts: Sequence[dict], key: str) -> str | None:
value: Final = next((metadata[key] for metadata in metadata_dicts if metadata.get(key) is not None), None)
return None if value is None else str(value)
@classmethod
def _get_user_key_from_request_kwargs(cls, request_kwargs: dict) -> str | None:
"""
Extract a stable affinity key from request kwargs.
Source (proxy): `metadata.user_api_key_hash`
Source (proxy): `metadata.user_api_key_hash` for virtual-key callers. JWT-authenticated
callers carry no key hash, so their `metadata.user_api_key_user_id` stands in for it,
namespaced under `USER_ID_AFFINITY_PREFIX` so a user id can never alias a key hash.
Note: the OpenAI `user` parameter is an end-user identifier and is intentionally
not used for deployment affinity.
"""
# Check metadata dicts (Proxy usage)
for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs):
user_key = DeploymentAffinityCheck._get_user_key_from_metadata_dict(metadata=metadata)
if user_key is not None:
return user_key
return None
metadata_dicts: Final = cls._iter_metadata_dicts(request_kwargs)
user_api_key_hash: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_hash")
if user_api_key_hash is not None:
return user_api_key_hash
user_id: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_user_id")
if user_id is None:
return None
return f"{cls.USER_ID_AFFINITY_PREFIX}{user_id}"
@staticmethod
def _get_session_id_from_request_kwargs(request_kwargs: dict) -> str | None:
@ -533,9 +533,9 @@ class DeploymentAffinityCheck(CustomLogger):
return typed_healthy_deployments
verbose_router_logger.debug(
"DeploymentAffinityCheck: api-key affinity hit -> deployment=%s user_key=%s",
"DeploymentAffinityCheck: caller affinity hit -> deployment=%s user_key=%s",
model_id,
self._shorten_for_logs(user_key),
self._shorten_for_logs(self._hash_user_key(user_key)),
)
return [deployment]
@ -626,7 +626,7 @@ class DeploymentAffinityCheck(CustomLogger):
deployment_model_name,
model_id,
self.ttl_seconds,
self._shorten_for_logs(user_key),
self._shorten_for_logs(self._hash_user_key(user_key)),
)
else:
verbose_router_logger.debug(

View file

@ -1,11 +1,12 @@
import asyncio
import itertools
import json
from collections.abc import Sequence
from typing import Final
from unittest.mock import AsyncMock, patch
import pytest
import json
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
@ -102,11 +103,10 @@ async def test_async_user_key_affinity_routes_to_same_deployment():
# Deterministic routing: first selection uses seq[0], second selection attempts seq[1]
# unless the list has been filtered to length=1 by deployment affinity.
choice_calls = {"count": 0}
choice_calls: Final = itertools.count(1)
def deterministic_choice(seq):
choice_calls["count"] += 1
if choice_calls["count"] == 1:
def deterministic_choice(seq: Sequence[dict[str, object]]) -> dict[str, object]:
if next(choice_calls) == 1:
return seq[0]
return seq[1] if len(seq) > 1 else seq[0]
@ -998,3 +998,212 @@ async def test_model_group_affinity_config_overrides_global():
)
# All deployments returned (user-key affinity disabled for this group)
assert len(filtered) == 2
def _jwt_metadata(user_id: str) -> dict[str, str | None]:
return {"user_api_key_hash": None, "user_api_key_user_id": user_id}
def _two_deployments(model_group: str) -> list[dict]:
return [
{
"model_name": model_group,
"litellm_params": {"model": "openai/gpt-5.4-mini"},
"model_info": {"id": "openai-deployment-a"},
},
{
"model_name": model_group,
"litellm_params": {"model": "openai/gpt-5.4-mini"},
"model_info": {"id": "openai-deployment-b"},
},
]
@pytest.mark.asyncio
async def test_async_jwt_user_affinity_routes_to_same_deployment():
"""
JWT-authenticated proxy requests carry no `user_api_key_hash`, only `user_api_key_user_id`.
They must still pin to one deployment per user.
"""
model_group = "gpt-5.4-mini"
router = litellm.Router(
model_list=[
{
"model_name": model_group,
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "mock-api-key-a"},
"model_info": {"id": "openai-deployment-a"},
},
{
"model_name": model_group,
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "mock-api-key-b"},
"model_info": {"id": "openai-deployment-b"},
},
],
optional_pre_call_checks=["deployment_affinity"],
)
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( # test-quality-ok: simple-shuffle has no injectable RNG; forcing the other pick is what proves the pin overrides the strategy
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):
first_response = await router.acompletion(
model=model_group,
messages=[{"role": "user", "content": "Reply with the single word ok"}],
mock_response="ok",
metadata=_jwt_metadata("jwt-user-alice"),
)
second_response = await router.acompletion(
model=model_group,
messages=[{"role": "user", "content": "Reply with the single word ok"}],
mock_response="ok",
metadata=_jwt_metadata("jwt-user-alice"),
)
first_model_id = first_response._hidden_params["model_id"]
assert first_model_id in ("openai-deployment-a", "openai-deployment-b")
assert second_response._hidden_params["model_id"] == first_model_id
@pytest.mark.asyncio
async def test_proxy_jwt_auth_metadata_pins_per_user():
"""
The metadata the proxy stamps for a JWT caller (`UserAPIKeyAuth(api_key=None, user_id=<sub>)`)
must claim a pin and be read back by the filter, and another JWT user must not inherit it.
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
model_group = "gpt-5.4-mini"
healthy_deployments = _two_deployments(model_group)
callback = DeploymentAffinityCheck(
cache=DualCache(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
)
def proxy_request(user_id: str) -> dict[str, object]:
return LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data={"model": model_group, "messages": [{"role": "user", "content": "hi"}], "metadata": {}},
user_api_key_dict=UserAPIKeyAuth(api_key=None, user_id=user_id),
_metadata_variable_name="metadata",
)
alice_request = proxy_request("jwt-user-alice")
alice_metadata = alice_request["metadata"]
assert isinstance(alice_metadata, dict)
assert alice_metadata["user_api_key_hash"] is None
await callback.async_pre_call_deployment_hook(
kwargs={
**alice_request,
"metadata": {**alice_metadata, "deployment_model_name": model_group},
"model_info": {"id": "openai-deployment-b"},
},
call_type=None,
)
alice_pinned = await callback.async_filter_deployments(
model=model_group,
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs=alice_request,
parent_otel_span=None,
)
assert [deployment["model_info"]["id"] for deployment in alice_pinned] == ["openai-deployment-b"]
bob_filtered = await callback.async_filter_deployments(
model=model_group,
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs=proxy_request("jwt-user-bob"),
parent_otel_span=None,
)
assert bob_filtered == healthy_deployments
@pytest.mark.asyncio
async def test_jwt_user_id_never_reads_a_virtual_key_pin():
"""
A JWT user id that happens to equal a virtual key's 64-hex hash must not read that key's pin.
"""
model_group = "gpt-5.4-mini"
healthy_deployments = _two_deployments(model_group)
callback = DeploymentAffinityCheck(
cache=DualCache(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
)
key_hash = "a" * 64
await callback.async_pre_call_deployment_hook(
kwargs={
"metadata": {"user_api_key_hash": key_hash, "deployment_model_name": model_group},
"model_info": {"id": "openai-deployment-b"},
},
call_type=None,
)
key_pinned = await callback.async_filter_deployments(
model=model_group,
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"metadata": {"user_api_key_hash": key_hash}},
parent_otel_span=None,
)
assert [deployment["model_info"]["id"] for deployment in key_pinned] == ["openai-deployment-b"]
lookalike_jwt_user = await callback.async_filter_deployments(
model=model_group,
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"metadata": _jwt_metadata(key_hash)},
parent_otel_span=None,
)
assert lookalike_jwt_user == healthy_deployments
@pytest.mark.asyncio
async def test_virtual_key_hash_wins_over_user_id_for_affinity():
"""
A virtual-key caller with a user id pins on the key hash, so two keys owned by one user
keep independent pins.
"""
model_group = "gpt-5.4-mini"
healthy_deployments = _two_deployments(model_group)
callback = DeploymentAffinityCheck(
cache=DualCache(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
)
await callback.async_pre_call_deployment_hook(
kwargs={
"metadata": {
"user_api_key_hash": "key-one",
"user_api_key_user_id": "shared-user",
"deployment_model_name": model_group,
},
"model_info": {"id": "openai-deployment-b"},
},
call_type=None,
)
other_key_same_user = await callback.async_filter_deployments(
model=model_group,
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"metadata": {"user_api_key_hash": "key-two", "user_api_key_user_id": "shared-user"}},
parent_otel_span=None,
)
assert other_key_same_user == healthy_deployments