From 7d01090d5dac66da519f830303db5fc352348c66 Mon Sep 17 00:00:00 2001 From: Julien Ambrosio Date: Thu, 20 Aug 2026 13:51:31 -0300 Subject: [PATCH] fix: load team model_aliases on JWT user-direct auth path Co-authored-by: Cursor --- litellm/proxy/auth/user_api_key_auth.py | 55 ++++++++++++ .../proxy/auth/test_user_api_key_auth.py | 83 ++++++++++++++++--- 2 files changed, 126 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 99592d44f9b..842c669eea7 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ab7e3d9701c..7b16674b93c 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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