Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py: 26%

96 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2MCP End User Permission Guardrail Hook 

3 

4Enforces end user permissions for MCP server access via apply_guardrail: 

5- input_type="request" → filter tools the end user cannot access 

6 

7Permission logic: 

8- No end_user_id → allow all (key/team-level permissions apply) 

9- end_user_id, no mcp_servers → allow all (default) 

10- end_user_id + mcp_servers → allow only those servers 

11""" 

12 

13from typing import TYPE_CHECKING, Any, Final, Literal 

14 

15from litellm._logging import verbose_proxy_logger 

16from litellm.integrations.custom_guardrail import ( 

17 CustomGuardrail, 

18 log_guardrail_information, 

19) 

20from litellm.proxy._types import LiteLLM_ObjectPermissionTable 

21from litellm.types.guardrails import GuardrailEventHooks 

22from litellm.types.utils import GenericGuardrailAPIInputs 

23 

24if TYPE_CHECKING: 24 ↛ 25line 24 didn't jump to line 25 because the condition on line 24 was never true

25 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

26 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

27 

28GUARDRAIL_NAME: Final = "mcp_end_user_permission" 

29 

30 

31class MCPEndUserPermissionGuardrail(CustomGuardrail): 

32 """ 

33 Guardrail that enforces end user permissions for MCP server access. 

34 

35 Runs on input only (pre-call). Filters tools in the request that the 

36 end user is not permitted to call based on their object_permission. 

37 

38 end_user_object_permission is populated on UserAPIKeyAuth during auth. 

39 The guardrail resolves it via a cached get_end_user_object lookup — 

40 no extra DB round-trip when the cache is warm. 

41 """ 

42 

43 def __init__(self, **kwargs): 

44 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

45 super().__init__(**kwargs) 

46 verbose_proxy_logger.debug("MCP End User Permission Guardrail initialized") 

47 

48 # ------------------------------------------------------------------ 

49 # apply_guardrail — filters MCP tools on the request side only 

50 # ------------------------------------------------------------------ 

51 

52 @log_guardrail_information 

53 async def apply_guardrail( 

54 self, 

55 inputs: GenericGuardrailAPIInputs, 

56 request_data: dict, 

57 input_type: Literal["request", "response"] = "request", 

58 logging_obj: "LiteLLMLoggingObj | None" = None, 

59 ) -> GenericGuardrailAPIInputs: 

60 """ 

61 Filters MCP tools the end user cannot access based on their 

62 object_permission.mcp_servers / mcp_access_groups settings. 

63 """ 

64 object_permission: Final = await self._resolve_end_user_object_permission(request_data) 

65 return await self._check_request_tools(inputs, object_permission) 

66 

67 # ------------------------------------------------------------------ 

68 # Private — request-side tool filtering 

69 # ------------------------------------------------------------------ 

70 

71 async def _check_request_tools( 

72 self, 

73 inputs: GenericGuardrailAPIInputs, 

74 object_permission: LiteLLM_ObjectPermissionTable | None, 

75 ) -> GenericGuardrailAPIInputs: 

76 tools: Final = inputs.get("tools") 

77 if not tools: 

78 return inputs 

79 

80 allowed_mcp_servers: Final = await self._get_allowed_mcp_servers_from_object_permission(object_permission) 

81 if allowed_mcp_servers is None: 

82 return inputs # No restrictions → pass through unchanged 

83 

84 verbose_proxy_logger.debug("MCP guardrail: end user restricted to MCP servers: %s", allowed_mcp_servers) 

85 

86 filtered_tools: Final = [] 

87 removed_tools: Final = [] 

88 

89 for tool in tools: 

90 tool_name = self._get_tool_name_from_definition(tool) 

91 server_name = self._extract_mcp_server_name(tool_name) if tool_name else None 

92 

93 if server_name is None: 

94 # Not an MCP tool (no prefix) or unrecognised format → keep 

95 filtered_tools.append(tool) 

96 elif server_name in allowed_mcp_servers: 

97 filtered_tools.append(tool) 

98 else: 

99 removed_tools.append(tool_name) 

100 verbose_proxy_logger.warning( 

101 "MCP guardrail: removing tool '%s' (server: '%s') — not in end user's allowed servers", 

102 tool_name, 

103 server_name, 

104 ) 

105 

106 if removed_tools: 

107 verbose_proxy_logger.debug( 

108 "MCP guardrail: removed %s unauthorized MCP tool(s): %s", len(removed_tools), removed_tools 

109 ) 

110 inputs["tools"] = filtered_tools 

111 

112 return inputs 

113 

114 # ------------------------------------------------------------------ 

115 # Private — end user permission resolution 

116 # ------------------------------------------------------------------ 

117 

118 @staticmethod 

119 async def _resolve_end_user_object_permission( 

120 request_data: dict, 

121 ) -> LiteLLM_ObjectPermissionTable | None: 

122 """ 

