From 526156ed9d065721b57be5490a28df6ff7f2640b Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 27 Sep 2025 19:36:11 -0700 Subject: [PATCH] feat(mcp/): allows admin to prevent llm's from accidentally deleting github repo's even if user is allowed to do this --- .../mcp_server/mcp_server_manager.py | 26 ++ litellm/proxy/_new_secret_config.yaml | 3 +- .../types/mcp_server/mcp_server_manager.py | 2 + .../mcp_server/test_mcp_server_manager.py | 237 +++++++++++++++++- 4 files changed, 265 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a3a5d93a3fd..6423a9ae153 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index e7494bd0dab..804cf2cf2cf 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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"] diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 4327bd5afe7..3e0c2b20e39 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 3237de37636..126a4ebb896 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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__])