From a60c036f17e0ecac971baa44ae50dc6fc76fcde7 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 10 Feb 2026 10:28:15 -0800 Subject: [PATCH] fix(mcp_server_manager.py): ensure only allowed MCP's are returned to the user, via rest endpoints --- .../mcp_server/mcp_server_manager.py | 40 ++++++-- litellm/proxy/_types.py | 1 + .../test_mcp_management_endpoints.py | 92 ++++++++++++++++--- 3 files changed, 111 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 532aea249bf..8878c52b077 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -70,9 +70,7 @@ try: from mcp.shared.tool_name_validation import ( validate_tool_name, # pyright: ignore[reportAssignmentType] ) - from mcp.shared.tool_name_validation import ( - SEP_986_URL, - ) + from mcp.shared.tool_name_validation import SEP_986_URL except ImportError: from pydantic import BaseModel @@ -671,24 +669,47 @@ class MCPServerManager: return [ server.server_id for server in self.get_registry().values() - if server.allow_all_keys + if server.allow_all_keys is True ] async def get_allowed_mcp_servers( self, user_api_key_auth: Optional[UserAPIKeyAuth] = None ) -> List[str]: """ - Get the allowed MCP Servers for the user + Get the allowed MCP Servers for the user. + + Priority: + 1. If object_permission.mcp_servers is explicitly set, use it (even for admins) + 2. If admin and no object_permission, return all servers + 3. Otherwise, use standard permission checks """ from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view - # If admin, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth): - return list(self.get_registry().keys()) - allow_all_server_ids = self.get_allow_all_keys_server_ids() try: + # Check if object_permission.mcp_servers is explicitly set + has_explicit_object_permission = False + if user_api_key_auth and user_api_key_auth.object_permission: + # Check if mcp_servers is explicitly set (not None, empty list is valid) + if user_api_key_auth.object_permission.mcp_servers is not None: + has_explicit_object_permission = True + verbose_logger.debug( + f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}" + ) + + # If admin but NO explicit object permission, get all servers + if ( + user_api_key_auth + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): + verbose_logger.debug( + "Admin user without explicit object_permission - returning all servers" + ) + return list(self.get_registry().keys()) + + # Get allowed servers from object permissions (respects object_permission even for admins) allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( user_api_key_auth ) @@ -2244,6 +2265,7 @@ class MCPServerManager: from litellm.proxy.proxy_server import ( general_settings as proxy_general_settings, ) + return proxy_general_settings except ImportError: # Fallback if proxy_server not available diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 25a7292442f..bb3fd748ef3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -422,6 +422,7 @@ class LiteLLMRoutes(enum.Enum): "/mcp/tools/call", "/mcp-rest/tools/list", "/mcp-rest/tools/call", + "/v1/mcp/server", ] agent_routes = [ diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 83914e30354..6331e99462b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -2,14 +2,15 @@ import json import os import sys import types -from types import SimpleNamespace from datetime import datetime, timedelta +from types import SimpleNamespace from typing import List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient + from litellm._uuid import uuid from litellm.proxy.management_endpoints import ( mcp_management_endpoints as mgmt_endpoints, @@ -204,9 +205,7 @@ class TestListMCPServers: transport="http", ), ] - mock_manager.get_all_allowed_mcp_servers = AsyncMock( - return_value=mock_servers - ) + mock_manager.get_all_allowed_mcp_servers = AsyncMock(return_value=mock_servers) for idx, server in enumerate(mock_servers): server.credentials = {"auth_value": f"secret_{idx}"} @@ -374,9 +373,7 @@ class TestListMCPServers: transport="http", ), ] - mock_manager.get_all_allowed_mcp_servers = AsyncMock( - return_value=mock_servers - ) + mock_manager.get_all_allowed_mcp_servers = AsyncMock(return_value=mock_servers) for idx, server in enumerate(mock_servers): server.credentials = {"auth_value": f"secret_{idx}"} @@ -494,9 +491,7 @@ class TestListMCPServers: url="https://actions.zapier.com/mcp/sse", ), ] - mock_manager.get_all_allowed_mcp_servers = AsyncMock( - return_value=mock_servers - ) + mock_manager.get_all_allowed_mcp_servers = AsyncMock(return_value=mock_servers) for idx, server in enumerate(mock_servers): server.credentials = {"auth_value": f"secret_{idx}"} @@ -540,6 +535,71 @@ class TestListMCPServers: assert server.alias == "Allowed Zapier MCP" assert server.url == "https://actions.zapier.com/mcp/sse" + @pytest.mark.asyncio + async def test_admin_user_with_object_permission_respects_mcp_servers(self): + """ + Test that admin users with explicit object_permission.mcp_servers + only see the servers specified in object_permission. + + Scenario: Admin user has object_permission.mcp_servers set to specific servers + Expected: Only those servers are returned, not all servers in the registry + """ + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # Create mock object permission with specific servers + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-obj-perm-id", + mcp_servers=["server-1", "server-2"], # Only these two servers + mcp_access_groups=[], + mcp_tool_permissions={}, + vector_stores=[], + agents=[], + agent_access_groups=[], + ) + + # Create admin user with object permission + mock_user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin_user_id", + api_key="admin_api_key", + object_permission=mock_object_permission, + object_permission_id="test-obj-perm-id", + ) + + # Mock servers that the user should see + server_1 = generate_mock_mcp_server_db_record( + server_id="server-1", alias="Server 1", url="https://server1.example.com" + ) + server_2 = generate_mock_mcp_server_db_record( + server_id="server-2", alias="Server 2", url="https://server2.example.com" + ) + + # Mock manager + mock_manager = MagicMock() + mock_manager.get_all_allowed_mcp_servers = AsyncMock( + return_value=[server_1, server_2] + ) + + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts", + AsyncMock(return_value=[mock_user_auth]), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_all_mcp_servers, + ) + + result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth) + + # Verify results - should only return the 2 servers in object_permission + assert len(result) == 2 + server_ids = {server.server_id for server in result} + assert server_ids == {"server-1", "server-2"} + + # Verify credentials are redacted + assert all(server.credentials is None for server in result) @pytest.mark.asyncio async def test_fetch_single_mcp_server_redacts_credentials(self): @@ -947,6 +1007,7 @@ class TestTemporaryMCPSessionEndpoints: fallback_client_id="server-1", ) + class TestUpdateMCPServer: """Test suite for update MCP server functionality""" @@ -954,7 +1015,7 @@ class TestUpdateMCPServer: async def test_update_mcp_server_respects_extra_headers(self): """ Test that updating an MCP server with extra_headers properly saves the field. - + This test ensures that extra_headers field in UpdateMCPServerRequest is properly handled and persisted when updating an MCP server. """ @@ -1030,7 +1091,10 @@ class TestUpdateMCPServer: # First arg is prisma_client, second is the payload (UpdateMCPServerRequest) called_payload = call_args[0][1] assert called_payload.server_id == "test-server-1" - assert called_payload.extra_headers == ["X-Custom-Header", "X-Another-Header"] + assert called_payload.extra_headers == [ + "X-Custom-Header", + "X-Another-Header", + ] assert called_payload.alias == "Updated Test Server" # Verify the result includes extra_headers @@ -1131,7 +1195,9 @@ class TestMCPRegistryEndpoint: mock_manager = MagicMock() mock_manager.get_registry.return_value = {mock_server.server_id: mock_server} # The registry endpoint uses get_filtered_registry (filters by client IP) - mock_manager.get_filtered_registry.return_value = {mock_server.server_id: mock_server} + mock_manager.get_filtered_registry.return_value = { + mock_server.server_id: mock_server + } with patch_proxy_general_settings({"enable_mcp_registry": True}), patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",