mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(auth): centralize optional team alias lookup
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
75a7188d1f
commit
9bd43ca125
5 changed files with 71 additions and 30 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue