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

1""" 

2MCP Security Guardrail for LiteLLM. 

3 

4Validates that MCP servers referenced in request tools are registered 

5on the LiteLLM gateway. Blocks or alerts when unregistered servers are found. 

6""" 

7 

8from typing import Any, Final, Literal 

9 

10from fastapi import HTTPException 

11 

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 

22 

23 

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 

33 

34 @classmethod 

35 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

36 return [GuardrailEventHooks.pre_call] 

37 

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 

48 

49 unregistered: Final = self._find_unregistered_mcp_servers(data) 

50 if not unregistered: 

51 return data 

52 

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 ) 

58 

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) 

71 

72 return data 

73 

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 

91 

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() 

98 

99 requested_servers: Final = MCPSecurityGuardrail._extract_mcp_server_names_from_tools(tools) 

100 if not requested_servers: 

101 return set() 

102 

103 from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( 

104 global_mcp_server_manager, 

105 ) 

106 

107 registry: Final = global_mcp_server_manager.get_registry() 

108 registered_names: Final = set(registry.keys()) 

109 

110 return requested_servers - registered_names