mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
fix(mcp_server_manager.py): ensure only allowed MCP's are returned to the user, via rest endpoints
This commit is contained in:
parent
c80ae8e185
commit
a60c036f17
3 changed files with 111 additions and 22 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -422,6 +422,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/mcp/tools/call",
|
||||
"/mcp-rest/tools/list",
|
||||
"/mcp-rest/tools/call",
|
||||
"/v1/mcp/server",
|
||||
]
|
||||
|
||||
agent_routes = [
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue