mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: load team model_aliases on JWT user-direct auth path
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
3357ec8d34
commit
7d01090d5d
2 changed files with 126 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue