fix: load team model_aliases on JWT user-direct auth path

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Julien Ambrosio 2026-08-20 13:51:31 -03:00
parent 3357ec8d34
commit 7d01090d5d
No known key found for this signature in database
GPG key ID: 7CC17BD06C342C97
2 changed files with 126 additions and 12 deletions

View file

@ -18,6 +18,7 @@ import fastapi
import orjson
from fastapi import HTTPException, Request, WebSocket, status
from fastapi.security.api_key import APIKeyHeader
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
@ -106,6 +107,11 @@ except ImportError as e:
enterprise_custom_auth = None
user_api_key_service_logger_obj: Final = ServiceLogging() # used for tracking latency on OTEL
_TEAM_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
class _TeamModelAliasesRow(BaseModel):
model_aliases: object = None
def _normalize_public_auth_route(route: str) -> str:
@ -1060,6 +1066,44 @@ async def _read_request_body_deferring_parse_failure(
return populate_request_with_path_params(request_data=parsed_body, request=request), None
async def _get_team_model_aliases(
model_id: int,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> dict[str, str] | None: # mutable-ok: UserAPIKeyAuth.team_model_aliases requires a dict
cache_key: Final = f"team_model_aliases:{model_id}"
cached: Final[object] = await user_api_key_cache.async_get_cache( # pyright: ignore[reportAny] # cache API is untyped
cache_key
)
if cached is not None:
return _TEAM_MODEL_ALIASES_ADAPTER.validate_python(cached)
row_value: Final[object] = ( # pyright: ignore[reportAny] # generated table result is dynamic
await prisma_client.db.litellm_modeltable.find_unique( # pyright: ignore[reportAny] # generated table API is dynamic
where={"id": model_id}, # mutable-ok: Prisma find_unique requires a dict
)
)
if row_value is None:
return None
row: Final = _TeamModelAliasesRow.model_validate(row_value, from_attributes=True)
raw_aliases: Final = row.model_aliases
if raw_aliases is None:
return None
aliases: Final = (
_TEAM_MODEL_ALIASES_ADAPTER.validate_json(raw_aliases)
if isinstance(raw_aliases, (str, bytes, bytearray))
else _TEAM_MODEL_ALIASES_ADAPTER.validate_python(raw_aliases)
)
await user_api_key_cache.async_set_cache( # pyright: ignore[reportUnknownMemberType] # cache API has untyped kwargs
key=cache_key,
value=aliases,
ttl=60,
)
return aliases
async def _record_unparsable_body_failure(
user_api_key_dict: UserAPIKeyAuth,
body_parse_exception: ProxyException,
@ -1378,6 +1422,17 @@ async def _user_api_key_auth_builder(
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
team_models=(team_object.models if team_object is not None else []),
team_model_aliases=(
await _get_team_model_aliases(
model_id=team_object.model_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
if team_object is not None
and team_object.model_id is not None
and prisma_client is not None
else None
),
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None

View file

@ -5,6 +5,7 @@ import sys
from contextlib import contextmanager
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import cast
from unittest.mock import ANY, AsyncMock, MagicMock, patch
sys.path.insert(
@ -34,6 +35,7 @@ from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_check_key_model_budget_with_fallback,
_get_team_model_aliases,
_PendingAutoRegister,
_matches_routing_override,
_reserve_budget_after_common_checks,
@ -1741,18 +1743,32 @@ def test_proxy_admin_jwt_auth_handles_no_team_object():
assert result.end_user_id is None
@pytest.mark.parametrize(
("model_id", "expected_aliases"),
[
pytest.param(None, None, id="without-team"),
pytest.param(1, {"claude-opus-5": "FW-Kimi-K3"}, id="with-team-aliases"),
],
)
@pytest.mark.asyncio
async def test_standard_jwt_auth_propagates_user_email():
"""
Standard (non-mapped) JWT auth must copy user_email from the resolved
LiteLLM_UserTable onto the returned UserAPIKeyAuth so spend logs attribute
user_api_key_user_email. Regression: this branch built
UserAPIKeyAuth(api_key=None, ...) with user_id but never user_email, so
the email was silently dropped even though the DB user row had it.
"""
async def test_standard_jwt_auth_propagates_user_identity_and_team_model_aliases(
model_id: int | None,
expected_aliases: dict[str, str] | None,
):
from litellm.models.team import LiteLLM_TeamTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import PrismaClient
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
user_api_key_cache = UserApiKeyCache()
find_model_table = AsyncMock(return_value=SimpleNamespace(model_aliases='{"claude-opus-5": "FW-Kimi-K3"}'))
prisma_client = cast(
PrismaClient,
SimpleNamespace(
db=SimpleNamespace(litellm_modeltable=SimpleNamespace(find_unique=find_model_table)),
),
)
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
@ -1762,14 +1778,22 @@ async def test_standard_jwt_auth_propagates_user_email():
user_email="human@example.com",
user_role="internal_user",
)
team_object = (
LiteLLM_TeamTable(
team_id="jwt-team",
model_id=model_id,
)
if model_id is not None
else None
)
mock_jwt_result = {
"is_proxy_admin": False,
"team_object": None,
"team_object": team_object,
"user_object": user_object,
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": None,
"team_id": team_object.team_id if team_object is not None else None,
"user_id": "jwt-human-user",
"user_email": "human@example.com",
"end_user_id": None,
@ -1789,7 +1813,7 @@ async def test_standard_jwt_auth_propagates_user_email():
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler),
@ -1812,6 +1836,41 @@ async def test_standard_jwt_auth_propagates_user_email():
assert result.user_id == "jwt-human-user"
assert result.user_email == "human@example.com"
assert result.api_key is None
assert cast(dict[str, str] | None, result.team_model_aliases) == expected_aliases
if model_id is None:
find_model_table.assert_not_awaited()
return
assert await _get_team_model_aliases(
model_id=model_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
) == {"claude-opus-5": "FW-Kimi-K3"}
find_model_table.assert_awaited_once_with(where={"id": 1})
@pytest.mark.asyncio
async def test_get_team_model_aliases_returns_none_when_model_row_missing():
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import PrismaClient
find_model_table = AsyncMock(return_value=None)
prisma_client = cast(
PrismaClient,
SimpleNamespace(
db=SimpleNamespace(litellm_modeltable=SimpleNamespace(find_unique=find_model_table)),
),
)
assert (
await _get_team_model_aliases(
model_id=99,
prisma_client=prisma_client,
user_api_key_cache=UserApiKeyCache(),
)
is None
)
find_model_table.assert_awaited_once_with(where={"id": 99})
@pytest.mark.asyncio