diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f137902f1b7..e137ad15f1b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -203,6 +203,69 @@ if MCP_AVAILABLE: 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 @@ -1347,36 +1410,12 @@ if MCP_AVAILABLE: # 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: - existing_permission_id = getattr( - manager_team_obj, "object_permission_id", None - ) - - # Read existing mcp_servers list and append the new server - existing_mcp_servers: list = [] - if manager_team_obj.object_permission is not None: - existing_mcp_servers = ( - manager_team_obj.object_permission.mcp_servers or [] - ) - updated_mcp_servers = list( - set(existing_mcp_servers + [new_mcp_server.server_id]) - ) - - new_permission_id = await handle_update_object_permission_common( - data_json={ - "object_permission": { - "mcp_servers": updated_mcp_servers, - } - }, - existing_object_permission_id=existing_permission_id, + 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, ) - - # If the team had no object_permission_id, link the new one - if existing_permission_id is None and new_permission_id is not None: - await prisma_client.db.litellm_teamtable.update( - where={"team_id": manager_team_id}, - data={"object_permission_id": new_permission_id}, - ) except Exception as e: verbose_proxy_logger.exception( f"MCP server created but failed to auto-assign to team {manager_team_id}: {str(e)}" @@ -1617,28 +1656,10 @@ if MCP_AVAILABLE: # Remove server from the manager's team permission list if manager_team_obj is not None: try: - existing_permission_id = getattr( - manager_team_obj, "object_permission_id", None + await _remove_mcp_server_from_team( + server_id=server_id, + team_obj=manager_team_obj, ) - if ( - existing_permission_id is not None - and manager_team_obj.object_permission is not None - ): - existing_mcp_servers = ( - manager_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, - ) except Exception as e: verbose_proxy_logger.exception( f"MCP server deleted but failed to remove from team permissions: {str(e)}" 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 index 67e9dda4cef..ed601dfffa4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_manager_role.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_manager_role.py @@ -233,7 +233,7 @@ class TestCreateMcpServerAsManager: created_server.server_id = "new_server_id" created_server.credentials = None - mock_handle_update = AsyncMock() + mock_auto_assign = AsyncMock() with ( patch( @@ -263,38 +263,23 @@ class TestCreateMcpServerAsManager: ), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.handle_update_object_permission_common", - mock_handle_update, + "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 handle_update_object_permission_common was called with merged server list - mock_handle_update.assert_called_once() - call_kwargs = mock_handle_update.call_args.kwargs - mcp_servers = call_kwargs["data_json"]["object_permission"]["mcp_servers"] - assert "existing_server" in mcp_servers - assert "new_server_id" in mcp_servers - assert call_kwargs["existing_object_permission_id"] == "perm1" + # 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_create_links_new_permission_to_team_when_none_exists(self): - """When team has no object_permission_id, create should link the new one.""" - from litellm.proxy._types import NewMCPServerRequest + 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 ( - 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", + _auto_assign_mcp_server_to_team, ) # Team with NO object_permission_id @@ -307,46 +292,24 @@ class TestCreateMcpServerAsManager: ) mock_team.object_permission = None - created_server = MagicMock() - created_server.server_id = "new_server_id" - created_server.credentials = None - mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.update = AsyncMock() - with ( - patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", - return_value=mock_prisma, - ), - 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.handle_update_object_permission_common", - AsyncMock(return_value="new_perm_id"), - ), + # 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 add_mcp_server(payload=payload, user_api_key_dict=user_auth) + 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(