diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index d134c39c91b..366fe9e58fd 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -22,7 +22,9 @@ from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import delete_cached_project_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.management.teams.access import is_team_admin +from litellm.proxy.management.teams.access import TEAM_OR_ORG_ADMIN, TeamAccess, is_team_admin, is_team_member +from litellm.proxy.management.teams.dependencies import get_team_access +from litellm.proxy.management.users.service import org_admin_org_ids from litellm.proxy.management_endpoints.common_utils import _set_object_metadata_field from litellm.proxy.management_endpoints.team_admin_field_permissions import team_admin_may_manage_projects from litellm.proxy.management_helpers.utils import ( @@ -118,6 +120,35 @@ async def _check_user_permission_for_project( return is_team_admin(user_api_key_dict, team) or user_api_key_dict.user_id in (team.admins or []) +async def _can_view_team_projects( + user_api_key_dict: UserAPIKeyAuth, + team_id: str | None, + prisma_client: PrismaClient, + team_access: TeamAccess, +) -> bool: + if user_api_key_has_admin_view(user_api_key_dict): + return True + if not team_id or not user_api_key_dict.user_id: + return False + team_row: Final = await _team_table(prisma_client).find_unique(where={"team_id": team_id}) + if team_row is None: + return False + team: Final = LiteLLM_TeamTable.model_validate(team_row.model_dump()) + return is_team_member(user_api_key_dict, team) or await team_access.allows( + user_api_key_dict, team, TEAM_OR_ORG_ADMIN + ) + + +def _visible_projects_where(team_ids: list[str], admin_org_ids: frozenset[str]) -> dict[str, object]: + member_scope: Final[dict[str, object]] = {"team_id": {"in": team_ids}} + if not admin_org_ids: + return member_scope + org_scope: Final[dict[str, object]] = { + "litellm_team_table": {"is": {"organization_id": {"in": sorted(admin_org_ids)}}} + } + return {"OR": [member_scope, org_scope]} + + async def _validate_team_exists( team_id: str, prisma_client: PrismaClient, @@ -973,6 +1004,7 @@ async def delete_project( async def project_info( project_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + team_access: TeamAccess = Depends(get_team_access), ): """ Get information about a specific project @@ -1009,21 +1041,7 @@ async def project_info( param="project_id", ) - # Check if user has access to this project (admin or team member) - is_admin = user_api_key_has_admin_view(user_api_key_dict) - is_team_member = False - - if project.team_id and user_api_key_dict.user_id: - team = await _team_table(prisma_client).find_unique(where={"team_id": project.team_id}) - if team: - caller_user_id = user_api_key_dict.user_id - for m in team.members_with_roles or []: - m_user_id = m.get("user_id") if isinstance(m, dict) else getattr(m, "user_id", None) - if m_user_id == caller_user_id: - is_team_member = True - break - - if not (is_admin or is_team_member): + if not await _can_view_team_projects(user_api_key_dict, project.team_id, prisma_client, team_access): raise HTTPException( status_code=403, detail={"error": "You don't have access to this project"}, @@ -1072,16 +1090,13 @@ async def list_projects( include={"litellm_budget_table": True, "object_permission": True} ) else: - # Look up the user's team memberships via the reverse-index on - # LiteLLM_UserTable.teams (maintained by team_member_add alongside - # members_with_roles). This avoids a full scan of all team rows. user_record: Final = await _user_table(prisma_client).find_unique( where={"user_id": user_api_key_dict.user_id}, + include={"organization_memberships": True}, ) - user_team_ids: list[str] = user_record.teams if user_record is not None and user_record.teams else [] - + user: Final = None if user_record is None else LiteLLM_UserTable.model_validate(user_record.model_dump()) projects = await _project_table(prisma_client).find_many( - where={"team_id": {"in": user_team_ids}}, + where=_visible_projects_where(user.teams if user is not None else [], org_admin_org_ids(user)), include={"litellm_budget_table": True, "object_permission": True}, ) diff --git a/litellm/proxy/management/teams/access.py b/litellm/proxy/management/teams/access.py index 77af588c636..869ee01a302 100644 --- a/litellm/proxy/management/teams/access.py +++ b/litellm/proxy/management/teams/access.py @@ -44,6 +44,12 @@ class TeamAccess: return await self.org_roles.is_org_admin(caller.user_id, team.organization_id) +def is_team_member(user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable) -> bool: + return user_api_key_dict.user_id is not None and any( + member.user_id == user_api_key_dict.user_id for member in team_obj.members_with_roles + ) + + def is_team_admin(user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable) -> bool: return any( member.user_id is not None and member.user_id == user_api_key_dict.user_id and member.role == "admin" diff --git a/litellm/proxy/management/users/service.py b/litellm/proxy/management/users/service.py index 5bf19c0c885..8a9c30ddbeb 100644 --- a/litellm/proxy/management/users/service.py +++ b/litellm/proxy/management/users/service.py @@ -10,13 +10,20 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient, ProxyLogging -def holds_org_admin(user: LiteLLM_UserTable | None, organization_id: str) -> bool: - return user is not None and any( - membership.organization_id == organization_id and membership.user_role == LitellmUserRoles.ORG_ADMIN.value +def org_admin_org_ids(user: LiteLLM_UserTable | None) -> frozenset[str]: + if user is None: + return frozenset() + return frozenset( + membership.organization_id for membership in user.organization_memberships or [] + if membership.user_role == LitellmUserRoles.ORG_ADMIN.value ) +def holds_org_admin(user: LiteLLM_UserTable | None, organization_id: str) -> bool: + return organization_id in org_admin_org_ids(user) + + @dataclass(frozen=True, slots=True) class PrismaOrgRoles: prisma_client: PrismaClient | None diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 226755e7b8e..9c61a9fb551 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1,5 +1,9 @@ import os import traceback +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Final from litellm._uuid import uuid from unittest import mock @@ -37,8 +41,14 @@ from litellm.proxy._types import ( DeleteProjectRequest, NewTeamRequest, UserAPIKeyAuth, + LiteLLM_OrganizationMembershipTable, + LiteLLM_TeamTable, + LiteLLM_UserTable, + Member, ProxyException, ) +from litellm.proxy.management.teams.access import TeamAccess +from litellm.proxy.management.users.service import PrismaOrgRoles proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) @@ -1380,3 +1390,206 @@ async def test_new_project_flag_on_access_group_model_returns_400(monkeypatch): assert "prod-models" in str(exc_info.value) assert "expand to multiple models at request time" in str(exc_info.value) + + +_ACCESS_NOW: Final = datetime.now(timezone.utc) + + +@dataclass(frozen=True, slots=True) +class _ProjectRow: + project_id: str + team_id: str | None + + +def _membership(user_id: str, org_id: str, role: LitellmUserRoles) -> LiteLLM_OrganizationMembershipTable: + return LiteLLM_OrganizationMembershipTable( + user_id=user_id, organization_id=org_id, user_role=role.value, created_at=_ACCESS_NOW, updated_at=_ACCESS_NOW + ) + + +def _user(user_id: str, teams: tuple[str, ...] = (), admin_of: tuple[str, ...] = ()) -> LiteLLM_UserTable: + return LiteLLM_UserTable( + user_id=user_id, + teams=list(teams), + organization_memberships=[_membership(user_id, "org-a", LitellmUserRoles.INTERNAL_USER)] + + [_membership(user_id, org_id, LitellmUserRoles.ORG_ADMIN) for org_id in admin_of], + ) + + +_TEAMS: Final = { + team.team_id: team + for team in ( + LiteLLM_TeamTable( + team_id="team-a1", + organization_id="org-a", + members_with_roles=[Member(user_id="member", role="user"), Member(user_id="team-admin", role="admin")], + ), + LiteLLM_TeamTable(team_id="team-a2", organization_id="org-a"), + LiteLLM_TeamTable( + team_id="team-b1", + organization_id="org-b", + members_with_roles=[Member(user_id="org-admin-a-member-b1", role="user")], + ), + LiteLLM_TeamTable(team_id="team-orgless"), + ) +} +_PROJECTS: Final = ( + _ProjectRow("p-a1", "team-a1"), + _ProjectRow("p-a2", "team-a2"), + _ProjectRow("p-b1", "team-b1"), + _ProjectRow("p-orgless", "team-orgless"), + _ProjectRow("p-teamless", None), +) +_USERS: Final = { + user.user_id: user + for user in ( + _user("member", teams=("team-a1",)), + _user("team-admin", teams=("team-a1",)), + _user("org-admin-a", admin_of=("org-a",)), + _user("org-admin-b", admin_of=("org-b",)), + _user("org-admin-a-member-b1", teams=("team-b1",), admin_of=("org-a",)), + ) +} + + +def _row_matches(row: object, where: Mapping[str, object]) -> bool: + return all(_condition_holds(row, key, condition) for key, condition in where.items()) + + +def _condition_holds(row: object, key: str, condition) -> bool: + if key == "OR": + return any(_row_matches(row, branch) for branch in condition) + if key == "litellm_team_table": + team = _TEAMS.get(getattr(row, "team_id")) + return team is not None and _row_matches(team, condition["is"]) + value = getattr(row, key) + return value in condition["in"] if isinstance(condition, dict) else value == condition + + +class _FakeTeamTable: + async def find_unique(self, where: Mapping[str, str], include: object = None) -> LiteLLM_TeamTable | None: + return _TEAMS.get(where["team_id"]) + + +class _FakeProjectTable: + async def find_unique(self, where: Mapping[str, str], include: object = None) -> _ProjectRow | None: + return next((p for p in _PROJECTS if p.project_id == where["project_id"]), None) + + async def find_many(self, where: Mapping[str, object] | None = None, include: object = None) -> list[_ProjectRow]: + return [p for p in _PROJECTS if where is None or _row_matches(p, where)] + + +class _FakeUserTable: + async def find_unique( + self, where: Mapping[str, str], include: Mapping[str, bool] | None = None + ) -> LiteLLM_UserTable | None: + user = _USERS.get(where["user_id"]) + if user is None or (include or {}).get("organization_memberships"): + return user + return user.model_copy(update={"organization_memberships": None}) + + +@pytest.fixture +def project_access_db(monkeypatch): + db = mock.MagicMock( + litellm_teamtable=_FakeTeamTable(), litellm_projecttable=_FakeProjectTable(), litellm_usertable=_FakeUserTable() + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock.MagicMock(db=db)) + + +async def _team_access_over_cached_users() -> TeamAccess: + cache = UserApiKeyCache() + for user in _USERS.values(): + await cache.async_set_cache(key=user.user_id, value=user) + return TeamAccess(org_roles=PrismaOrgRoles(None, cache, proxy_logging_obj)) + + +def _caller(user_id: str, role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER) -> UserAPIKeyAuth: + return UserAPIKeyAuth(user_role=role, api_key="sk-caller", user_id=user_id) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_id", "expected"), + [ + ("member", {"p-a1"}), + ("team-admin", {"p-a1"}), + ("org-admin-a", {"p-a1", "p-a2"}), + ("org-admin-b", {"p-b1"}), + ("org-admin-a-member-b1", {"p-a1", "p-a2", "p-b1"}), + ("unknown-user", set()), + ], +) +async def test_list_projects_scopes_to_teams_the_caller_can_view(project_access_db, user_id, expected): + from litellm_enterprise.proxy.management_endpoints.project_endpoints import list_projects + + projects = await list_projects(user_api_key_dict=_caller(user_id)) + + assert {p.project_id for p in projects} == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +async def test_list_projects_admin_view_sees_every_project(project_access_db, role): + from litellm_enterprise.proxy.management_endpoints.project_endpoints import list_projects + + projects = await list_projects(user_api_key_dict=_caller("someone", role)) + + assert {p.project_id for p in projects} == {p.project_id for p in _PROJECTS} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_id", "project_id"), + [ + ("member", "p-a1"), + ("team-admin", "p-a1"), + ("org-admin-a", "p-a1"), + ("org-admin-a", "p-a2"), + ("org-admin-a-member-b1", "p-b1"), + ], +) +async def test_project_info_allows_team_members_and_org_admins_of_the_team_org(project_access_db, user_id, project_id): + project = await project_info( + project_id=project_id, + user_api_key_dict=_caller(user_id), + team_access=await _team_access_over_cached_users(), + ) + + assert project.project_id == project_id + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_id", "project_id"), + [ + ("member", "p-a2"), + ("team-admin", "p-b1"), + ("org-admin-b", "p-a2"), + ("org-admin-a", "p-b1"), + ("org-admin-a", "p-orgless"), + ("org-admin-a", "p-teamless"), + ], +) +async def test_project_info_denies_callers_who_cannot_view_the_team(project_access_db, user_id, project_id): + with pytest.raises(ProxyException) as exc_info: + await project_info( + project_id=project_id, + user_api_key_dict=_caller(user_id), + team_access=await _team_access_over_cached_users(), + ) + + assert str(exc_info.value.code) == "403" + assert "You don't have access to this project" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_project_info_missing_project_is_404_even_for_org_admins(project_access_db): + with pytest.raises(ProxyException) as exc_info: + await project_info( + project_id="p-missing", + user_api_key_dict=_caller("org-admin-a"), + team_access=await _team_access_over_cached_users(), + ) + + assert str(exc_info.value.code) == "404"