diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index fbef33c32ed..0bc102889bc 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -40,8 +40,8 @@ def _prepare_mcp_server_data( """ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - # Convert model to dict - data_dict = data.model_dump(exclude_none=True) + # Convert model to dict, excluding fields not in the DB schema + data_dict = data.model_dump(exclude_none=True, exclude={"team_id"}) # Ensure alias is always present in the dict (even if None) if "alias" not in data_dict: data_dict["alias"] = getattr(data, "alias", None) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9e86680e355..67ae31bea8e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1144,6 +1144,10 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): None, description="Server-managed: set by the endpoint; caller values are overridden.", ) + team_id: Optional[str] = Field( + None, + description="Team ID to assign the MCP server to. Required for team MCP managers.", + ) @model_validator(mode="before") @classmethod @@ -1608,8 +1612,9 @@ class Member(MemberBase): role: Literal[ "admin", "user", + "mcp_server_manager", ] = Field( - description="The role of the user within the team. 'admin' users can manage team settings and members, 'user' is a regular team member" + description="The role of the user within the team. 'admin' users can manage team settings and members, 'user' is a regular team member, 'mcp_server_manager' can manage MCP servers for the team" ) diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index efc42d3355c..092b22cee6f 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -41,6 +41,18 @@ def _is_user_team_admin( return False +def _is_user_team_mcp_manager( + user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable +) -> bool: + 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 + ) and member.role == "mcp_server_manager": + return True + + return False + + async def _is_user_org_admin_for_team( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable ) -> bool: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 3e5b729cea6..8c8bd629884 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -21,7 +21,7 @@ import json import os from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import Any, Dict, Iterable, List, Literal, Optional +from typing import Any, Dict, Iterable, List, Literal, Optional, Tuple from fastapi import ( APIRouter, @@ -112,6 +112,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, + LiteLLM_TeamTableCachedObj, LitellmUserRoles, MakeMCPServersPublicRequest, MCPApprovalStatus, @@ -130,11 +131,142 @@ if MCP_AVAILABLE: ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import _read_request_body - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.management_endpoints.common_utils import ( + _is_user_team_mcp_manager, + _user_has_admin_view, + ) + from litellm.proxy.management_helpers.object_permission_utils import ( + _get_team_allowed_mcp_servers, + handle_update_object_permission_common, + ) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.types.mcp import MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer + async def _assert_can_manage_team_mcp_server( + user_api_key_dict: UserAPIKeyAuth, + team_id: Optional[str] = None, + server_id: Optional[str] = None, + ) -> Tuple[str, LiteLLM_TeamTableCachedObj]: + """ + Verify that the caller is an MCP server manager for a team and (for edit/delete) + that the target server belongs to that team. + + Returns a tuple of (team_id, team_obj) for downstream use. + Raises HTTPException(400) if no team_id can be determined. + Raises HTTPException(403) if the caller is not an MCP manager or server not in team. + """ + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + resolved_team_id = team_id or user_api_key_dict.team_id + if not resolved_team_id: + raise HTTPException( + status_code=400, + detail={"error": "team_id is required for MCP server manager operations."}, + ) + + # When the API key is team-scoped, ensure the request team_id matches + if ( + team_id + and user_api_key_dict.team_id + and team_id != user_api_key_dict.team_id + ): + raise HTTPException( + status_code=403, + detail={"error": "team_id does not match the API key's team."}, + ) + + team_obj = await get_team_object( + team_id=resolved_team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, # bypass cache to get fresh object_permission + ) + + if not _is_user_team_mcp_manager(user_api_key_dict, team_obj): + raise HTTPException( + status_code=403, + detail={ + "error": f"User does not have mcp_server_manager role in team {resolved_team_id}." + }, + ) + + if server_id is not None: + team_server_ids = await _get_team_allowed_mcp_servers(team_obj) + if server_id not in team_server_ids: + raise HTTPException( + status_code=403, + detail={ + "error": f"MCP server {server_id} is not assigned to team {resolved_team_id}." + }, + ) + + return resolved_team_id, team_obj + + async def _auto_assign_mcp_server_to_team( + server_id: str, + team_id: str, + team_obj: LiteLLM_TeamTableCachedObj, + prisma_client: Any, + ) -> None: + """ + Add an MCP server to a team's ObjectPermissionTable and link the + permission back to the team if it didn't have one yet. + + Uses handle_update_object_permission (the team-endpoint helper) so + the object_permission_id linkage follows the same pattern as + team_endpoints.update_team. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + handle_update_object_permission, + ) + + existing_mcp_servers: list = [] + if team_obj.object_permission is not None: + existing_mcp_servers = team_obj.object_permission.mcp_servers or [] + updated_mcp_servers = list(set(existing_mcp_servers + [server_id])) + + # Build the data dict in the same shape team_endpoints uses + data_json: Dict[str, Any] = { + "object_permission": {"mcp_servers": updated_mcp_servers}, + } + data_json = await handle_update_object_permission( + data_json=data_json, + existing_team_row=team_obj, + ) + + # If handle_update_object_permission produced an object_permission_id, + # persist it on the team row (it sets data_json["object_permission_id"]). + if "object_permission_id" in data_json: + await prisma_client.db.litellm_teamtable.update( + where={"team_id": team_id}, + data={"object_permission_id": data_json["object_permission_id"]}, + ) + + async def _remove_mcp_server_from_team( + server_id: str, + team_obj: LiteLLM_TeamTableCachedObj, + ) -> None: + """Remove a server ID from a team's ObjectPermissionTable.mcp_servers list.""" + from litellm.proxy.proxy_server import prisma_client + + existing_permission_id = getattr(team_obj, "object_permission_id", None) + if ( + existing_permission_id is None + or team_obj.object_permission is None + ): + return + + existing_mcp_servers = team_obj.object_permission.mcp_servers or [] + updated_mcp_servers = [s for s in existing_mcp_servers if s != server_id] + + await handle_update_object_permission_common( + data_json={"object_permission": {"mcp_servers": updated_mcp_servers}}, + existing_object_permission_id=existing_permission_id, + prisma_client=prisma_client, + ) + @dataclass class _TemporaryMCPServerEntry: server: MCPServer @@ -1206,16 +1338,27 @@ if MCP_AVAILABLE: # Validate and normalize payload fields validate_and_normalize_mcp_server_payload(payload) - # AuthZ - restrict only proxy admins to create mcp servers - if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={ - "error": "User does not have permission to create mcp servers. You can only create mcp servers if you are a PROXY_ADMIN." - }, + # AuthZ - proxy admins or team MCP managers can create MCP servers + is_proxy_admin = LitellmUserRoles.PROXY_ADMIN == user_api_key_dict.user_role + manager_team_id: Optional[str] = None + + manager_team_obj = None + if not is_proxy_admin: + # Check if the user is an MCP manager for the specified team + if payload.team_id is None: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "team_id is required when creating MCP servers as a team MCP manager." + }, + ) + manager_team_id, manager_team_obj = await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_api_key_dict, + team_id=payload.team_id, ) - elif payload.server_id is not None: - # fail if the mcp server with id already exists + + # Fail if the MCP server with this id already exists + if payload.server_id is not None: mcp_server = await get_mcp_server(prisma_client, payload.server_id) if mcp_server is not None: raise HTTPException( @@ -1224,7 +1367,8 @@ if MCP_AVAILABLE: "error": f"MCP Server with id {payload.server_id} already exists. Cannot create another." }, ) - elif ( + + if ( SpecialMCPServerName.all_team_servers == payload.server_id or SpecialMCPServerName.all_proxy_servers == payload.server_id ): @@ -1255,12 +1399,28 @@ if MCP_AVAILABLE: # Ensure registry is up to date by reloading from database await global_mcp_server_manager.reload_servers_from_database() + except Exception as e: verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Error creating mcp server: {str(e)}"}, ) + + # Auto-assign the server to the manager's team (separate from create + # so a failure here doesn't mask the successfully created server). + if manager_team_id is not None and manager_team_obj is not None: + try: + await _auto_assign_mcp_server_to_team( + server_id=new_mcp_server.server_id, + team_id=manager_team_id, + team_obj=manager_team_obj, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception( + f"MCP server created but failed to auto-assign to team {manager_team_id}: {str(e)}" + ) return _redact_mcp_credentials(new_mcp_server) @router.post( @@ -1473,15 +1633,12 @@ if MCP_AVAILABLE: "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - # Authz - restrict only admins to delete mcp servers + # Authz - proxy admins or team MCP managers can delete MCP servers + manager_team_obj = None if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={ - "error": "Call not allowed to delete MCP server. User is not a proxy admin. route={}".format( - "DELETE /v1/mcp/server" - ) - }, + _, manager_team_obj = await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_api_key_dict, + server_id=server_id, ) # try to delete the mcp server @@ -1497,6 +1654,18 @@ if MCP_AVAILABLE: # Ensure registry is up to date by reloading from database await global_mcp_server_manager.reload_servers_from_database() + # Remove server from the manager's team permission list + if manager_team_obj is not None: + try: + await _remove_mcp_server_from_team( + server_id=server_id, + team_obj=manager_team_obj, + ) + except Exception as e: + verbose_proxy_logger.exception( + f"MCP server deleted but failed to remove from team permissions: {str(e)}" + ) + # TODO: Enterprise: Finish audit log trail if litellm.store_audit_logs: pass @@ -1798,15 +1967,11 @@ if MCP_AVAILABLE: # Validate and normalize payload fields validate_and_normalize_mcp_server_payload(payload) - # Authz - restrict only admins to delete mcp servers + # Authz - proxy admins or team MCP managers can update MCP servers if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={ - "error": "Call not allowed to update MCP server. User is not a proxy admin. route={}".format( - "PUT /v1/mcp/server" - ) - }, + await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_api_key_dict, + server_id=payload.server_id, ) # try to update the mcp server diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_manager_role.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_manager_role.py new file mode 100644 index 00000000000..86b0419fb22 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_manager_role.py @@ -0,0 +1,559 @@ +import pytest + +pytest.importorskip("mcp", reason="mcp package not installed") + +from litellm.proxy._types import ( + LiteLLM_TeamTable, + LitellmUserRoles, + Member, + UserAPIKeyAuth, +) +from litellm.proxy.management_endpoints.common_utils import ( + _is_user_team_mcp_manager, +) +from litellm.proxy.management_endpoints import ( + mcp_management_endpoints as mgmt_endpoints, +) + + +class TestIsUserTeamMcpManager: + def test_mcp_server_manager_role_returns_true(self): + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + ) + team = LiteLLM_TeamTable( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager") + ], + ) + assert _is_user_team_mcp_manager(user_auth, team) is True + + def test_regular_user_role_returns_false(self): + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + ) + team = LiteLLM_TeamTable( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="user")], + ) + assert _is_user_team_mcp_manager(user_auth, team) is False + + def test_admin_role_returns_false(self): + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + ) + team = LiteLLM_TeamTable( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="admin")], + ) + assert _is_user_team_mcp_manager(user_auth, team) is False + + def test_user_not_in_team_returns_false(self): + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user2", + api_key="sk-test", + ) + team = LiteLLM_TeamTable( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager") + ], + ) + assert _is_user_team_mcp_manager(user_auth, team) is False + + +from unittest.mock import AsyncMock, MagicMock, patch +from fastapi import HTTPException +from litellm.proxy._types import LiteLLM_TeamTableCachedObj + + +@pytest.mark.asyncio +@pytest.mark.skipif( + not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed" +) +class TestAssertCanManageTeamMcpServer: + async def test_mcp_manager_with_team_id_succeeds(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _assert_can_manage_team_mcp_server, + ) + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", + ) + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="mcp_server_manager")], + ) + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object", + AsyncMock(return_value=mock_team), + ): + team_id, team_obj = await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_auth, team_id="team1" + ) + assert team_id == "team1" + assert team_obj == mock_team + + async def test_regular_user_gets_403(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _assert_can_manage_team_mcp_server, + ) + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", + ) + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="user")], + ) + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object", + AsyncMock(return_value=mock_team), + ): + with pytest.raises(Exception) as exc_info: + await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_auth, team_id="team1" + ) + assert exc_info.value.status_code == 403 + + async def test_admin_gets_403(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _assert_can_manage_team_mcp_server, + ) + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", + ) + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="admin")], + ) + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object", + AsyncMock(return_value=mock_team), + ): + with pytest.raises(Exception) as exc_info: + await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_auth, team_id="team1" + ) + assert exc_info.value.status_code == 403 + + async def test_mcp_manager_server_in_team_succeeds(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _assert_can_manage_team_mcp_server, + ) + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", team_id="team1", + ) + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="mcp_server_manager")], + ) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object", + AsyncMock(return_value=mock_team), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_allowed_mcp_servers", + AsyncMock(return_value={"server1", "server2"}), + ), + ): + team_id, team_obj = await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_auth, server_id="server1" + ) + assert team_id == "team1" + assert team_obj == mock_team + + async def test_mcp_manager_server_not_in_team_gets_403(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _assert_can_manage_team_mcp_server, + ) + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", team_id="team1", + ) + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[Member(user_id="user1", role="mcp_server_manager")], + ) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object", + AsyncMock(return_value=mock_team), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_allowed_mcp_servers", + AsyncMock(return_value={"server2", "server3"}), + ), + ): + with pytest.raises(Exception) as exc_info: + await _assert_can_manage_team_mcp_server( + user_api_key_dict=user_auth, server_id="server1" + ) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.skipif( + not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed" +) +class TestCreateMcpServerAsManager: + async def test_create_auto_assigns_to_team(self): + """MCP manager creating a server should auto-assign it to their team's ObjectPermissionTable.""" + from litellm.proxy._types import NewMCPServerRequest + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + add_mcp_server, + ) + + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + team_id="team1", + ) + + payload = NewMCPServerRequest( + server_name="test-server", + url="https://example.com/mcp", + team_id="team1", + ) + + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager"), + ], + object_permission_id="perm1", + ) + mock_team.object_permission = MagicMock(mcp_servers=["existing_server"]) + + created_server = MagicMock() + created_server.server_id = "new_server_id" + created_server.credentials = None + + mock_auto_assign = AsyncMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server", + AsyncMock(return_value=("team1", mock_team)), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + AsyncMock(return_value=created_server), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + MagicMock( + add_server=AsyncMock(), + reload_servers_from_database=AsyncMock(), + ), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._auto_assign_mcp_server_to_team", + mock_auto_assign, + ), + ): + await add_mcp_server(payload=payload, user_api_key_dict=user_auth) + + # Verify _auto_assign_mcp_server_to_team was called with the right args + mock_auto_assign.assert_called_once() + call_kwargs = mock_auto_assign.call_args.kwargs + assert call_kwargs["server_id"] == "new_server_id" + assert call_kwargs["team_id"] == "team1" + assert call_kwargs["team_obj"] == mock_team + + async def test_auto_assign_links_new_permission_to_team(self): + """_auto_assign_mcp_server_to_team should create permission and link to team.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _auto_assign_mcp_server_to_team, + ) + + # Team with NO object_permission_id + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager"), + ], + object_permission_id=None, + ) + mock_team.object_permission = None + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.update = AsyncMock() + + # handle_update_object_permission sets object_permission_id in data_json + async def fake_handle(data_json, existing_team_row): + data_json["object_permission_id"] = "new_perm_id" + return data_json + + with patch( + "litellm.proxy.management_endpoints.team_endpoints.handle_update_object_permission", + side_effect=fake_handle, + ): + await _auto_assign_mcp_server_to_team( + server_id="new_server_id", + team_id="team1", + team_obj=mock_team, + prisma_client=mock_prisma, + ) + + # Verify the team was updated with the new object_permission_id + mock_prisma.db.litellm_teamtable.update.assert_called_once_with( + where={"team_id": "team1"}, + data={"object_permission_id": "new_perm_id"}, + ) + + +@pytest.mark.asyncio +@pytest.mark.skipif( + not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed" +) +class TestEditMcpServerAsManager: + async def test_edit_succeeds_for_mcp_manager(self): + """MCP manager should be able to edit a server assigned to their team.""" + from litellm.proxy._types import UpdateMCPServerRequest + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + edit_mcp_server, + ) + + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + team_id="team1", + ) + + payload = UpdateMCPServerRequest( + server_id="server1", + description="Updated description", + ) + + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager"), + ], + ) + + updated_server = MagicMock() + updated_server.server_id = "server1" + updated_server.credentials = None + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server", + AsyncMock(return_value=("team1", mock_team)), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server", + AsyncMock(return_value=updated_server), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + MagicMock( + update_server=AsyncMock(), + reload_servers_from_database=AsyncMock(), + ), + ), + ): + result = await edit_mcp_server(payload=payload, user_api_key_dict=user_auth) + # Should not raise — edit succeeded + assert result is not None + + async def test_edit_fails_for_server_not_in_team(self): + """MCP manager should get 403 when editing a server not in their team.""" + from litellm.proxy._types import UpdateMCPServerRequest + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + edit_mcp_server, + ) + + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + team_id="team1", + ) + + payload = UpdateMCPServerRequest( + server_id="server_not_in_team", + description="Updated description", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload", + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server", + AsyncMock(side_effect=HTTPException(status_code=403, detail="Not in team")), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await edit_mcp_server(payload=payload, user_api_key_dict=user_auth) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.skipif( + not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed" +) +class TestDeleteMcpServerAsManager: + async def test_delete_succeeds_and_cleans_up_team(self): + """MCP manager deleting a server should also remove it from team permissions.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + remove_mcp_server, + ) + + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + team_id="team1", + ) + + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager"), + ], + ) + + deleted_server = MagicMock() + deleted_server.server_id = "server1" + + mock_remove_from_team = AsyncMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server", + AsyncMock(return_value=("team1", mock_team)), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_mcp_server", + AsyncMock(return_value=deleted_server), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + MagicMock( + remove_server=MagicMock(), + reload_servers_from_database=AsyncMock(), + ), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._remove_mcp_server_from_team", + mock_remove_from_team, + ), + ): + response = await remove_mcp_server( + server_id="server1", user_api_key_dict=user_auth + ) + assert response.status_code == 202 + + # Verify team cleanup was called + mock_remove_from_team.assert_called_once() + call_kwargs = mock_remove_from_team.call_args.kwargs + assert call_kwargs["server_id"] == "server1" + assert call_kwargs["team_obj"] == mock_team + + async def test_delete_fails_for_server_not_in_team(self): + """MCP manager should get 403 when deleting a server not in their team.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + remove_mcp_server, + ) + + user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user1", + api_key="sk-test", + team_id="team1", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server", + AsyncMock(side_effect=HTTPException(status_code=403, detail="Not in team")), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await remove_mcp_server( + server_id="server_not_in_team", user_api_key_dict=user_auth + ) + assert exc_info.value.status_code == 403 + + async def test_remove_mcp_server_from_team_helper(self): + """_remove_mcp_server_from_team should update the permission list without the deleted server.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _remove_mcp_server_from_team, + ) + + mock_team = LiteLLM_TeamTableCachedObj( + team_id="team1", + members_with_roles=[ + Member(user_id="user1", role="mcp_server_manager"), + ], + object_permission_id="perm1", + ) + mock_team.object_permission = MagicMock( + mcp_servers=["server1", "server2", "server3"] + ) + + mock_handle_common = AsyncMock() + + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.handle_update_object_permission_common", + mock_handle_common, + ), + ): + await _remove_mcp_server_from_team( + server_id="server2", + team_obj=mock_team, + ) + + mock_handle_common.assert_called_once() + call_kwargs = mock_handle_common.call_args.kwargs + updated_servers = call_kwargs["data_json"]["object_permission"]["mcp_servers"] + assert "server2" not in updated_servers + assert "server1" in updated_servers + assert "server3" in updated_servers + assert call_kwargs["existing_object_permission_id"] == "perm1"