diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index cd538ad8c8d..e4822195bec 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -2012,4 +2012,5 @@ class PanwPrismaAirsHandler(CustomGuardrail): GuardrailEventHooks.logging_only, GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.during_mcp_call, + GuardrailEventHooks.post_mcp_call, ] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 1f52fa224ee..5db3e11ac06 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -21,7 +21,9 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from fastapi import HTTPException +from mcp.types import CallToolResult, TextContent +import litellm from litellm.caching import DualCache from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import UserAPIKeyAuth @@ -29,6 +31,7 @@ from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( PanwPrismaAirsHandler, initialize_guardrail, ) +from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks, LitellmParams from litellm.types.utils import ( ChatCompletionCustomToolCallPayload, @@ -2025,6 +2028,30 @@ class TestPanwAirsShouldRunGuardrail: True, id="explicit_pre_mcp_call_mode", ), + pytest.param( + True, + "post_call", + _simple_data(), + GuardrailEventHooks.post_mcp_call, + False, + id="post_call_mode_does_not_run_for_post_mcp_call", + ), + pytest.param( + True, + "post_mcp_call", + _simple_data(), + GuardrailEventHooks.post_mcp_call, + True, + id="explicit_post_mcp_call_mode", + ), + pytest.param( + True, + "post_mcp_call", + _simple_data(), + GuardrailEventHooks.post_call, + False, + id="post_mcp_call_mode_does_not_run_for_regular_post_call", + ), pytest.param( True, "pre_call", @@ -2048,6 +2075,66 @@ class TestPanwAirsShouldRunGuardrail: assert handler.should_run_guardrail(data, query_event) is expected +class TestPanwAirsPostMcpCall: + """Explicit MCP output scans use the existing AIRS response contract.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize("action", ["allow", "block", "mask"]) + async def test_post_mcp_call_scans_tool_result(self, monkeypatch: pytest.MonkeyPatch, action: str) -> None: + original: Final = "ssn 123-45-6789" + masked: Final = "ssn ***********" + + def respond(request: httpx.Request) -> httpx.Response: + payload: Final = json.loads(request.content) + assert request.url.path.endswith("/v1/scan/sync/request") + assert payload["contents"] == [{"response": original}] + assert payload["ai_profile"] == {"profile_name": "test_profile"} + return httpx.Response( + 200, + json={ + "action": "block" if action == "block" else "allow", + "category": "malicious" if action == "block" else "benign", + "scan_id": "s1", + "report_id": "r1", + "profile_name": "test_profile", + **({"response_masked_data": {"data": masked}} if action == "mask" else {}), + }, + ) + + transport_handler: Final = MagicMock(side_effect=respond) + http_client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(transport_handler)) + handler: Final = make_handler( + event_hook="post_mcp_call", + default_on=True, + mask_response_content=True, + http_client=http_client, + ) + monkeypatch.setattr(litellm, "callbacks", [handler]) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + result: Final = CallToolResult(content=[TextContent(type="text", text=original)], isError=False) + try: + if action == "block": + with pytest.raises(HTTPException) as exc_info: + await proxy_logging.post_mcp_call_hook( + response=result, + request_data={"litellm_call_id": "c1"}, + user_api_key_dict=None, + ) + assert exc_info.value.status_code == 400 + transport_handler.assert_called_once() + return + returned: Final = await proxy_logging.post_mcp_call_hook( + response=result, + request_data={"litellm_call_id": "c1"}, + user_api_key_dict=None, + ) + transport_handler.assert_called_once() + assert returned.model_dump(by_alias=True)["isError"] is False + assert returned.content == [TextContent(type="text", text=masked if action == "mask" else original)] + finally: + await http_client.client.aclose() + + class TestPanwAirsToolEventIsResponseFix: """Tests for Bug A fix: tool_event scans must not set is_response metadata.""" diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index 79d91db902c..fcd7e537937 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -1,5 +1,5 @@ import json -from typing import Literal +from typing import Final, Literal from unittest.mock import MagicMock, patch import pytest @@ -11,6 +11,38 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import Mode, SupportedGuardrailIntegrations +def test_init_guardrails_v2_registers_panw_mcp_output_scanner(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.proxy.guardrails import guardrail_registry + from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import PanwPrismaAirsHandler + from litellm.types.guardrails import GuardrailEventHooks + + monkeypatch.setenv("LITELLM_STRICT_GUARDRAIL_MODES", "true") + monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", InMemoryGuardrailHandler()) + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "panw-mcp-output", + "litellm_params": { + "guardrail": "panw_prisma_airs", + "mode": "post_mcp_call", + "default_on": True, + "api_key": "test-panw-key", + "profile_name": "test-profile", + }, + } + ] + ) + scanners: Final = tuple( + callback + for callback in litellm.callbacks + if isinstance(callback, PanwPrismaAirsHandler) and callback.guardrail_name == "panw-mcp-output" + ) + assert len(scanners) == 1, "PANW MCP output scanning must be registered at startup" + assert scanners[0].should_run_guardrail({}, GuardrailEventHooks.post_mcp_call) is True + assert scanners[0].should_run_guardrail({}, GuardrailEventHooks.post_call) is False + + def test_initialize_presidio_guardrail(): """ Test that initialize_guardrail correctly uses registered initializers