refactor(auth): centralize optional team alias lookup

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Julien Ambrosio 2026-08-20 14:32:31 -03:00
parent 75a7188d1f
commit 9bd43ca125
No known key found for this signature in database
GPG key ID: 7CC17BD06C342C97
5 changed files with 71 additions and 30 deletions

View file

@ -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,

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

@ -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,