123 Resolve the end user's object_permission via the cached auth lookup. 

124 

125 Uses get_end_user_object (same path as auth) so no extra DB round-trip 

126 when the cache is warm. 

127 """ 

128 end_user_id: Final = MCPEndUserPermissionGuardrail._get_end_user_id_from_request_data(request_data) 

129 if not end_user_id: 

130 return None 

131 

132 end_user_object: Final = await MCPEndUserPermissionGuardrail._fetch_end_user_object(end_user_id) 

133 return end_user_object.object_permission if end_user_object is not None else None 

134 

135 @staticmethod 

136 def _get_end_user_id_from_request_data(request_data: dict) -> str | None: 

137 return request_data.get("user_api_key_end_user_id") or request_data.get("litellm_metadata", {}).get( 

138 "user_api_key_end_user_id" 

139 ) 

140 

141 @staticmethod 

142 async def _fetch_end_user_object(end_user_id: str): 

143 """ 

144 Fetch end user object via the same cached path used during auth. 

145 No extra DB round-trip when the cache is warm. 

146 """ 

147 from litellm.proxy.auth.auth_checks import get_end_user_object 

148 from litellm.proxy.proxy_server import ( 

149 prisma_client, 

150 proxy_logging_obj, 

151 user_api_key_cache, 

152 ) 

153 

154 if prisma_client is None: 

155 return None 

156 

157 try: 

158 return await get_end_user_object( 

159 end_user_id=end_user_id, 

160 prisma_client=prisma_client, 

161 user_api_key_cache=user_api_key_cache, 

162 parent_otel_span=None, 

163 proxy_logging_obj=proxy_logging_obj, 

164 route="/mcp", 

165 ) 

166 except Exception as e: 

167 verbose_proxy_logger.warning("MCP guardrail: failed to fetch end_user_object for '%s': %s", end_user_id, e) 

168 return None 

169 

170 # ------------------------------------------------------------------ 

171 # Private — permission derivation 

172 # ------------------------------------------------------------------ 

173 

174 @staticmethod 

175 async def _get_allowed_mcp_servers_from_object_permission( 

176 object_permission: LiteLLM_ObjectPermissionTable | None, 

177 ) -> list[str] | None: 

178 """ 

179 Returns: 

180 None — no restrictions configured, allow all MCP servers 

181 list — restrict to exactly these server names 

182 """ 

183 if object_permission is None: 

184 return None 

185 

186 direct_mcp_servers: Final = object_permission.mcp_servers or [] 

187 mcp_access_groups: Final = object_permission.mcp_access_groups or [] 

188 

189 if not direct_mcp_servers and not mcp_access_groups: 

190 return None # Both empty → no restrictions 

191 

192 from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( 

193 MCPRequestHandler, 

194 ) 

195 

196 access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups) 

197 

198 return list(set(direct_mcp_servers + access_group_servers)) 

199 

200 # ------------------------------------------------------------------ 

201 # Config model — exposes this guardrail in the UI 

202 # ------------------------------------------------------------------ 

203 

204 @staticmethod 

205 def get_config_model() -> type["GuardrailConfigModel"] | None: 

206 from litellm.types.proxy.guardrails.guardrail_hooks.mcp_end_user_permission import ( 

207 MCPEndUserPermissionGuardrailConfigModel, 

208 ) 

209 

210 return MCPEndUserPermissionGuardrailConfigModel 

211 

212 @classmethod 

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

214 return [ 

215 GuardrailEventHooks.pre_call, 

216 ] 

217 

218 # ------------------------------------------------------------------ 

219 # Private — tool name extraction 

220 # ------------------------------------------------------------------ 

221 

222 @staticmethod 

223 def _extract_mcp_server_name(tool_name: str) -> str | None: 

224 """ 

225 Split "github-create_issue" → "github". 

226 Returns None if the tool name has no '-' prefix (not an MCP tool). 

227 """ 

228 if not tool_name or "-" not in tool_name: 

229 return None 

230 return tool_name.split("-", 1)[0] 

231 

232 @staticmethod 

233 def _get_tool_name_from_definition(tool: Any) -> str | None: 

234 """ 

235 Extract tool name from a definition dict. 

236 

237 OpenAI format: {"type": "function", "function": {"name": "..."}} 

238 Anthropic format: {"name": "...", "input_schema": {...}} 

239 """ 

240 if not isinstance(tool, dict): 

241 return None 

242 function_def: Final = tool.get("function") 

243 if isinstance(function_def, dict): 

244 name: Final = function_def.get("name") 

245 if name: 

246 return name 

247 return tool.get("name")