From 9bd43ca125ef87f8b7adfb7a63632370aec414f3 Mon Sep 17 00:00:00 2001 From: Julien Ambrosio Date: Thu, 20 Aug 2026 14:32:31 -0300 Subject: [PATCH] refactor(auth): centralize optional team alias lookup Co-authored-by: Cursor --- litellm/proxy/auth/auth_checks.py | 14 +++++++ litellm/proxy/auth/handle_jwt.py | 26 ++++-------- litellm/proxy/auth/user_api_key_auth.py | 16 +++----- .../proxy/auth/test_auth_checks.py | 41 +++++++++++++++++++ .../proxy/auth/test_handle_jwt.py | 4 +- 5 files changed, 71 insertions(+), 30 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e7b46ceaf51..8ade208c2b1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2480,6 +2480,20 @@ async def get_team_model_aliases( return aliases +async def get_team_model_aliases_for_team( + team_object: LiteLLM_TeamTable | None, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, +) -> dict[str, str] | None: # mutable-ok: delegates UserAPIKeyAuth.team_model_aliases dict contract + if team_object is None or team_object.model_id is None or prisma_client is None: + return None + return await get_team_model_aliases( + model_id=team_object.model_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + + async def _get_team_object_from_user_api_key_cache( team_id: str, prisma_client: PrismaClient, diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 0c359dbf8b3..fdbfa44ef7a 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -64,7 +64,7 @@ from .auth_checks import ( get_role_based_models, get_role_based_routes, get_team_membership, - get_team_model_aliases, + get_team_model_aliases_for_team, get_team_object, get_team_object_by_alias, get_user_object, @@ -1378,14 +1378,10 @@ class JWTAuthManager: model=requested_model, team_object=team_object, llm_router=llm_router, - 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.model_id is not None and prisma_client is not None - else None + team_model_aliases=await get_team_model_aliases_for_team( + team_object=team_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, ), ) ): @@ -1923,14 +1919,10 @@ class JWTAuthManager: model=requested_model, team_object=team_object, llm_router=llm_router, - 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.model_id is not None and prisma_client is not None - else None + team_model_aliases=await get_team_model_aliases_for_team( + team_object=team_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, ), ) except ProxyException: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 389cddbbc1f..19cd86249e8 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -49,7 +49,7 @@ from litellm.proxy.auth.auth_checks import ( get_end_user_object, get_jwt_key_mapping_object, get_project_object, - get_team_model_aliases, + get_team_model_aliases_for_team, get_team_object, get_user_object, is_valid_fallback_model, @@ -1379,16 +1379,10 @@ 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 + team_model_aliases=await get_team_model_aliases_for_team( + team_object=team_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, ), user_role=( LitellmUserRoles(user_object.user_role) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 1d1bd9ebf8a..109a5f3b2b3 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -47,6 +47,7 @@ from litellm.proxy.auth.auth_checks import ( _virtual_key_max_budget_check, _virtual_key_soft_budget_check, get_key_object, + get_team_model_aliases_for_team, get_user_object, vector_store_access_check, ) @@ -1315,6 +1316,46 @@ async def test_get_team_db_check_does_not_call_new_team_if_exists( mock_new_team.assert_not_called() +@pytest.mark.asyncio +async def test_get_team_model_aliases_for_team_guards_incomplete_context_and_delegates(): + cache = UserApiKeyCache() + prisma_client = MagicMock() + team_without_model_id = LiteLLM_TeamTable(team_id="plain-team") + team_with_model_id = LiteLLM_TeamTable(team_id="alias-team", model_id=7) + + with patch( + "litellm.proxy.auth.auth_checks.get_team_model_aliases", + new_callable=AsyncMock, + return_value={"requested": "target"}, + ) as mock_loader: + for team_object, client in ( + (None, prisma_client), + (team_without_model_id, prisma_client), + (team_with_model_id, None), + ): + assert ( + await get_team_model_aliases_for_team( + team_object=team_object, + prisma_client=client, + user_api_key_cache=cache, + ) + is None + ) + + result = await get_team_model_aliases_for_team( + team_object=team_with_model_id, + prisma_client=prisma_client, + user_api_key_cache=cache, + ) + + assert result == {"requested": "target"} + mock_loader.assert_awaited_once_with( + model_id=7, + prisma_client=prisma_client, + user_api_key_cache=cache, + ) + + # Vector Store Auth Check Tests diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index e491536a279..f42e6994148 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -4630,7 +4630,7 @@ async def test_find_team_with_model_access_loads_aliases_before_restricted_team_ return_value=team, ), patch( - "litellm.proxy.auth.handle_jwt.get_team_model_aliases", + "litellm.proxy.auth.handle_jwt.get_team_model_aliases_for_team", new_callable=AsyncMock, return_value={"claude-opus-5": "FW-Kimi-K3"}, ) as mock_get_aliases, @@ -5010,7 +5010,7 @@ async def test_resolve_db_team_fallback_loads_aliases_before_restricted_team_sel return_value=team, ), patch( - "litellm.proxy.auth.handle_jwt.get_team_model_aliases", + "litellm.proxy.auth.handle_jwt.get_team_model_aliases_for_team", new_callable=AsyncMock, return_value={"claude-opus-5": "FW-Kimi-K3"}, ) as mock_get_aliases,