From a477c3e8415718b98a72dd1d900defa8fdc309d1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 15 Aug 2026 18:32:04 +0000 Subject: [PATCH] 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> --- .../panw_prisma_airs/panw_prisma_airs.py | 100 +++++++++++++----- .../guardrail_hooks/test_panw_prisma_airs.py | 99 ++++++++++++----- 2 files changed, 148 insertions(+), 51 deletions(-) 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 bd1f440fd36..51b00af4ed0 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 @@ -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( 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 6b28ccc54c5..0e27c53a9fa 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 @@ -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(