Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py: 23%
55 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2MCP Security Guardrail for LiteLLM.
4Validates that MCP servers referenced in request tools are registered
5on the LiteLLM gateway. Blocks or alerts when unregistered servers are found.
6"""
8from typing import Any, Final, Literal
10from fastapi import HTTPException
12from litellm._logging import verbose_proxy_logger
13from litellm.integrations.custom_guardrail import (
14 CustomGuardrail,
15 log_guardrail_information,
16)
17from litellm.proxy._types import UserAPIKeyAuth
18from litellm.responses.mcp.litellm_proxy_mcp_handler import (
19 LITELLM_PROXY_MCP_SERVER_URL_PREFIX,
20)
21from litellm.types.guardrails import GuardrailEventHooks
24class MCPSecurityGuardrail(CustomGuardrail):
25 def __init__(
26 self,
27 on_violation: Literal["block", "alert"] = "block",
28 **kwargs,
29 ):
30 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
31 super().__init__(**kwargs)
32 self.on_violation = on_violation
34 @classmethod
35 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
36 return [GuardrailEventHooks.pre_call]
38 @log_guardrail_information
39 async def async_pre_call_hook(
40 self,
41 user_api_key_dict: UserAPIKeyAuth,
42 cache: Any,
43 data: dict,
44 call_type: str,
45 ) -> Exception | str | dict | None:
46 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True:
47 return data
49 unregistered: Final = self._find_unregistered_mcp_servers(data)
50 if not unregistered:
51 return data
53 message: Final = (
54 f"MCP Security: request references unregistered MCP server(s): "
55 f"{', '.join(sorted(unregistered))}. "
56 f"Only servers registered on this gateway are allowed."
57 )
59 if self.on_violation == "block":
60 raise HTTPException(
61 status_code=400,
62 detail={
63 "error": "Violated guardrail policy",
64 "guardrail": "mcp_security",
65 "unregistered_servers": sorted(unregistered),
66 "detection_message": message,
67 },
68 )
69 else:
70 verbose_proxy_logger.warning(message)
72 return data
74 @staticmethod
75 def _extract_mcp_server_names_from_tools(tools: list[dict]) -> set[str]:
76 """Extract MCP server names from tools with type=mcp and litellm_proxy server_url."""
77 server_names: Final[set[str]] = set()
78 for tool in tools:
79 if not isinstance(tool, dict):
80 continue
81 if tool.get("type") != "mcp":
82 continue
83 server_url = tool.get("server_url", "")
84 if not isinstance(server_url, str):
85 continue
86 if server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX):
87 name = server_url[len(LITELLM_PROXY_MCP_SERVER_URL_PREFIX) :]
88 if name:
89 server_names.add(name)
90 return server_names
92 @staticmethod
93 def _find_unregistered_mcp_servers(data: dict) -> set[str]:
94 """Check tools in data against the MCP server registry. Returns set of unregistered server names."""
95 tools: Final = data.get("tools")
96 if not tools or not isinstance(tools, list):
97 return set()
99 requested_servers: Final = MCPSecurityGuardrail._extract_mcp_server_names_from_tools(tools)
100 if not requested_servers:
101 return set()
103 from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
104 global_mcp_server_manager,
105 )
107 registry: Final = global_mcp_server_manager.get_registry()
108 registered_names: Final = set(registry.keys())
110 return requested_servers - registered_names