mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(panw_prisma_airs): scan tool call args as plain text, not a tool_event
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7a5b98e6a9
commit
dcabd457e5
2 changed files with 67 additions and 116 deletions
|
|
@ -1390,53 +1390,19 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
) -> None:
|
||||
"""Scan tool call arguments with allow/block/mask treatment (in-place modification).
|
||||
|
||||
Each tool call is sent as a ``tool_event`` using the canonical PANW
|
||||
AIRS schema::
|
||||
|
||||
{
|
||||
"metadata": {
|
||||
"ecosystem": "openai",
|
||||
"method": "tools/call",
|
||||
"server_name": "litellm",
|
||||
"tool_invoked": "<function_name>",
|
||||
},
|
||||
"input": "<args_json>", # optional, omitted for empty args
|
||||
}
|
||||
|
||||
Empty-arg invocations are still reported (without ``input``) so AIRS
|
||||
can enforce tool-name-based policies.
|
||||
Arguments go out as plain prompt/response text: the AIRS ``tool_event`` schema
|
||||
only accepts ``ecosystem: "mcp"``, which OpenAI-format tool calls are not.
|
||||
"""
|
||||
for tool_call in tool_calls:
|
||||
# --- extract tool_name and args_text --------------------------
|
||||
tool_name: str | None = None
|
||||
args_text: str | None = None
|
||||
|
||||
if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"):
|
||||
args_text = tool_call.function.arguments
|
||||
tool_name = getattr(tool_call.function, "name", None)
|
||||
elif isinstance(tool_call, dict):
|
||||
func = tool_call.get("function", {})
|
||||
if isinstance(func, dict):
|
||||
args_text = func.get("arguments")
|
||||
tool_name = func.get("name")
|
||||
|
||||
# --- build tool_event payload (canonical PANW schema) -----------
|
||||
tool_event: dict[str, object] = {
|
||||
"metadata": {
|
||||
"ecosystem": "openai",
|
||||
"method": "tools/call",
|
||||
"server_name": "litellm",
|
||||
"tool_invoked": tool_name or "unknown",
|
||||
},
|
||||
}
|
||||
if args_text and args_text.strip():
|
||||
tool_event["input"] = args_text
|
||||
args_text = self._get_tool_call_arguments(tool_call)
|
||||
if not args_text or not args_text.strip():
|
||||
continue
|
||||
|
||||
scan_result = await self._call_panw_api(
|
||||
is_response=False, # tool_event is always request-side in AIRS schema
|
||||
content=args_text,
|
||||
is_response=is_response,
|
||||
metadata=metadata,
|
||||
call_id=call_id,
|
||||
tool_event=tool_event,
|
||||
)
|
||||
|
||||
if scan_result.get("_is_transient") or scan_result.get("_always_block"):
|
||||
|
|
@ -1451,10 +1417,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
continue # fallback_on_error="allow" — leave args unchanged
|
||||
|
||||
action = scan_result.get("action", "block")
|
||||
# Always is_response=False for masked data lookup because
|
||||
# tool_event scans are request-side in AIRS schema and
|
||||
# AIRS returns prompt_masked_data for them.
|
||||
masked_text = self._get_masked_text(scan_result, is_response=False)
|
||||
masked_text = self._get_masked_text(scan_result, is_response=is_response)
|
||||
|
||||
if action == "allow":
|
||||
if masked_text:
|
||||
|
|
@ -1467,6 +1430,17 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
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) -> str | None:
|
||||
"""Read a tool call's function arguments, handling both object and dict forms."""
|
||||
if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"):
|
||||
return tool_call.function.arguments
|
||||
if isinstance(tool_call, dict):
|
||||
func: Final = tool_call.get("function")
|
||||
if isinstance(func, dict):
|
||||
return func.get("arguments")
|
||||
return 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."""
|
||||
|
|
|
|||
|
|
@ -2144,8 +2144,8 @@ class TestPanwAirsToolEventIsResponseFix:
|
|||
"""Tests for Bug A fix: tool_event scans must not set is_response metadata."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_tool_calls_post_call_uses_request_mode_for_tool_event(self):
|
||||
"""_scan_tool_calls_for_guardrail(is_response=True) must call _call_panw_api with is_response=False."""
|
||||
async def test_scan_tool_calls_post_call_scans_args_as_response_text(self):
|
||||
"""_scan_tool_calls_for_guardrail(is_response=True) scans args as response text, never as a tool_event."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_key",
|
||||
|
|
@ -2173,7 +2173,9 @@ class TestPanwAirsToolEventIsResponseFix:
|
|||
start_time=datetime.now(),
|
||||
)
|
||||
mock_api.assert_called_once()
|
||||
assert mock_api.call_args.kwargs.get("is_response") is False
|
||||
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("tool_event") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_panw_api_tool_event_omits_is_response_metadata(self):
|
||||
|
|
@ -2686,8 +2688,8 @@ class TestPanwAirsToolEventPayload:
|
|||
mock_panw_client.client.post.assert_called_once()
|
||||
|
||||
|
||||
class TestPanwAirsToolCallToolEvent:
|
||||
"""Test _scan_tool_calls_for_guardrail sends tool_event payloads."""
|
||||
class TestPanwAirsToolCallContentScan:
|
||||
"""Test _scan_tool_calls_for_guardrail scans arguments as plain prompt/response text."""
|
||||
|
||||
@pytest.fixture
|
||||
def handler(self):
|
||||
|
|
@ -2698,8 +2700,8 @@ class TestPanwAirsToolCallToolEvent:
|
|||
return make_handler(mask_request_content=True)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_event_includes_metadata_and_input(self, handler):
|
||||
"""_scan_tool_calls_for_guardrail sends canonical tool_event with metadata + input."""
|
||||
async def test_tool_call_args_sent_as_prompt_content(self, handler):
|
||||
"""Regression (LIT-5279): args go out as prompt text, not as an ecosystem=openai tool_event."""
|
||||
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
id="call_1",
|
||||
|
|
@ -2725,19 +2727,13 @@ class TestPanwAirsToolCallToolEvent:
|
|||
)
|
||||
|
||||
call_kwargs = mock_api.call_args.kwargs
|
||||
te = call_kwargs["tool_event"]
|
||||
assert_canonical_tool_event(
|
||||
te,
|
||||
ecosystem="openai",
|
||||
server_name="litellm",
|
||||
tool_invoked="get_weather",
|
||||
)
|
||||
# input field carries args
|
||||
assert te["input"] == '{"city": "San Francisco"}'
|
||||
assert call_kwargs["content"] == '{"city": "San Francisco"}'
|
||||
assert call_kwargs["is_response"] is False
|
||||
assert call_kwargs.get("tool_event") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_event_empty_args_omits_input(self, handler):
|
||||
"""Empty args → tool_event has metadata but no input key."""
|
||||
async def test_empty_args_are_not_scanned(self, handler):
|
||||
"""Empty args carry nothing to scan, so no AIRS call is made."""
|
||||
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
id="call_1",
|
||||
|
|
@ -2762,17 +2758,11 @@ class TestPanwAirsToolCallToolEvent:
|
|||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
# Empty args → tool_event still sent for name-based policies
|
||||
mock_api.assert_called_once()
|
||||
te = mock_api.call_args.kwargs["tool_event"]
|
||||
assert_canonical_tool_event(
|
||||
te, ecosystem="openai", server_name="litellm", tool_invoked="list_items"
|
||||
)
|
||||
assert "input" not in te
|
||||
mock_api.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_block_still_raises(self, handler):
|
||||
"""Tool call block with tool_event raises HTTPException(400)."""
|
||||
"""Tool call block raises HTTPException(400)."""
|
||||
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
id="call_1",
|
||||
|
|
@ -2801,8 +2791,8 @@ class TestPanwAirsToolCallToolEvent:
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_mask_with_tool_event(self, handler_mask_request):
|
||||
"""Tool call masking still works with tool_event payloads."""
|
||||
async def test_tool_call_mask_applies_masked_args(self, handler_mask_request):
|
||||
"""Tool call masking still rewrites the arguments in place."""
|
||||
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
id="call_1",
|
||||
|
|
@ -2834,8 +2824,8 @@ class TestPanwAirsToolCallToolEvent:
|
|||
assert tool_call.function.arguments == '{"ssn": "XXXXXXXXXX"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dict_tool_call_extracts_name(self, handler):
|
||||
"""Dict-style tool calls also extract tool_name for tool_event."""
|
||||
async def test_dict_tool_call_extracts_args(self, handler):
|
||||
"""Dict-style tool calls also have their arguments scanned."""
|
||||
|
||||
tool_call = {
|
||||
"function": {
|
||||
|
|
@ -2859,11 +2849,8 @@ class TestPanwAirsToolCallToolEvent:
|
|||
)
|
||||
|
||||
call_kwargs = mock_api.call_args.kwargs
|
||||
te = call_kwargs["tool_event"]
|
||||
assert_canonical_tool_event(
|
||||
te, ecosystem="openai", server_name="litellm", tool_invoked="search"
|
||||
)
|
||||
assert te["input"] == '{"query": "test"}'
|
||||
assert call_kwargs["content"] == '{"query": "test"}'
|
||||
assert call_kwargs.get("tool_event") is None
|
||||
|
||||
|
||||
class TestPanwAirsMcpToolEventScan:
|
||||
|
|
@ -3291,25 +3278,20 @@ class TestPanwAirsDuplicateScanRegression:
|
|||
|
||||
# Expected calls:
|
||||
# 1. text scan for "Hello"
|
||||
# 2. tool_calls scan for get_weather (with tool_event)
|
||||
# 2. tool_calls scan for get_weather (plain prompt text)
|
||||
# 3. MCP scan for file_reader (with tool_event)
|
||||
assert mock_api.call_count == 3
|
||||
|
||||
# Verify ordering: first is text (no tool_event), second is tool_call, third is MCP
|
||||
# Verify ordering: first is text, second is tool_call args, third is MCP
|
||||
calls = mock_api.call_args_list
|
||||
|
||||
# First call: text scan (content="Hello", no tool_event)
|
||||
assert calls[0].kwargs.get("content") == "Hello"
|
||||
assert calls[0].kwargs.get("tool_event") is None
|
||||
|
||||
# Second call: tool_calls scan (tool_event with get_weather)
|
||||
assert (
|
||||
calls[1].kwargs["tool_event"]["metadata"]["tool_invoked"]
|
||||
== "get_weather"
|
||||
)
|
||||
assert calls[1].kwargs["tool_event"]["metadata"]["ecosystem"] == "openai"
|
||||
assert calls[1].kwargs["tool_event"]["metadata"]["method"] == "tools/call"
|
||||
assert "tool_name" not in calls[1].kwargs["tool_event"]
|
||||
# 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"}'
|
||||
|
||||
# Third call: MCP scan (tool_event with file_reader)
|
||||
assert (
|
||||
|
|
@ -3839,11 +3821,11 @@ class TestPanwAirsDeveloperRoleGuardrail:
|
|||
|
||||
|
||||
class TestPanwAirsEmptyToolArgsBlock:
|
||||
"""Test empty-arg tool call blocking by name policy."""
|
||||
"""Test empty-arg tool call handling."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_empty_args_block_by_name_policy(self):
|
||||
"""Empty-args tool call where PANW returns block raises HTTPException."""
|
||||
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."""
|
||||
|
||||
handler = make_handler()
|
||||
|
||||
|
|
@ -3861,17 +3843,16 @@ class TestPanwAirsEmptyToolArgsBlock:
|
|||
) as mock_api:
|
||||
mock_api.return_value = {"action": "block", "category": "dangerous"}
|
||||
|
||||
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(),
|
||||
)
|
||||
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(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_api.assert_not_called()
|
||||
|
||||
|
||||
class TestPanwAirsDictChunkStreaming:
|
||||
|
|
@ -4149,13 +4130,10 @@ class TestPanwAirsUnifiedToolsScan:
|
|||
# Exactly 1 API call: the tool_call invocation, not the definitions
|
||||
assert mock_api.call_count == 1
|
||||
|
||||
te = mock_api.call_args.kwargs["tool_event"]
|
||||
# Must carry the exact function name — not "unknown"
|
||||
assert te["metadata"]["tool_invoked"] == "get_weather"
|
||||
# Must NOT carry definition-shaped keys
|
||||
assert "type" not in te
|
||||
assert "server_label" not in te
|
||||
assert "server_url" not in te
|
||||
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.get("tool_event") is None
|
||||
|
||||
|
||||
class TestPanwAirsMcpRestToolInvoked:
|
||||
|
|
@ -5239,16 +5217,16 @@ class TestPanwAirsMcpMasking:
|
|||
|
||||
|
||||
class TestPanwAirsResponseToolCallMasking:
|
||||
"""Tests for response-side tool-call masking using prompt_masked_data."""
|
||||
"""Tests for response-side tool-call masking using response_masked_data."""
|
||||
|
||||
@pytest.fixture
|
||||
def handler(self):
|
||||
return make_handler(mask_response_content=True)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_side_tool_call_uses_prompt_masked_data(self, handler):
|
||||
"""_scan_tool_calls_for_guardrail(is_response=True) should look up
|
||||
prompt_masked_data (not response_masked_data) and mask instead of blocking."""
|
||||
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"}'
|
||||
|
|
@ -5260,8 +5238,7 @@ class TestPanwAirsResponseToolCallMasking:
|
|||
mock_api.return_value = {
|
||||
"action": "block",
|
||||
"category": "dlp",
|
||||
# AIRS returns prompt_masked_data for tool_event scans
|
||||
"prompt_masked_data": {"data": '{"query": "****"}'},
|
||||
"response_masked_data": {"data": '{"query": "****"}'},
|
||||
}
|
||||
|
||||
await handler._scan_tool_calls_for_guardrail(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue