mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
db27600ce9
commit
526156ed9d
4 changed files with 265 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue