feat(mcp/): allows admin to prevent llm's from accidentally deleting github repo's even if user is allowed to do this

This commit is contained in:
Krrish Dholakia 2025-09-27 19:36:11 -07:00
parent db27600ce9
commit 526156ed9d
4 changed files with 265 additions and 3 deletions

View file

@ -213,6 +213,8 @@ class MCPServerManager:
),
mcp_info=mcp_info,
extra_headers=server_config.get("extra_headers", None),
allowed_tools=server_config.get("allowed_tools", None),
disallowed_tools=server_config.get("disallowed_tools", None),
access_groups=server_config.get("access_groups", None),
)
self.config_mcp_servers[server_id] = new_server
@ -277,6 +279,8 @@ class MCPServerManager:
args=getattr(mcp_server, "args", None) or [],
env=env_dict,
access_groups=getattr(mcp_server, "mcp_access_groups", None),
allowed_tools=getattr(mcp_server, "allowed_tools", None),
disallowed_tools=getattr(mcp_server, "disallowed_tools", None),
)
self.registry[mcp_server.server_id] = new_server
verbose_logger.debug(f"Added MCP Server: {name_for_prefix}")
@ -569,6 +573,16 @@ class MCPServerManager:
)
return prefixed_tools
def check_allowed_or_banned_tools(self, tool_name: str, server: MCPServer) -> bool:
"""
Check if the tool is allowed or banned for the given server
"""
if server.allowed_tools:
return tool_name in server.allowed_tools
if server.disallowed_tools:
return tool_name not in server.disallowed_tools
return True
async def pre_call_tool_check(
self,
name: str,
@ -576,7 +590,18 @@ class MCPServerManager:
server_name_from_prefix: str,
user_api_key_auth: Optional[UserAPIKeyAuth],
proxy_logging_obj: ProxyLogging,
server: MCPServer,
):
## check if the tool is allowed or banned for the given server
if not self.check_allowed_or_banned_tools(name, server):
raise HTTPException(
status_code=403,
detail={
"error": f"Tool {name} is not allowed for server {server.name}. Contact proxy admin to allow this tool."
},
)
pre_hook_kwargs = {
"name": name,
"arguments": arguments,
@ -700,6 +725,7 @@ class MCPServerManager:
server_name_from_prefix=server_name_from_prefix,
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=mcp_server,
)
# Get server-specific auth header if available

View file

@ -25,7 +25,6 @@ mcp_servers:
client_id: os.environ/GITHUB_OAUTH_CLIENT_ID
client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET
scopes: ["public_repo", "user:email"]
extra_headers: ["custom_key"]
# allowed_tools: ["list_tools"]
allowed_tools: ["list_tools"]
# disallowed_tools: ["repo_delete"]

View file

@ -23,6 +23,8 @@ class MCPServer(BaseModel):
extra_headers: Optional[List[str]] = (
None # allow admin to specify which headers to forward to the MCP server
)
allowed_tools: Optional[List[str]] = None
disallowed_tools: Optional[List[str]] = None
# OAuth-specific fields
client_id: Optional[str] = None
client_secret: Optional[str] = None

View file

@ -1,8 +1,9 @@
import sys
from datetime import datetime
from unittest.mock import MagicMock, AsyncMock
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
# Add the parent directory to the path so we can import litellm
sys.path.insert(0, "../../../../../")
@ -420,6 +421,240 @@ class TestMCPServerManager:
assert result["status"] == "healthy"
assert result["tools_count"] == 1
@pytest.mark.asyncio
async def test_pre_call_tool_check_allowed_tools_list_allows_tool(self):
"""Test pre_call_tool_check allows tool when it's in allowed_tools list"""
manager = MCPServerManager()
# Create server with allowed_tools list
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=["allowed_tool", "another_allowed_tool"],
disallowed_tools=None,
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
# This should not raise an exception
await manager.pre_call_tool_check(
name="allowed_tool",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_allowed_tools_list_blocks_tool(self):
"""Test pre_call_tool_check blocks tool when it's not in allowed_tools list"""
manager = MCPServerManager()
# Create server with allowed_tools list
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=["allowed_tool", "another_allowed_tool"],
disallowed_tools=None,
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# This should raise an HTTPException
with pytest.raises(HTTPException) as exc_info:
await manager.pre_call_tool_check(
name="blocked_tool",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
assert exc_info.value.status_code == 403
assert (
"Tool blocked_tool is not allowed for server test-server"
in exc_info.value.detail["error"]
)
assert (
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_disallowed_tools_list_allows_tool(self):
"""Test pre_call_tool_check allows tool when it's not in disallowed_tools list"""
manager = MCPServerManager()
# Create server with disallowed_tools list
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=None,
disallowed_tools=["banned_tool", "another_banned_tool"],
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
# This should not raise an exception
await manager.pre_call_tool_check(
name="allowed_tool",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_disallowed_tools_list_blocks_tool(self):
"""Test pre_call_tool_check blocks tool when it's in disallowed_tools list"""
manager = MCPServerManager()
# Create server with disallowed_tools list
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=None,
disallowed_tools=["banned_tool", "another_banned_tool"],
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# This should raise an HTTPException
with pytest.raises(HTTPException) as exc_info:
await manager.pre_call_tool_check(
name="banned_tool",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
assert exc_info.value.status_code == 403
assert (
"Tool banned_tool is not allowed for server test-server"
in exc_info.value.detail["error"]
)
assert (
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_no_restrictions_allows_any_tool(self):
"""Test pre_call_tool_check allows any tool when no restrictions are set"""
manager = MCPServerManager()
# Create server with no tool restrictions
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=None,
disallowed_tools=None,
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
# This should not raise an exception
await manager.pre_call_tool_check(
name="any_tool",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_allowed_tools_takes_precedence(self):
"""Test that allowed_tools list takes precedence over disallowed_tools list"""
manager = MCPServerManager()
# Create server with both allowed_tools and disallowed_tools
# Note: The logic in check_allowed_or_banned_tools prioritizes allowed_tools
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.stdio,
allowed_tools=["tool1", "tool2"],
disallowed_tools=["tool2", "tool3"], # tool2 is in both lists
)
# Mock dependencies
user_api_key_auth = MagicMock()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
# tool2 should be allowed since it's in allowed_tools (takes precedence)
await manager.pre_call_tool_check(
name="tool2",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
# tool3 should be blocked since it's not in allowed_tools
with pytest.raises(HTTPException) as exc_info:
await manager.pre_call_tool_check(
name="tool3",
arguments={"param": "value"},
server_name_from_prefix="test-server",
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=server,
)
assert exc_info.value.status_code == 403
assert (
"Tool tool3 is not allowed for server test-server"
in exc_info.value.detail["error"]
)
if __name__ == "__main__":
pytest.main([__file__])