mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
fix(router): pin JWT-authenticated callers by user id in deployment_affinity
This commit is contained in:
parent
2c30fe16b0
commit
68ffa1db23
2 changed files with 225 additions and 18 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue