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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2MCP End User Permission Guardrail Hook
4Enforces end user permissions for MCP server access via apply_guardrail:
5- input_type="request" → filter tools the end user cannot access
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"""
13from typing import TYPE_CHECKING, Any, Final, Literal
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
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
28GUARDRAIL_NAME: Final = "mcp_end_user_permission"
31class MCPEndUserPermissionGuardrail(CustomGuardrail):
32 """
33 Guardrail that enforces end user permissions for MCP server access.
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.
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 """
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")
48 # ------------------------------------------------------------------
49 # apply_guardrail — filters MCP tools on the request side only
50 # ------------------------------------------------------------------
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)
67 # ------------------------------------------------------------------
68 # Private — request-side tool filtering
69 # ------------------------------------------------------------------
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
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
84 verbose_proxy_logger.debug("MCP guardrail: end user restricted to MCP servers: %s", allowed_mcp_servers)
86 filtered_tools: Final = []
87 removed_tools: Final = []
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
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 )
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
112 return inputs
114 # ------------------------------------------------------------------
115 # Private — end user permission resolution
116 # ------------------------------------------------------------------
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.
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
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
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 )
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 )
154 if prisma_client is None:
155 return None
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
170 # ------------------------------------------------------------------
171 # Private — permission derivation
172 # ------------------------------------------------------------------
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
186 direct_mcp_servers: Final = object_permission.mcp_servers or []
187 mcp_access_groups: Final = object_permission.mcp_access_groups or []
189 if not direct_mcp_servers and not mcp_access_groups:
190 return None # Both empty → no restrictions
192 from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
193 MCPRequestHandler,
194 )
196 access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups)
198 return list(set(direct_mcp_servers + access_group_servers))
200 # ------------------------------------------------------------------
201 # Config model — exposes this guardrail in the UI
202 # ------------------------------------------------------------------
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 )
210 return MCPEndUserPermissionGuardrailConfigModel
212 @classmethod
213 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
214 return [
215 GuardrailEventHooks.pre_call,
216 ]
218 # ------------------------------------------------------------------
219 # Private — tool name extraction
220 # ------------------------------------------------------------------
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]
232 @staticmethod
233 def _get_tool_name_from_definition(tool: Any) -> str | None:
234 """
235 Extract tool name from a definition dict.
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")