mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(guardrails): enable explicit PANW MCP output scanning (#43109)
* fix(guardrails): declare post_mcp_call for PANW Prisma AIRS Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): exercise post_mcp_call_hook dispatch in PANW tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-catalog): add fal_ai resolution-tiered image cost fields Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): keep post_mcp_call opt-in for PANW Prisma AIRS Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-catalog): add fal_ai resolution-tiered image cost keys to cost map schema Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: joshua <joshua@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
1bfa3d4fa6
commit
fb74957ddd
3 changed files with 121 additions and 1 deletions
|
|
@ -2012,4 +2012,5 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
GuardrailEventHooks.logging_only,
|
||||
GuardrailEventHooks.pre_mcp_call,
|
||||
GuardrailEventHooks.during_mcp_call,
|
||||
GuardrailEventHooks.post_mcp_call,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue