From 5da6ec035b0aec0ed2b6845a77676d02c8310549 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 7 Apr 2026 07:55:37 -0400 Subject: [PATCH] add test --- .../guardrail_hooks/tool_permission.py | 40 +++++++++++ .../guardrail_hooks/test_tool_permission.py | 71 +++++++++++++++++++ 2 files changed, 111 insertions(+) diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 79b5bd7adbe..a67eaacf875 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -537,6 +537,9 @@ class ToolPermissionGuardrail(CustomGuardrail): "detection_message": message, }, ) + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) return data new_tools: Optional[List[ChatCompletionToolParam]] = data.get("tools") @@ -580,6 +583,43 @@ class ToolPermissionGuardrail(CustomGuardrail): ) return data + @log_guardrail_information + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: CallTypesLiteral, + ) -> None: + """ + Enforce tool permission rules during MCP tool execution (during_mcp_call). + + The data dict is the synthetic MCP payload built by _convert_mcp_to_llm_format, + which carries mcp_tool_name as the namespaced '{server}-{tool}' string. + """ + if self.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.during_mcp_call + ) is not True: + return + + mcp_tool_name: Optional[str] = data.get("mcp_tool_name") + if mcp_tool_name is None: + return + + is_allowed, _, message = self._check_tool_permission(mcp_tool_name) + if not is_allowed and message is not None: + verbose_proxy_logger.warning(f"Tool Permission Guardrail: {message}") + raise HTTPException( + status_code=400, + detail={ + "error": "Violated guardrail policy", + "detection_message": message, + }, + ) + + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) + @log_guardrail_information async def async_post_call_success_hook( self, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py index 03f914fccc3..8d9d436ce7d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py @@ -643,6 +643,23 @@ class TestToolPermissionGuardrailMCPPreCall: assert isinstance(result, dict) assert "tools" in result + @pytest.mark.asyncio + async def test_mcp_allowed_tool_sets_applied_guardrails_header(self): + """Allowed MCP tool must call add_guardrail_to_applied_guardrails_header + so the guardrail appears in x-litellm-applied-guardrails.""" + guardrail = self._make_guardrail(decision="allow", pattern=r"exa-.*") + data = {"mcp_tool_name": "exa-web_search_exa"} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=DualCache(), + data=data, + call_type="call_mcp_tool", + ) + + applied = (data.get("metadata") or {}).get("applied_guardrails", []) + assert "test-mcp-guardrail" in applied + def test_supported_event_hooks_includes_mcp(self): """ToolPermissionGuardrail must declare pre_mcp_call and during_mcp_call so that mode=['pre_mcp_call','during_mcp_call'] passes _validate_event_hook.""" @@ -650,6 +667,60 @@ class TestToolPermissionGuardrailMCPPreCall: assert GuardrailEventHooks.pre_mcp_call in (guardrail.supported_event_hooks or []) assert GuardrailEventHooks.during_mcp_call in (guardrail.supported_event_hooks or []) + @pytest.mark.asyncio + async def test_during_mcp_call_denied_raises_http_exception(self): + """during_mcp_call: denied tool raises HTTPException(400) via async_moderation_hook.""" + guardrail = self._make_guardrail(decision="deny", pattern=r"exa-.*") + data = {"mcp_tool_name": "exa-web_search_exa", "mcp_arguments": {"query": "test"}} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_moderation_hook( + data=data, + user_api_key_dict=UserAPIKeyAuth(), + call_type="call_mcp_tool", + ) + assert exc_info.value.status_code == 400 + assert "Violated guardrail policy" in exc_info.value.detail["error"] + + @pytest.mark.asyncio + async def test_during_mcp_call_allowed_passes(self): + """during_mcp_call: allowed tool does not raise.""" + guardrail = self._make_guardrail(decision="allow", pattern=r"exa-.*") + data = {"mcp_tool_name": "exa-web_search_exa", "mcp_arguments": {"query": "test"}} + + await guardrail.async_moderation_hook( + data=data, + user_api_key_dict=UserAPIKeyAuth(), + call_type="call_mcp_tool", + ) + + @pytest.mark.asyncio + async def test_during_mcp_call_allowed_sets_applied_guardrails_header(self): + """during_mcp_call: allowed tool must appear in applied-guardrails header.""" + guardrail = self._make_guardrail(decision="allow", pattern=r"exa-.*") + data = {"mcp_tool_name": "exa-web_search_exa"} + + await guardrail.async_moderation_hook( + data=data, + user_api_key_dict=UserAPIKeyAuth(), + call_type="call_mcp_tool", + ) + + applied = (data.get("metadata") or {}).get("applied_guardrails", []) + assert "test-mcp-guardrail" in applied + + @pytest.mark.asyncio + async def test_during_mcp_call_no_mcp_tool_name_is_noop(self): + """during_mcp_call: missing mcp_tool_name is a no-op (does not raise).""" + guardrail = self._make_guardrail(decision="deny", pattern=r".*") + data = {"mcp_arguments": {"query": "test"}} + + await guardrail.async_moderation_hook( + data=data, + user_api_key_dict=UserAPIKeyAuth(), + call_type="call_mcp_tool", + ) + class TestToolPermissionGuardrailIntegration: """Integration tests for Tool Permission Guardrail"""