fix(router): pin JWT-authenticated callers by user id in deployment_affinity

This commit is contained in:
mateo-berri 2026-09-03 10:43:41 -07:00
parent 2c30fe16b0
commit 68ffa1db23
2 changed files with 225 additions and 18 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:

View file

@ -998,3 +998,210 @@ 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:
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:
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")
assert alice_request["metadata"]["user_api_key_hash"] is None
await callback.async_pre_call_deployment_hook(
kwargs={
**alice_request,
"metadata": {**alice_request["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