fix(panw_prisma_airs): scan tool names with args and tolerate custom tool calls

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-15 18:32:04 +00:00
parent 380c449a62
commit a477c3e841
2 changed files with 148 additions and 51 deletions

View file

@ -10,11 +10,12 @@ import os
import re
from collections.abc import AsyncIterable, Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias
from urllib.parse import urlparse
import httpx
from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -38,6 +39,9 @@ from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
CallTypes,
CallTypesLiteral,
ChatCompletionDeltaCustomToolCall,
ChatCompletionDeltaToolCall,
ChatCompletionMessageCustomToolCall,
ChatCompletionMessageToolCall,
ChatCompletionToolCallChunk,
Choices,
@ -46,6 +50,28 @@ from litellm.types.utils import (
ModelResponseStream,
)
ToolCallLike: TypeAlias = (
ChatCompletionMessageToolCall
| ChatCompletionDeltaToolCall
| ChatCompletionMessageCustomToolCall
| ChatCompletionDeltaCustomToolCall
| ChatCompletionToolCallChunk
)
class _ToolCallFunctionSlice(BaseModel):
model_config = ConfigDict(from_attributes=True, extra="ignore")
name: str | None = None
arguments: str | None = None
class _ToolCallSlice(BaseModel):
model_config = ConfigDict(from_attributes=True, extra="ignore")
function: _ToolCallFunctionSlice | None = None
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -1390,18 +1416,22 @@ class PanwPrismaAirsHandler(CustomGuardrail):
request_data: dict,
start_time: datetime,
) -> None:
"""Scan tool call arguments with allow/block/mask treatment (in-place modification).
"""Scan tool calls with allow/block/mask treatment (in-place modification).
Arguments go out as plain prompt/response text: the AIRS ``tool_event`` schema
only accepts ``ecosystem: "mcp"``, which OpenAI-format tool calls are not.
Tool name and arguments go out as plain prompt/response text, newline separated:
the AIRS ``tool_event`` schema only accepts ``ecosystem: "mcp"``, which
OpenAI-format tool calls are not. A name-only call is still scanned so
tool-name policies keep firing on empty arguments.
"""
for tool_call in tool_calls:
args_text = self._get_tool_call_arguments(tool_call)
if not args_text or not args_text.strip():
tool_name, args_text = self._get_tool_call_function(tool_call)
scanned_args = args_text if args_text and args_text.strip() else None
scan_text = "\n".join(part for part in (tool_name, scanned_args) if part)
if not scan_text.strip():
continue
scan_result = await self._call_panw_api(
content=args_text,
content=scan_text,
is_response=is_response,
metadata=metadata,
call_id=call_id,
@ -1419,36 +1449,58 @@ class PanwPrismaAirsHandler(CustomGuardrail):
continue
action = scan_result.get("action", "block")
masked_text = self._get_masked_text(scan_result, is_response=is_response)
masked_args = self._masked_tool_call_arguments(
self._get_masked_text(scan_result, is_response=is_response),
scanned_name=bool(tool_name),
scanned_args=scanned_args,
)
if action == "allow":
if masked_text:
self._set_tool_call_arguments(tool_call, masked_text)
elif masked_text and (
if masked_args:
self._set_tool_call_arguments(tool_call, masked_args)
elif masked_args and (
(is_response and self.mask_response_content) or (not is_response and self.mask_request_content)
):
self._set_tool_call_arguments(tool_call, masked_text)
self._set_tool_call_arguments(tool_call, masked_args)
else:
error_detail = self._build_error_detail(scan_result, is_response=is_response)
raise HTTPException(status_code=400, detail=error_detail)
@staticmethod
def _get_tool_call_arguments(
tool_call: ChatCompletionMessageToolCall | ChatCompletionToolCallChunk,
def _masked_tool_call_arguments(
masked_text: str | None,
*,
scanned_name: bool,
scanned_args: str | None,
) -> str | None:
"""Read a tool call's function arguments, handling both object and dict forms."""
if isinstance(tool_call, dict):
func: Final = tool_call.get("function")
return func.get("arguments") if isinstance(func, dict) else None
return tool_call.function.arguments
"""Recover the arguments slice of a masked scan, or None when it cannot be applied."""
if masked_text is None or scanned_args is None:
return None
if not scanned_name:
return masked_text
_, separator, masked_args = masked_text.partition("\n")
return masked_args if separator else None
@staticmethod
def _set_tool_call_arguments(tool_call, masked_text: str) -> None:
"""Set masked text on a tool call's function arguments, handling both object and dict forms."""
if hasattr(tool_call, "function"):
tool_call.function.arguments = masked_text
elif isinstance(tool_call, dict) and isinstance(tool_call.get("function"), dict):
def _get_tool_call_function(tool_call: ToolCallLike) -> tuple[str | None, str | None]:
"""Read a tool call's function name and arguments; (None, None) for non-function shapes."""
try:
parsed: Final = _ToolCallSlice.model_validate(tool_call, from_attributes=True)
except ValidationError:
return (None, None)
if parsed.function is None:
return (None, None)
return (parsed.function.name, parsed.function.arguments)
@staticmethod
def _set_tool_call_arguments(tool_call: ToolCallLike, masked_text: str) -> None:
"""Set masked text on the function arguments of a call that _get_tool_call_function accepted."""
if isinstance(tool_call, dict):
tool_call["function"]["arguments"] = masked_text
return
if isinstance(tool_call, ChatCompletionMessageCustomToolCall | ChatCompletionDeltaCustomToolCall):
return
tool_call.function.arguments = masked_text
@staticmethod
def _is_anthropic_request(

View file

@ -26,6 +26,8 @@ from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import (
)
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
from litellm.types.utils import (
ChatCompletionCustomToolCallPayload,
ChatCompletionMessageCustomToolCall,
ChatCompletionMessageToolCall,
Choices,
Delta,
@ -1811,7 +1813,7 @@ class TestPanwAirsApplyGuardrail:
mock_api.return_value = {
"action": "block",
"category": "dlp",
"prompt_masked_data": {"data": '{"ssn": "XXXXXXXXXX"}'},
"prompt_masked_data": {"data": 'get_user\n{"ssn": "XXXXXXXXXX"}'},
}
await handler_mask_request.apply_guardrail(
@ -2174,7 +2176,7 @@ class TestPanwAirsToolEventIsResponseFix:
)
mock_api.assert_called_once()
assert mock_api.call_args.kwargs.get("is_response") is True
assert mock_api.call_args.kwargs.get("content") == '{"city": "Paris"}'
assert mock_api.call_args.kwargs.get("content") == 'get_weather\n{"city": "Paris"}'
assert mock_api.call_args.kwargs.get("tool_event") is None
@pytest.mark.asyncio
@ -2727,13 +2729,13 @@ class TestPanwAirsToolCallContentScan:
)
call_kwargs = mock_api.call_args.kwargs
assert call_kwargs["content"] == '{"city": "San Francisco"}'
assert call_kwargs["content"] == 'get_weather\n{"city": "San Francisco"}'
assert call_kwargs["is_response"] is False
assert call_kwargs.get("tool_event") is None
@pytest.mark.asyncio
async def test_empty_args_are_not_scanned(self, handler):
"""Empty args carry nothing to scan, so no AIRS call is made."""
async def test_empty_args_still_scan_the_tool_name(self, handler):
"""A name-only call is still scanned so tool-name policies keep firing."""
tool_call = ChatCompletionMessageToolCall(
id="call_1",
@ -2758,6 +2760,46 @@ class TestPanwAirsToolCallContentScan:
start_time=datetime.now(),
)
assert mock_api.call_args.kwargs["content"] == "list_items"
@pytest.mark.asyncio
async def test_custom_tool_call_is_skipped(self, handler):
"""Custom tool calls carry no function payload, so they are skipped instead of crashing."""
tool_call = ChatCompletionMessageCustomToolCall(
id="call_1",
type="custom",
custom=ChatCompletionCustomToolCallPayload(name="run_sql", input="select 1"),
)
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
await handler._scan_tool_calls_for_guardrail(
tool_calls=[tool_call],
is_response=False,
metadata={"user": "test", "model": "gpt-4"},
call_id="test-call-id",
request_data={"litellm_call_id": "test-call-id"},
start_time=datetime.now(),
)
mock_api.assert_not_called()
@pytest.mark.asyncio
async def test_tool_call_without_function_is_skipped(self, handler):
"""A tool call with no function payload is skipped instead of raising AttributeError."""
tool_call = ChatCompletionMessageToolCall(id="call_1", type="function", function=None)
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
await handler._scan_tool_calls_for_guardrail(
tool_calls=[tool_call],
is_response=False,
metadata={"user": "test", "model": "gpt-4"},
call_id="test-call-id",
request_data={"litellm_call_id": "test-call-id"},
start_time=datetime.now(),
)
mock_api.assert_not_called()
@pytest.mark.asyncio
@ -2809,7 +2851,7 @@ class TestPanwAirsToolCallContentScan:
mock_api.return_value = {
"action": "block",
"category": "dlp",
"prompt_masked_data": {"data": '{"ssn": "XXXXXXXXXX"}'},
"prompt_masked_data": {"data": 'get_user\n{"ssn": "XXXXXXXXXX"}'},
}
await handler_mask_request._scan_tool_calls_for_guardrail(
@ -2849,7 +2891,7 @@ class TestPanwAirsToolCallContentScan:
)
call_kwargs = mock_api.call_args.kwargs
assert call_kwargs["content"] == '{"query": "test"}'
assert call_kwargs["content"] == 'search\n{"query": "test"}'
assert call_kwargs.get("tool_event") is None
@pytest.mark.asyncio
@ -2869,7 +2911,7 @@ class TestPanwAirsToolCallContentScan:
mock_api.return_value = {
"action": "allow",
"category": "dlp",
"prompt_masked_data": {"data": '{"ssn": "XXXXXXXXXX"}'},
"prompt_masked_data": {"data": 'get_user\n{"ssn": "XXXXXXXXXX"}'},
}
await handler._scan_tool_calls_for_guardrail(
@ -2893,7 +2935,7 @@ class TestPanwAirsToolCallContentScan:
mock_api.return_value = {
"action": "block",
"category": "dlp",
"prompt_masked_data": {"data": '{"ssn": "XXXXXXXXXX"}'},
"prompt_masked_data": {"data": 'get_user\n{"ssn": "XXXXXXXXXX"}'},
}
await handler_mask_request._scan_tool_calls_for_guardrail(
@ -3411,7 +3453,7 @@ class TestPanwAirsDuplicateScanRegression:
# Second call: tool_calls scan (args as prompt text, no tool_event)
assert calls[1].kwargs.get("tool_event") is None
assert calls[1].kwargs["content"] == '{"city": "NYC"}'
assert calls[1].kwargs["content"] == 'get_weather\n{"city": "NYC"}'
# Third call: MCP scan (tool_event with file_reader)
assert (
@ -3944,8 +3986,8 @@ class TestPanwAirsEmptyToolArgsBlock:
"""Test empty-arg tool call handling."""
@pytest.mark.asyncio
async def test_tool_call_empty_args_not_scanned(self):
"""Empty-args tool call has no text to scan, so no AIRS call and no block."""
async def test_tool_call_empty_args_block_by_name_policy(self):
"""An empty-args call is still scanned by name, so a name policy can block it."""
handler = make_handler()
@ -3963,16 +4005,18 @@ class TestPanwAirsEmptyToolArgsBlock:
) as mock_api:
mock_api.return_value = {"action": "block", "category": "dangerous"}
await handler._scan_tool_calls_for_guardrail(
tool_calls=[tool_call],
is_response=False,
metadata={"user": "test", "model": "gpt-4"},
call_id="test-call-id",
request_data={"litellm_call_id": "test-call-id"},
start_time=datetime.now(),
)
with pytest.raises(HTTPException) as exc_info:
await handler._scan_tool_calls_for_guardrail(
tool_calls=[tool_call],
is_response=False,
metadata={"user": "test", "model": "gpt-4"},
call_id="test-call-id",
request_data={"litellm_call_id": "test-call-id"},
start_time=datetime.now(),
)
mock_api.assert_not_called()
assert exc_info.value.status_code == 400
assert mock_api.call_args.kwargs["content"] == "dangerous_tool"
class TestPanwAirsDictChunkStreaming:
@ -4252,7 +4296,7 @@ class TestPanwAirsUnifiedToolsScan:
call_kwargs = mock_api.call_args.kwargs
# Must carry the invocation arguments, not definition-shaped payloads
assert call_kwargs["content"] == '{"location": "NYC"}'
assert call_kwargs["content"] == 'get_weather\n{"location": "NYC"}'
assert call_kwargs.get("tool_event") is None
@ -5347,10 +5391,11 @@ class TestPanwAirsResponseToolCallMasking:
async def test_response_side_tool_call_uses_response_masked_data(self, handler):
"""_scan_tool_calls_for_guardrail(is_response=True) scans args as response text,
so masked output comes from response_masked_data and masks instead of blocking."""
tool_call = MagicMock()
tool_call.function = MagicMock()
tool_call.function.arguments = '{"query": "sensitive-data"}'
tool_call.function.name = "search"
tool_call = ChatCompletionMessageToolCall(
id="call_1",
type="function",
function=Function(name="search", arguments='{"query": "sensitive-data"}'),
)
with patch.object(
handler, "_call_panw_api", new_callable=AsyncMock
@ -5358,7 +5403,7 @@ class TestPanwAirsResponseToolCallMasking:
mock_api.return_value = {
"action": "block",
"category": "dlp",
"response_masked_data": {"data": '{"query": "****"}'},
"response_masked_data": {"data": 'search\n{"query": "****"}'},
}
await handler._scan_tool_calls_for_guardrail(