mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(projects): let org admins see projects of every team in their orgs
/project/list and /project/info only knew proxy admins and team members, so an org admin saw projects only for teams they had personally joined. Both now use the same team access check as /team/info.
This commit is contained in:
parent
42d50c8db9
commit
0b4fd60aee
4 changed files with 266 additions and 25 deletions
|
|
@ -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},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue