mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
380c449a62
commit
a477c3e841
2 changed files with 148 additions and 51 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue