From edd41cf88d0dae0d4529f20c8db8ddde2b58dfb3 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 20 Mar 2026 22:11:17 -0700 Subject: [PATCH] address greptile review feedback (greploop iteration 1) - Wrap add/remove_mcp_server_to_team in DB transactions (race condition fix) - Consolidate role + extra_permissions into single DB write (atomicity) - Remove silent try/except on team linking (surface errors to caller) - Add email fallback to _find_member_in_team (email-only members) - Add None guard on payload.server_id in update endpoint - Add field_validator on Member.extra_permissions (resource:action format) Co-Authored-By: Claude Opus 4.6 (1M context) --- litellm/proxy/_types.py | 16 +++ .../management_endpoints/common_utils.py | 21 +++- .../mcp_management_endpoints.py | 25 ++-- .../management_endpoints/team_endpoints.py | 59 ++++------ .../object_permission_utils.py | 111 +++++++++--------- 5 files changed, 122 insertions(+), 110 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 62547b309d0..8c8e5d37f00 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1622,6 +1622,22 @@ class Member(MemberBase): description="Granular permissions granted to this member (e.g. 'mcp:create', 'mcp:delete'). Used for per-member permission grants.", ) + @field_validator("extra_permissions", mode="before") + @classmethod + def validate_permission_format(cls, v): + """Validate that all permission strings follow the resource:action format.""" + if v is None: + return v + if not isinstance(v, list): + raise ValueError("extra_permissions must be a list of strings") + for perm in v: + if not isinstance(perm, str) or ":" not in perm: + raise ValueError( + f"Invalid permission format: '{perm}'. " + "Must follow 'resource:action' format (e.g. 'mcp:create')." + ) + return v + class OrgMember(MemberBase): role: Literal[ diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index fec6239813d..1776a4b5264 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -45,12 +45,25 @@ def _is_user_team_admin( def _find_member_in_team( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable ) -> Optional[Member]: - """Find and return the Member object for the given user in the team, or None.""" - if not user_api_key_dict.user_id: - return None + """Find and return the Member object for the given user in the team, or None. + + Matches by user_id first, then falls back to user_email for email-only members. + """ for member in team_obj.members_with_roles: - if member.user_id is not None and member.user_id == user_api_key_dict.user_id: + if ( + user_api_key_dict.user_id + and member.user_id is not None + and member.user_id == user_api_key_dict.user_id + ): return member + # Fallback: match by email for email-only members + if user_api_key_dict.user_email: + for member in team_obj.members_with_roles: + if ( + member.user_email is not None + and member.user_email == user_api_key_dict.user_email + ): + return member return None diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f9f0f2db318..e27e5b0bdbe 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1294,14 +1294,9 @@ if MCP_AVAILABLE: # Auto-assign server to team's ObjectPermissionTable if team-scoped if team_id and new_mcp_server.server_id: - try: - await add_mcp_server_to_team( - prisma_client, team_id, new_mcp_server.server_id - ) - except Exception as e: - verbose_proxy_logger.warning( - f"Failed to auto-assign MCP server {new_mcp_server.server_id} to team {team_id}: {e}" - ) + await add_mcp_server_to_team( + prisma_client, team_id, new_mcp_server.server_id + ) return _redact_mcp_credentials(new_mcp_server) @@ -1577,12 +1572,7 @@ if MCP_AVAILABLE: # Remove server from team's ObjectPermissionTable if team_id: - try: - await remove_mcp_server_from_team(prisma_client, team_id, server_id) - except Exception as e: - verbose_proxy_logger.warning( - f"Failed to remove MCP server {server_id} from team {team_id}: {e}" - ) + await remove_mcp_server_from_team(prisma_client, team_id, server_id) # TODO: Enterprise: Finish audit log trail if litellm.store_audit_logs: @@ -1913,6 +1903,13 @@ if MCP_AVAILABLE: ) team_servers = await _get_team_allowed_mcp_servers(team_obj) + if payload.server_id is None: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "server_id is required to update an MCP server." + }, + ) if payload.server_id not in team_servers: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index d6b0a4a6836..f36ed83e6f5 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2434,31 +2434,7 @@ async def team_member_update( rpm_limit=data.rpm_limit, ) - ### update team member role - if data.role is not None: - team_members: List[Member] = [] - for member in team_table.members_with_roles: - if member.user_id == received_user_id: - team_members.append( - Member( - user_id=member.user_id, - role=data.role, - user_email=data.user_email or member.user_email, - extra_permissions=member.extra_permissions, - ) - ) - else: - team_members.append(member) - - team_table.members_with_roles = team_members - - _db_team_members: List[dict] = [m.model_dump() for m in team_members] - await prisma_client.db.litellm_teamtable.update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore - ) - - ### update team member extra_permissions + ### Validate extra_permissions before applying any changes if data.extra_permissions is not None: from litellm.proxy.auth.permissions import VALID_PERMISSIONS @@ -2472,34 +2448,39 @@ async def team_member_update( }, ) - team_members_updated: List[Member] = [] + ### Apply role and extra_permissions updates in-memory, then do a single DB write + members_changed = data.role is not None or data.extra_permissions is not None + if members_changed: + updated_members: List[Member] = [] for member in team_table.members_with_roles: if member.user_id == received_user_id: - team_members_updated.append( + updated_members.append( Member( user_id=member.user_id, - role=member.role, - user_email=member.user_email, - extra_permissions=data.extra_permissions, + role=data.role if data.role is not None else member.role, + user_email=data.user_email or member.user_email, + extra_permissions=( + data.extra_permissions + if data.extra_permissions is not None + else member.extra_permissions + ), ) ) else: - team_members_updated.append(member) + updated_members.append(member) - team_table.members_with_roles = team_members_updated + team_table.members_with_roles = updated_members - _db_team_members_perms: List[dict] = [ - m.model_dump() for m in team_members_updated - ] + _db_team_members: List[dict] = [m.model_dump() for m in updated_members] await prisma_client.db.litellm_teamtable.update( where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members_perms)}, # type: ignore + data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore ) - # Invalidate team cache so permission changes take effect immediately - from litellm.proxy.proxy_server import user_api_key_cache + # Invalidate team cache so changes take effect immediately + from litellm.proxy.proxy_server import user_api_key_cache - user_api_key_cache.delete_cache(key="team_id:{}".format(data.team_id)) + user_api_key_cache.delete_cache(key="team_id:{}".format(data.team_id)) return TeamMemberUpdateResponse( team_id=data.team_id, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index d7e4e1fa173..1f2bef0d191 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -373,51 +373,54 @@ async def add_mcp_server_to_team( Add an MCP server ID to a team's ObjectPermissionTable.mcp_servers. If the team has no ObjectPermissionTable yet, one is created. + Uses a DB transaction to prevent race conditions when multiple + concurrent calls try to initialize the ObjectPermissionTable. """ - team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) - if team is None: - raise ValueError(f"Team {team_id} not found") + async with prisma_client.db.tx() as tx: + team = await tx.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + if team is None: + raise ValueError(f"Team {team_id} not found") - object_permission_id = team.object_permission_id or str(uuid.uuid4()) + object_permission_id = team.object_permission_id or str(uuid.uuid4()) - # Get existing object permission or start fresh - existing_mcp_servers: List[str] = [] - if team.object_permission_id: - existing_perm = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": team.object_permission_id} + # Get existing object permission or start fresh + existing_mcp_servers: List[str] = [] + if team.object_permission_id: + existing_perm = ( + await tx.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": team.object_permission_id} + ) ) - ) - if existing_perm and existing_perm.mcp_servers: - existing_mcp_servers = list(existing_perm.mcp_servers) + if existing_perm and existing_perm.mcp_servers: + existing_mcp_servers = list(existing_perm.mcp_servers) - # Add server if not already present - if server_id not in existing_mcp_servers: - existing_mcp_servers.append(server_id) + # Add server if not already present + if server_id not in existing_mcp_servers: + existing_mcp_servers.append(server_id) - # Upsert the ObjectPermissionTable - await prisma_client.db.litellm_objectpermissiontable.upsert( - where={"object_permission_id": object_permission_id}, - data={ - "create": { - "object_permission_id": object_permission_id, - "mcp_servers": existing_mcp_servers, + # Upsert the ObjectPermissionTable + await tx.litellm_objectpermissiontable.upsert( + where={"object_permission_id": object_permission_id}, + data={ + "create": { + "object_permission_id": object_permission_id, + "mcp_servers": existing_mcp_servers, + }, + "update": { + "mcp_servers": existing_mcp_servers, + }, }, - "update": { - "mcp_servers": existing_mcp_servers, - }, - }, - ) - - # Link the team to the ObjectPermissionTable if not already linked - if not team.object_permission_id: - await prisma_client.db.litellm_teamtable.update( - where={"team_id": team_id}, - data={"object_permission_id": object_permission_id}, ) + # Link the team to the ObjectPermissionTable if not already linked + if not team.object_permission_id: + await tx.litellm_teamtable.update( + where={"team_id": team_id}, + data={"object_permission_id": object_permission_id}, + ) + async def remove_mcp_server_from_team( prisma_client: PrismaClient, team_id: str, server_id: str @@ -426,24 +429,26 @@ async def remove_mcp_server_from_team( Remove an MCP server ID from a team's ObjectPermissionTable.mcp_servers. No-op if the team has no ObjectPermissionTable or the server isn't in the list. + Uses a DB transaction for consistency. """ - team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) - if team is None or not team.object_permission_id: - return - - existing_perm = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": team.object_permission_id} + async with prisma_client.db.tx() as tx: + team = await tx.litellm_teamtable.find_unique( + where={"team_id": team_id} ) - ) - if existing_perm is None or not existing_perm.mcp_servers: - return + if team is None or not team.object_permission_id: + return - updated_servers = [s for s in existing_perm.mcp_servers if s != server_id] - if len(updated_servers) != len(existing_perm.mcp_servers): - await prisma_client.db.litellm_objectpermissiontable.update( - where={"object_permission_id": team.object_permission_id}, - data={"mcp_servers": updated_servers}, + existing_perm = ( + await tx.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": team.object_permission_id} + ) ) + if existing_perm is None or not existing_perm.mcp_servers: + return + + updated_servers = [s for s in existing_perm.mcp_servers if s != server_id] + if len(updated_servers) != len(existing_perm.mcp_servers): + await tx.litellm_objectpermissiontable.update( + where={"object_permission_id": team.object_permission_id}, + data={"mcp_servers": updated_servers}, + )