fix(mcp_server_manager.py): ensure only allowed MCP's are returned to the user, via rest endpoints

This commit is contained in:
Krrish Dholakia 2026-02-10 10:28:15 -08:00
parent c80ae8e185
commit a60c036f17
3 changed files with 111 additions and 22 deletions

View file

@ -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

View file

@ -422,6 +422,7 @@ class LiteLLMRoutes(enum.Enum):
"/mcp/tools/call",
"/mcp-rest/tools/list",
"/mcp-rest/tools/call",
"/v1/mcp/server",
]
agent_routes = [

View file

@ -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",