From 317430db4ec5d0bfd844632220ab51d8e04e2348 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:16:18 +0000 Subject: [PATCH] fix(panw_prisma_airs): honor experimental_use_latest_role_message_only on every request shape (#42447) * fix(panw_prisma_airs): apply experimental_use_latest_role_message_only to every request shape Explicit true/false now applies to chat completions, Anthropic /v1/messages and /v1/responses alike; unset keeps latest-only for Anthropic and full history otherwise. Text indices are mapped back to their source message by value instead of by count, so Responses instructions, function_call_output and reasoning items no longer derail the alignment and silently rescan the whole history Co-authored-by: scthornton Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(panw_prisma_airs): type latest-message helpers against AllMessageValues Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(panw_prisma_airs): require forward and reverse text attribution to agree A Responses function_call_output whose text equals the latest user turn could claim that turn's slot in a forward-only walk and demote the latest-only scan to an earlier message. Walk both directions and fall back to the full role-filter scan when they disagree Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(panw_prisma_airs): pick the latest human turn from messages, not from aligned texts An image-only latest user turn no longer promotes an earlier user turn into the latest-only scan; it scans nothing on the request side, as the Anthropic path did before. A latest user/developer message whose text never reached texts (a trailing Responses reasoning item) falls back to the role-filter scan instead of narrowing. Types the test helpers, drops the narrating docstrings and adds regressions for both shapes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(panw_prisma_airs): log when latest-only selection leaves nothing to scan An image-only latest user turn with experimental_use_latest_role_message_only=true intentionally yields zero scanner calls. Emit a debug line naming the call_id so operators can tell this apart from the guardrail not firing. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(panw_prisma_airs): keep Responses reasoning items out of latest-turn selection The Responses translation handler gives reasoning input items the default user role, so a reasoning item with text content after the latest prompt was picked as the latest human turn and the real prompt went unscanned under experimental_use_latest_role_message_only. Map reasoning items back to their texts positions from the raw input and exclude them; fall back to the role-filter scan when the raw items do not account for every text Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: scthornton Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../panw_prisma_airs/panw_prisma_airs.py | 268 +++--- .../guardrail_hooks/panw_prisma_airs.py | 9 +- .../observability/test_guardrail_effects.py | 108 +++ .../guardrail_hooks/test_panw_prisma_airs.py | 803 +++++++++--------- 4 files changed, 669 insertions(+), 519 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 df5a265bb72..cd538ad8c8d 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 @@ -5,6 +5,8 @@ Palo Alto Networks Prisma AI Runtime Security (AIRS) Guardrail Integration for L Provides real-time threat detection, DLP, URL filtering, content masking, and policy enforcement for AI applications. """ +import functools +import itertools import json import os import re @@ -15,7 +17,7 @@ from urllib.parse import urlparse import httpx from fastapi import HTTPException -from pydantic import BaseModel, ConfigDict, ValidationError, field_validator +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError, field_validator from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -38,6 +40,7 @@ from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, ) from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ( CallTypes, CallTypesLiteral, @@ -90,6 +93,32 @@ class _ToolCallSlice(BaseModel): function: _ToolCallFunctionSlice | None = None +class _ResponsesContentPart(BaseModel): + model_config = ConfigDict(extra="ignore") + + text: str | None = None + + +class _ResponsesInputItem(BaseModel): + """The slice of a raw Responses ``input`` item that decides which ``texts`` it flattens to.""" + + model_config = ConfigDict(extra="ignore") + + type: str | None = None + content: str | tuple[_ResponsesContentPart, ...] | None = None + + def text_count(self) -> int: + if isinstance(self.content, str): + return 1 + if self.content is None: + return 0 + return sum(part.text is not None for part in self.content) + + +_ResponsesInput: TypeAlias = str | tuple[_ResponsesInputItem, ...] | None +_RESPONSES_INPUT: Final[TypeAdapter[_ResponsesInput]] = TypeAdapter(_ResponsesInput) + + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -194,7 +223,6 @@ class PanwPrismaAirsHandler(CustomGuardrail): # internal '<=' comparison and surfaces as a misleading api_error. self.timeout = float(timeout) if timeout is not None else 10.0 - # Tri-state: None = not set (default-on for Anthropic), True = explicit on, False = explicit off self.experimental_use_latest_role_message_only: bool | None = kwargs.get( "experimental_use_latest_role_message_only" ) @@ -1578,119 +1606,149 @@ class PanwPrismaAirsHandler(CustomGuardrail): request_data: Mapping[str, object], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> bool: - """Resolve whether to scan only the latest user message. + """Resolve whether to scan only the latest user/developer message. - - Non-Anthropic requests: always False (existing behavior) - - Anthropic requests: - - Flag explicitly True/False: respect it - - Flag None (not set): default to True + - Flag explicitly True/False: respect it for every request shape, + matching the bedrock guardrail's semantics for the same flag + - Flag None (not set): True for Anthropic /v1/messages requests, False otherwise """ - if not self._is_anthropic_request(request_data, logging_obj): - return False - if self.experimental_use_latest_role_message_only is None: - return True # Default-on for Anthropic - return self.experimental_use_latest_role_message_only + if self.experimental_use_latest_role_message_only is not None: + return self.experimental_use_latest_role_message_only + return self._is_anthropic_request(request_data, logging_obj) @staticmethod - def _get_latest_user_text_indices( + def _message_texts(message: AllMessageValues) -> tuple[str, ...]: + """Text entries the framework flattens out of one structured message.""" + content: Final = message.get("content") + if isinstance(content, str): + return (content,) + if not isinstance(content, list): + return () + return tuple(text for item in content if isinstance(item, dict) and isinstance(text := item.get("text"), str)) + + @classmethod + def _text_source_message_indices( + cls, texts: Sequence[str], - messages: Sequence[object], - ) -> set | None: + messages: Sequence[AllMessageValues], + ) -> tuple[int, ...] | None: + """Map every ``texts`` entry to the index of the structured message it was flattened from. + + A message's texts are consumed only when they sit at the running position of + ``texts``; messages the translation handler added without a counterpart in + ``texts`` (Responses ``instructions``, ``function_call_output``, ``reasoning``) + are skipped. The walk runs front-to-back and back-to-front and both must agree, + so an added message whose text happens to equal a neighbouring real message's + text cannot steal that text's attribution. Returns None otherwise. + """ + runs: Final = tuple(cls._message_texts(message) for message in messages) + + def walk(ordered_runs: Sequence[tuple[str, ...]], ordered_texts: Sequence[str]) -> tuple[int, ...]: + def consume(sources: tuple[int, ...], item: tuple[int, tuple[str, ...]]) -> tuple[int, ...]: + position, run = item + start: Final = len(sources) + if run and tuple(ordered_texts[start : start + len(run)]) == run: + return sources + (position,) * len(run) + return sources + + return functools.reduce(consume, enumerate(ordered_runs), ()) + + forward: Final = walk(runs, texts) + last: Final = len(runs) - 1 + backward: Final = tuple( + last - position for position in walk(tuple(run[::-1] for run in runs[::-1]), texts[::-1])[::-1] + ) + return forward if len(forward) == len(texts) and forward == backward else None + + @classmethod + def _reasoning_item_text_indices( + cls, + texts: Sequence[str], + request_data: Mapping[str, object], + ) -> frozenset[int] | None: + """Return the ``texts`` indices flattened from Responses ``reasoning`` input items. + + The Responses translation handler gives those model-authored items the default + ``user`` role, so the latest-turn selection must not mistake one for a human turn. + Empty for requests without a Responses ``input`` item list; None when the raw items + do not account for every entry of ``texts``. + """ + try: + raw_input: Final = _RESPONSES_INPUT.validate_python(request_data.get("input")) + except ValidationError: + return None + if not isinstance(raw_input, tuple): + return frozenset() + counts: Final = tuple(item.text_count() for item in raw_input) + if sum(counts) != len(texts): + return None + starts: Final = itertools.accumulate(counts, initial=0) + return frozenset( + text_idx + for item, count, start in zip(raw_input, counts, starts) + if item.type == "reasoning" + for text_idx in range(start, start + count) + ) + + @classmethod + def _get_latest_user_text_indices( + cls, + texts: Sequence[str], + messages: Sequence[AllMessageValues], + request_data: Mapping[str, object], + ) -> frozenset[int] | None: """Return text indices belonging to only the latest scannable human-authored (user or developer) message. - Args: - texts: Flattened text entries from the framework. - messages: The structured messages the framework flattened into ``texts``, - hoisted top-level system prompt included, so positions line up. - - Returns a set of scannable indices, or None on count mismatch or no user/developer - message (safety fallback to existing role-filter behavior). + The latest user/developer message is chosen from ``messages`` itself, so a latest turn + without text (image only) yields an empty set rather than promoting an earlier turn. + Messages flattened from Responses ``reasoning`` items are never that turn. + Returns None when ``texts`` cannot be aligned with ``messages`` or ``request_data``, no + user/developer message exists, or the latest one carries text that never reached + ``texts`` (safety fallback to the role-filter scan). """ - last_human_msg_idx: int | None = None - for idx in range(len(messages) - 1, -1, -1): - msg = messages[idx] - if isinstance(msg, dict) and msg.get("role") in ("user", "developer"): - last_human_msg_idx = idx - break - - if last_human_msg_idx is None: - return None # No user/developer message → fallback to existing role-filter scan - - scannable: Final[set] = set() - text_idx = 0 - for msg_idx, msg in enumerate(messages): - if not isinstance(msg, dict): - continue - content = msg.get("content") - is_latest_human = msg_idx == last_human_msg_idx - - if content is None: - pass - elif isinstance(content, str): - if is_latest_human: - scannable.add(text_idx) - text_idx += 1 - elif isinstance(content, list): - for item in content: - if isinstance(item, dict) and item.get("text") is not None: - if is_latest_human: - scannable.add(text_idx) - text_idx += 1 - - if text_idx != len(texts): - return None # Count mismatch → safety fallback - - return scannable + sources: Final = cls._text_source_message_indices(texts, messages) + if sources is None: + return None + reasoning: Final = cls._reasoning_item_text_indices(texts, request_data) + if reasoning is None: + return None + reasoning_messages: Final = frozenset(sources[text_idx] for text_idx in reasoning) + latest_human: Final = max( + ( + idx + for idx, message in enumerate(messages) + if idx not in reasoning_messages and message.get("role") in ("user", "developer") + ), + default=None, + ) + if latest_human is None: + return None + if latest_human not in sources and cls._message_texts(messages[latest_human]): + return None + return frozenset(text_idx for text_idx, source in enumerate(sources) if source == latest_human) def supports_scan_only_tool_results(self) -> bool: return False - @staticmethod + @classmethod def _get_scannable_text_indices( + cls, texts: Sequence[str], - structured_messages: Sequence[object], - ) -> set | None: - """Derive which ``texts`` indices originate from user/system messages. + structured_messages: Sequence[AllMessageValues], + ) -> frozenset[int] | None: + """Derive which ``texts`` indices originate from user/system/developer messages. - The unified guardrail framework flattens message content into ``texts`` - without preserving role info. This helper re-walks - ``structured_messages`` using the **same** extraction logic the - framework uses (string content → 1 entry, list content → 1 per text - item, None → 0) and records the running text index for each entry - whose source role is ``"user"``, ``"system"``, or ``"developer"``. - - Returns a set of scannable indices, or ``None`` if the count doesn't - match ``len(texts)`` (safety fallback → scan everything). + Returns None when ``texts`` cannot be aligned with ``structured_messages`` + (safety fallback: scan everything). """ - scannable: Final[set] = set() - text_idx = 0 - for msg in structured_messages: - if not isinstance(msg, dict): - continue - role = msg.get("role", "") - content = msg.get("content") - is_scannable = role in ("user", "system", "developer") - - if content is None: - # No content → 0 text entries - pass - elif isinstance(content, str): - if is_scannable: - scannable.add(text_idx) - text_idx += 1 - elif isinstance(content, list): - for item in content: - if isinstance(item, dict) and item.get("text") is not None: - if is_scannable: - scannable.add(text_idx) - text_idx += 1 - # Ignore other content types (shouldn't happen) - - if text_idx != len(texts): - # Count mismatch → safety fallback: scan all + sources: Final = cls._text_source_message_indices(texts, structured_messages) + if sources is None: return None - - return scannable + return frozenset( + text_idx + for text_idx, source in enumerate(sources) + if structured_messages[source].get("role") in ("user", "system", "developer") + ) @staticmethod def _mcp_name_fallback(rd: dict) -> str | None: @@ -1783,16 +1841,18 @@ class PanwPrismaAirsHandler(CustomGuardrail): # On request side, determine which text indices correspond to scannable # messages so we can skip scanning assistant/tool history text. - scannable_indices: set | None = None + scannable_indices: frozenset[int] | None = None if input_type == "request": structured_messages: Final = inputs.get("structured_messages") if structured_messages: - # For Anthropic /v1/messages: default to latest-user-only scanning. if self._use_latest_user_only(request_data, logging_obj): - scannable_indices = self._get_latest_user_text_indices(texts, structured_messages) - # Fall through to existing role filtering if: - # - not Anthropic, OR flag explicitly False, OR - # - latest-user extraction returned None (no user / count mismatch) + scannable_indices = self._get_latest_user_text_indices(texts, structured_messages, request_data) + if scannable_indices is not None and not scannable_indices: + verbose_proxy_logger.debug( + "PANW Prisma AIRS: latest user message has no text, so " + "experimental_use_latest_role_message_only leaves nothing to scan for call_id=%s", + call_id, + ) if scannable_indices is None: scannable_indices = self._get_scannable_text_indices(texts, structured_messages) if ( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py b/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py index 606210d3b8b..14d541b2bf0 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py @@ -54,9 +54,12 @@ class PanwPrismaAirsGuardrailConfigModel(GuardrailConfigModel): experimental_use_latest_role_message_only: bool | None = Field( default=None, - description="Anthropic /v1/messages only. When unset: scans only latest user/developer " - "message on request side. Set false to scan all user/system/developer messages. " - "Non-Anthropic unaffected.", + description="Scan only the latest user/developer message on the request side instead of " + "the full conversation history. Set true to enable for every request shape (chat completions, " + "Anthropic /v1/messages, /v1/responses); set false to always scan all user/system/developer " + "messages. When unset: latest-only for Anthropic /v1/messages, full history otherwise. " + "Latest-only trusts caller-supplied history: earlier turns are not rescanned, so enable it only " + "where each turn was scanned when it was the latest message or history is server-controlled.", ) @staticmethod diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 9f5f3da4302..c448473391f 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -235,6 +235,114 @@ def test_guardrail_denial_prevents_provider_and_preserves_allowed_control(gatewa assert len(policy.drain()) == 2 +def test_panw_latest_role_message_only_scans_only_latest_turn_on_responses_input( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + latest: Final = "latest turn " + uuid.uuid4().hex + history: Final = ({"role": "user", "content": "first turn"}, {"role": "assistant", "content": "first reply"}) + shapes: Final = { + "plain": {"input": [*history, {"role": "user", "content": latest}]}, + "instructions": {"instructions": "answer briefly", "input": [*history, {"role": "user", "content": latest}]}, + "function_call_output": { + "input": [ + *history, + {"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": "tool result"}, + {"role": "user", "content": latest}, + ] + }, + "reasoning": { + "input": [ + *history, + {"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]}, + {"role": "user", "content": latest}, + ] + }, + "tool_loop_after_latest": { + "input": [ + *history, + {"role": "user", "content": latest}, + {"type": "reasoning", "id": "rs_2", "content": [{"type": "reasoning_text", "text": "thinking"}]}, + {"type": "function_call", "call_id": "call_2", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_2", "output": "tool result"}, + ] + }, + } + + def scanner(request: Request) -> Reply: + assert request.target == "/v1/scan/sync/request" + body: Final = json.loads(request.body) + return Reply( + body=json.dumps( + { + "action": "allow", + "category": "benign", + "profile_name": "synthetic-profile", + "report_id": "R" + body["tr_id"], + "scan_id": "S" + body["tr_id"], + "tr_id": body["tr_id"], + "prompt_detected": {"injection": False, "url_cats": False, "dlp": False}, + "response_detected": {}, + } + ).encode() + ) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/responses" + return Reply( + body=json.dumps( + { + "id": "resp_" + identity, + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "type": "message", + "id": "msg_" + identity, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "permitted response", "annotations": []}], + } + ], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + ).encode() + ) + + with wire_server(scanner) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "panw_prisma_airs", + "mode": "pre_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-panw-key", + "profile_name": "synthetic-profile", + "experimental_use_latest_role_message_only": True, + }, + } + ] + path: Final = tmp_path / "panw.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key" + ) + for name, shape in shapes.items(): + response = candidate.request("POST", "/v1/responses", {"model": model, **shape}) + assert response.status_code == 200, response.text + assert response.json()["output"][0]["content"][0]["text"] == "permitted response" + scanned = [json.loads(scan.body)["contents"][0]["prompt"] for scan in policy.drain()] + assert scanned == [latest], f"{name}: latest-only scanned {scanned}" + assert json.loads(upstream.drain()[0].body)["input"] == shape["input"] + + @pytest.mark.covers("other.observability.guardrails.bedrock_passthrough_converse_scans_only_caller_content") def test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_definition( gateway: Gateway, tmp_path: Path 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 f25727ebd9a..1f52fa224ee 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 @@ -11,6 +11,9 @@ This test file follows LiteLLM's testing patterns and covers: import copy import json +import logging +from collections.abc import Mapping, Sequence +from contextlib import AbstractContextManager from datetime import datetime from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -200,9 +203,7 @@ class TestPanwAirsInitialization: default_on=True, ) assert handler.api_key == "test_api_key_with_linked_profile" - assert ( - handler.profile_name is None - ) # Should be None, PANW API will use linked profile + assert handler.profile_name is None # Should be None, PANW API will use linked profile class TestPanwAirsPromptScanning: @@ -311,9 +312,7 @@ class TestPanwAirsResponseScanning: ("block", "harmful", True), ], ) - async def test_response_scanning( - self, base_handler, user_api_key_dict, action, category, should_block - ): + async def test_response_scanning(self, base_handler, user_api_key_dict, action, category, should_block): """Test response scanning with allow and block responses.""" request_data = { "model": "gpt-3.5-turbo", @@ -341,9 +340,7 @@ class TestPanwAirsResponseScanning: response=response, ) assert exc_info.value.status_code == 400 - assert "Response blocked by PANW Prisma AI Security policy" in str( - exc_info.value.detail - ) + assert "Response blocked by PANW Prisma AI Security policy" in str(exc_info.value.detail) else: result = await base_handler.async_post_call_success_hook( data=request_data, @@ -381,14 +378,10 @@ class TestPanwAirsAPIIntegration: ) as mock_client: mock_async_client = AsyncMock() mock_async_client.client = MagicMock() - mock_async_client.client.post = AsyncMock( - side_effect=Exception("API Error") - ) + mock_async_client.client.post = AsyncMock(side_effect=Exception("API Error")) mock_client.return_value = mock_async_client - result = await handler._call_panw_api( - "test content", call_id="test-call-id" - ) + result = await handler._call_panw_api("test content", call_id="test-call-id") assert result["action"] == "block" assert result["category"] == "api_error" @@ -408,9 +401,7 @@ class TestPanwAirsAPIIntegration: mock_async_client.client.post = AsyncMock(return_value=mock_response) mock_client.return_value = mock_async_client - result = await handler._call_panw_api( - "test content", call_id="test-call-id" - ) + result = await handler._call_panw_api("test content", call_id="test-call-id") assert result["action"] == "block" assert result["category"] == "api_error" @@ -592,9 +583,7 @@ class TestPanwAirsMaskingFunctionality: assert data["messages"][0]["content"][0]["text"] == "My SSN is XXXXXXXXXX" # Image should remain unchanged assert data["messages"][0]["content"][1]["type"] == "image" - assert ( - data["messages"][0]["content"][1]["url"] == "data:image/jpeg;base64,abc123" - ) + assert data["messages"][0]["content"][1]["url"] == "data:image/jpeg;base64,abc123" @pytest.mark.asyncio async def test_response_masking_on_block(self): @@ -641,9 +630,7 @@ class TestPanwAirsMaskingFunctionality: "litellm_call_id": "test-call-id", } - with patch.object( - handler, "_call_panw_api", side_effect=Exception("API Error") - ): + with patch.object(handler, "_call_panw_api", side_effect=Exception("API Error")): with pytest.raises(HTTPException) as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -771,14 +758,10 @@ class TestPanwAirsAdvancedFeatures: mock_scan_result = { "action": "block", "category": "sensitive_data", - "response_masked_data": { - "data": '{"location": "San Francisco", "ssn": "XXXXXXXXXX"}' - }, + "response_masked_data": {"data": '{"location": "San Francisco", "ssn": "XXXXXXXXXX"}'}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result result = await handler.async_post_call_success_hook( @@ -808,9 +791,7 @@ class TestPanwAirsAdvancedFeatures: Choices( finish_reason="stop", index=1, - message=Message( - content="Another SSN: 987-65-4321", role="assistant" - ), + message=Message(content="Another SSN: 987-65-4321", role="assistant"), ), ], created=1234567890, @@ -831,9 +812,7 @@ class TestPanwAirsAdvancedFeatures: "response_masked_data": {"data": "SSN is XXXXXXXXXX"}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result result = await handler.async_post_call_success_hook( @@ -893,9 +872,7 @@ class TestPanwAirsAdvancedFeatures: mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: with patch( "litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs.panw_prisma_airs.add_guardrail_to_applied_guardrails_header" ) as mock_header: @@ -911,9 +888,7 @@ class TestPanwAirsAdvancedFeatures: # Verify header function was called assert mock_header.called - mock_header.assert_called_once_with( - request_data=request_data, guardrail_name="test_panw_airs" - ) + mock_header.assert_called_once_with(request_data=request_data, guardrail_name="test_panw_airs") class TestTextCompletionSupport: @@ -924,9 +899,7 @@ class TestTextCompletionSupport: """Test that guardrail can extract and scan text completion prompts.""" handler = make_handler() - user_api_key_dict = UserAPIKeyAuth( - api_key="test_key", user_id="test_user", team_id="test_team" - ) + user_api_key_dict = UserAPIKeyAuth(api_key="test_key", user_id="test_user", team_id="test_team") # Text completion request (no messages, just prompt) data = { @@ -938,9 +911,7 @@ class TestTextCompletionSupport: mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result result = await handler.async_pre_call_hook( @@ -953,9 +924,7 @@ class TestTextCompletionSupport: # Verify API was called with the prompt text mock_api.assert_called_once() call_args = mock_api.call_args - assert ( - call_args.kwargs["content"] == "Complete this sentence: AI security is" - ) + assert call_args.kwargs["content"] == "Complete this sentence: AI security is" assert call_args.kwargs["is_response"] is False # Verify request was allowed through @@ -966,9 +935,7 @@ class TestTextCompletionSupport: """Test that masking works with text completion prompts.""" handler = make_handler(mask_request_content=True) - user_api_key_dict = UserAPIKeyAuth( - api_key="test_key", user_id="test_user", team_id="test_team" - ) + user_api_key_dict = UserAPIKeyAuth(api_key="test_key", user_id="test_user", team_id="test_team") data = { "prompt": "Send money to account 123-456-7890", @@ -983,9 +950,7 @@ class TestTextCompletionSupport: "prompt_masked_data": {"data": "Send money to account XXXXXXXXXX"}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result result = await handler.async_pre_call_hook( @@ -1004,9 +969,7 @@ class TestTextCompletionSupport: """Test that guardrail handles batch text completion (list of prompts).""" handler = make_handler() - user_api_key_dict = UserAPIKeyAuth( - api_key="test_key", user_id="test_user", team_id="test_team" - ) + user_api_key_dict = UserAPIKeyAuth(api_key="test_key", user_id="test_user", team_id="test_team") # Batch completion request data = { @@ -1017,9 +980,7 @@ class TestTextCompletionSupport: mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result await handler.async_pre_call_hook( @@ -1053,9 +1014,7 @@ class TestPanwAirsDeduplication: mock_response = {"action": "allow", "category": "benign"} - with patch.object( - handler, "_call_panw_api", return_value=mock_response - ) as mock_api: + with patch.object(handler, "_call_panw_api", return_value=mock_response) as mock_api: # First call - should scan await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1098,9 +1057,7 @@ class TestPanwAirsDeduplication: mock_response = {"action": "allow", "category": "benign"} - with patch.object( - handler, "_call_panw_api", return_value=mock_response - ) as mock_api: + with patch.object(handler, "_call_panw_api", return_value=mock_response) as mock_api: # First call await handler.async_post_call_success_hook( data=data, @@ -1153,9 +1110,7 @@ class TestPanwAirsDeduplication: mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result # First call - should scan @@ -1385,9 +1340,7 @@ class TestPanwAirsFailOpenBehavior: ("network", "allow", False), ], ) - async def test_transient_errors_respect_fallback_setting( - self, error_type, fallback_on_error, should_block - ): + async def test_transient_errors_respect_fallback_setting(self, error_type, fallback_on_error, should_block): """Test that transient errors respect fallback_on_error setting.""" handler = make_handler(fallback_on_error=fallback_on_error) @@ -1404,13 +1357,9 @@ class TestPanwAirsFailOpenBehavior: mock_async_client.client = MagicMock() if error_type == "timeout": - mock_async_client.client.post = AsyncMock( - side_effect=httpx.TimeoutException("Request timeout") - ) + mock_async_client.client.post = AsyncMock(side_effect=httpx.TimeoutException("Request timeout")) else: - mock_async_client.client.post = AsyncMock( - side_effect=httpx.RequestError("Network error") - ) + mock_async_client.client.post = AsyncMock(side_effect=httpx.RequestError("Network error")) mock_client.return_value = mock_async_client @@ -1612,9 +1561,7 @@ class TestPanwAirsAppUserMetadata: ) call_kwargs = mock_async_client.client.post.call_args.kwargs payload = call_kwargs["json"] - assert ( - payload["metadata"]["app_user"] == expected_app_user - ), f"Failed: {description}" + assert payload["metadata"]["app_user"] == expected_app_user, f"Failed: {description}" class TestPanwAirsDeduplicationMissingCallId: @@ -1633,10 +1580,7 @@ class TestPanwAirsDeduplicationMissingCallId: assert already_scanned is False assert data["litellm_call_id"] - assert ( - data["litellm_metadata"][f"_panw_pre_scanned_{data['litellm_call_id']}"] - is True - ) + assert data["litellm_metadata"][f"_panw_pre_scanned_{data['litellm_call_id']}"] is True @pytest.mark.asyncio async def test_call_panw_api_blocks_on_missing_call_id(self): @@ -1696,9 +1640,7 @@ class TestPanwAirsApplyGuardrail: assert result["texts"] == ["Hello world"] mock_api.assert_called_once() - mock_header.assert_called_once_with( - request_data=request_data, guardrail_name=handler.guardrail_name - ) + mock_header.assert_called_once_with(request_data=request_data, guardrail_name=handler.guardrail_name) @pytest.mark.asyncio async def test_apply_guardrail_warns_when_tool_results_scope_leaves_nothing_scannable(self, handler): @@ -1734,9 +1676,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["Malicious content"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "malicious"} with pytest.raises(HTTPException) as exc_info: @@ -1754,9 +1694,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["My SSN is 123-45-6789"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler_mask_request, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_mask_request, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -1777,9 +1715,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["Sensitive response data"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler_mask_response, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_mask_response, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -1809,9 +1745,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": [], "tool_calls": [tool_call]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler_mask_request, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_mask_request, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -1841,9 +1775,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": [], "tool_calls": [tool_call]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "dlp"} with pytest.raises(HTTPException) as exc_info: @@ -1861,9 +1793,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["", " "]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: result = await handler.apply_guardrail( inputs=inputs, request_data=request_data, @@ -1876,14 +1806,10 @@ class TestPanwAirsApplyGuardrail: @pytest.mark.asyncio async def test_apply_guardrail_multiple_texts(self, handler): """Test multiple texts all allowed pass through.""" - inputs: GenericGuardrailAPIInputs = { - "texts": ["Text one", "Text two", "Text three"] - } + inputs: GenericGuardrailAPIInputs = {"texts": ["Text one", "Text two", "Text three"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -1896,16 +1822,12 @@ class TestPanwAirsApplyGuardrail: assert mock_api.call_count == 3 @pytest.mark.asyncio - async def test_apply_guardrail_transient_error_fallback_allow( - self, handler_fail_open - ): + async def test_apply_guardrail_transient_error_fallback_allow(self, handler_fail_open): """Test transient error with fallback_on_error='allow' passes text unscanned.""" inputs: GenericGuardrailAPIInputs = {"texts": ["Test content"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler_fail_open, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_fail_open, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "timeout_error", @@ -1927,9 +1849,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["Test content"]} request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "timeout_error", @@ -1951,9 +1871,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["Test content"]} request_data = {"model": "gpt-4"} # No litellm_call_id - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -1969,16 +1887,12 @@ class TestPanwAirsApplyGuardrail: assert mock_api.call_count == 1 @pytest.mark.asyncio - async def test_apply_guardrail_synthesizes_call_id_for_direct_endpoint( - self, handler - ): + async def test_apply_guardrail_synthesizes_call_id_for_direct_endpoint(self, handler): """Direct /apply_guardrail with empty request_data: call_id synthesized.""" inputs: GenericGuardrailAPIInputs = {"texts": ["Test content"]} request_data: dict = {} # Exactly what guardrail_endpoints.py sends - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -1993,9 +1907,7 @@ class TestPanwAirsApplyGuardrail: assert len(request_data["litellm_call_id"]) == 36 # UUID4 format # PANW API called with synthesized call_id assert mock_api.call_count == 1 - assert ( - mock_api.call_args.kwargs["call_id"] == request_data["litellm_call_id"] - ) + assert mock_api.call_args.kwargs["call_id"] == request_data["litellm_call_id"] @pytest.mark.asyncio async def test_apply_guardrail_call_id_from_logging_obj(self, handler): @@ -2007,9 +1919,7 @@ class TestPanwAirsApplyGuardrail: logging_obj.litellm_call_id = "logging-call-id" logging_obj.model = "gpt-4" - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -2035,9 +1945,7 @@ class TestPanwAirsApplyGuardrail: inputs: GenericGuardrailAPIInputs = {"texts": ["Safe response"]} request_data: dict = {"response": response} # No litellm_call_id - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -2063,9 +1971,7 @@ class TestPanwAirsApplyGuardrail: ]: inputs: GenericGuardrailAPIInputs = {"texts": ["Test"]} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -2137,9 +2043,7 @@ class TestPanwAirsShouldRunGuardrail: ), ], ) - def test_should_run_guardrail( - self, default_on, event_hook, data, query_event, expected - ): + def test_should_run_guardrail(self, default_on, event_hook, data, query_event, expected): handler = make_handler(default_on=default_on, event_hook=event_hook) assert handler.should_run_guardrail(data, query_event) is expected @@ -2164,9 +2068,7 @@ class TestPanwAirsToolEventIsResponseFix: ) ] - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow"} await handler._scan_tool_calls_for_guardrail( tool_calls=tool_calls, @@ -2219,9 +2121,9 @@ class TestPanwAirsToolEventIsResponseFix: tool_event=tool_event, ) - sent_payload = mock_client.client.post.call_args.kwargs.get( - "json" - ) or mock_client.client.post.call_args[1].get("json") + sent_payload = mock_client.client.post.call_args.kwargs.get("json") or mock_client.client.post.call_args[ + 1 + ].get("json") assert "is_response" not in sent_payload["metadata"] assert sent_payload["contents"] == [{"tool_event": tool_event}] @@ -2254,9 +2156,9 @@ class TestPanwAirsToolEventIsResponseFix: tool_event=None, ) - sent_payload = mock_client.client.post.call_args.kwargs.get( - "json" - ) or mock_client.client.post.call_args[1].get("json") + sent_payload = mock_client.client.post.call_args.kwargs.get("json") or mock_client.client.post.call_args[ + 1 + ].get("json") assert sent_payload["metadata"]["is_response"] is True assert sent_payload["contents"] == [{"response": "Hello world"}] @@ -2323,12 +2225,8 @@ class TestPanwAirsMcpForceRun: ), ], ) - def test_should_run_guardrail( - self, guardrail_name, default_on, event_hook, data, query_event, expected - ): - handler = make_handler( - guardrail_name=guardrail_name, default_on=default_on, event_hook=event_hook - ) + def test_should_run_guardrail(self, guardrail_name, default_on, event_hook, data, query_event, expected): + handler = make_handler(guardrail_name=guardrail_name, default_on=default_on, event_hook=event_hook) assert handler.should_run_guardrail(data, query_event) is expected @@ -2359,9 +2257,7 @@ class TestPanwAirsStreamingBytesScan: mock_scan_result = {"action": action, "category": "benign"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result chunks_received = [] @@ -2431,9 +2327,7 @@ class TestPanwAirsStreamingBytesScan: guardrail_info_list = metadata.get("standard_logging_guardrail_information") assert guardrail_info_list is not None # Find the entry with guardrail_status == "success" from _scan_raw_streaming_text - success_entries = [ - g for g in guardrail_info_list if g["guardrail_status"] == "success" - ] + success_entries = [g for g in guardrail_info_list if g["guardrail_status"] == "success"] assert len(success_entries) >= 1 @@ -2494,9 +2388,7 @@ class TestPanwAirsStreamingPydanticEventsScan: mock_scan_result = {"action": action, "category": "benign"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result chunks_received = [] @@ -2568,9 +2460,7 @@ class TestPanwAirsStreamingPydanticEventsScan: guardrail_info_list = metadata.get("standard_logging_guardrail_information") assert guardrail_info_list is not None # Find the entry with guardrail_status == "success" from _scan_raw_streaming_text - success_entries = [ - g for g in guardrail_info_list if g["guardrail_status"] == "success" - ] + success_entries = [g for g in guardrail_info_list if g["guardrail_status"] == "success"] assert len(success_entries) >= 1 @@ -2592,14 +2482,10 @@ class TestPanwAirsApplyGuardrailMetadataEnrichment: logging_obj.litellm_call_id = "test-enrich-id" logging_obj.model = "gpt-4" logging_obj.model_call_details = { - "litellm_params": { - "metadata": {"profile_name": "prod", "app_user": "user-123"} - } + "litellm_params": {"metadata": {"profile_name": "prod", "app_user": "user-123"}} } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -2667,9 +2553,7 @@ class TestPanwAirsToolEventPayload: assert payload["contents"] == [{"response": "World"}] @pytest.mark.asyncio - async def test_tool_event_with_empty_content_still_scans( - self, handler, mock_panw_client - ): + async def test_tool_event_with_empty_content_still_scans(self, handler, mock_panw_client): """tool_event with empty content still sends scan request (not short-circuited).""" tool_event = { "metadata": { @@ -2716,9 +2600,7 @@ class TestPanwAirsToolCallContentScan: ), ) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler._scan_tool_calls_for_guardrail( @@ -2748,9 +2630,7 @@ class TestPanwAirsToolCallContentScan: ), ) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler._scan_tool_calls_for_guardrail( @@ -2878,9 +2758,7 @@ class TestPanwAirsToolCallContentScan: ), ) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "dangerous"} with pytest.raises(HTTPException) as exc_info: @@ -2908,9 +2786,7 @@ class TestPanwAirsToolCallContentScan: ), ) - with patch.object( - handler_mask_request, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_mask_request, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -2939,9 +2815,7 @@ class TestPanwAirsToolCallContentScan: } } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler._scan_tool_calls_for_guardrail( @@ -3135,9 +3009,7 @@ class TestPanwAirsMcpToolEventScan: "mcp_arguments": {"cmd": "rm -rf /"}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "dangerous"} with pytest.raises(HTTPException) as exc_info: @@ -3160,9 +3032,7 @@ class TestPanwAirsMcpToolEventScan: "mcp_arguments": {"path": "/etc/passwd"}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3183,9 +3053,7 @@ class TestPanwAirsMcpToolEventScan: "model": "gpt-4", } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3265,9 +3133,7 @@ class TestPanwAirsMcpToolEventScan: call_kwargs = mock_api.call_args.kwargs te = call_kwargs["tool_event"] - assert_canonical_tool_event( - te, ecosystem="mcp", server_name="test_server", tool_invoked="echo" - ) + assert_canonical_tool_event(te, ecosystem="mcp", server_name="test_server", tool_invoked="echo") assert te["input"] == "hello world" @pytest.mark.asyncio @@ -3373,9 +3239,7 @@ class TestPanwAirsRestMcpFallback: # No 'name', no 'mcp_tool_name' } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3439,9 +3303,7 @@ class TestPanwAirsRestMcpFallback: "name": "my_function", # stray — no "arguments" } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3519,16 +3381,10 @@ class TestPanwAirsDuplicateScanRegression: assert calls[1].kwargs["content"] == 'get_weather\n{"city": "NYC"}' # Third call: MCP scan (tool_event with file_reader) - assert ( - calls[2].kwargs["tool_event"]["metadata"]["server_name"] - == "test_server" - ) + assert calls[2].kwargs["tool_event"]["metadata"]["server_name"] == "test_server" assert calls[2].kwargs["tool_event"]["metadata"]["ecosystem"] == "mcp" assert calls[2].kwargs["tool_event"]["metadata"]["method"] == "tools/call" - assert ( - calls[2].kwargs["tool_event"]["metadata"]["tool_invoked"] - == "file_reader" - ) + assert calls[2].kwargs["tool_event"]["metadata"]["tool_invoked"] == "file_reader" assert "tool_name" not in calls[2].kwargs["tool_event"] @@ -3584,9 +3440,7 @@ class TestPanwAirsChatStreamingPostCall: mock_scan_result = {"action": action, "category": "safe"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result chunks_received = [] @@ -3632,9 +3486,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -3674,9 +3526,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3705,9 +3555,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3727,9 +3575,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3752,9 +3598,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -3780,9 +3624,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3808,9 +3650,7 @@ class TestPanwAirsRequestRoleFiltering: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3866,9 +3706,7 @@ class TestPanwAirsTrIdOverride: assert payload["metadata"]["litellm_trace_id"] == header_trace @pytest.mark.asyncio - async def test_tr_id_uses_call_id_with_requester_metadata_trace( - self, mock_panw_client - ): + async def test_tr_id_uses_call_id_with_requester_metadata_trace(self, mock_panw_client): """requester_metadata.litellm_trace_id is correlation-only, tr_id is always call_id.""" handler = PanwPrismaAirsHandler( guardrail_name="test_panw_airs", @@ -3906,9 +3744,7 @@ class TestPanwAirsTrIdOverride: assert payload["metadata"]["litellm_trace_id"] == trace_id @pytest.mark.asyncio - async def test_top_level_litellm_trace_id_is_correlation_only( - self, mock_panw_client - ): + async def test_top_level_litellm_trace_id_is_correlation_only(self, mock_panw_client): """Top-level data['litellm_trace_id'] is correlation-only, NOT a tr_id override.""" handler = PanwPrismaAirsHandler( guardrail_name="test_panw_airs", @@ -3963,9 +3799,7 @@ class TestPanwAirsDeveloperRoleGuardrail: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -3994,9 +3828,7 @@ class TestPanwAirsDeveloperRoleGuardrail: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "injection"} with pytest.raises(HTTPException) as exc_info: @@ -4025,9 +3857,7 @@ class TestPanwAirsDeveloperRoleGuardrail: "litellm_call_id": "test-call-id", } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.async_pre_call_hook( @@ -4063,9 +3893,7 @@ class TestPanwAirsEmptyToolArgsBlock: ), ) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "block", "category": "dangerous"} with pytest.raises(HTTPException) as exc_info: @@ -4137,9 +3965,7 @@ class TestPanwAirsDictChunkStreaming: for chunk in dict_chunks: yield chunk - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} chunks_received = [] @@ -4179,9 +4005,7 @@ class TestPanwAirsRawStreamingMaskingWarning: "response_masked_data": {"data": "XXXXXXXXX content"}, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = mock_scan_result with patch( @@ -4233,9 +4057,7 @@ class TestPanwAirsUnifiedToolsScan: ) request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4250,8 +4072,7 @@ class TestPanwAirsUnifiedToolsScan: openai_calls = [ c for c in mock_api.call_args_list - if c.kwargs.get("tool_event", {}).get("metadata", {}).get("ecosystem") - == "openai" + if c.kwargs.get("tool_event", {}).get("metadata", {}).get("ecosystem") == "openai" ] assert len(openai_calls) == 0 @@ -4273,9 +4094,7 @@ class TestPanwAirsUnifiedToolsScan: ) request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( inputs=inputs, @@ -4302,9 +4121,7 @@ class TestPanwAirsUnifiedToolsScan: } request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4344,9 +4161,7 @@ class TestPanwAirsUnifiedToolsScan: ) request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( inputs=inputs, @@ -4445,9 +4260,7 @@ class TestPanwAirsLatestRoleMessageOnly: ) @pytest.mark.asyncio - async def test_flag_unset_anthropic_defaults_latest_only( - self, anthropic_request_data, anthropic_inputs - ): + async def test_flag_unset_anthropic_defaults_latest_only(self, anthropic_request_data, anthropic_inputs): """Anthropic + flag None (not set): latest-user-only applied. Instantiate handler via the initializer path (model_dump(exclude_unset=True)) @@ -4474,9 +4287,7 @@ class TestPanwAirsLatestRoleMessageOnly: # Flag should be None (not set), not False assert handler.experimental_use_latest_role_message_only is None - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -4492,15 +4303,11 @@ class TestPanwAirsLatestRoleMessageOnly: assert result["texts"] == list(anthropic_inputs["texts"]) @pytest.mark.asyncio - async def test_flag_false_anthropic_full_scan( - self, anthropic_request_data, anthropic_inputs - ): + async def test_flag_false_anthropic_full_scan(self, anthropic_request_data, anthropic_inputs): """Anthropic + flag false: existing full role-filter behavior (user+system scanned).""" handler = make_handler(experimental_use_latest_role_message_only=False) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4518,15 +4325,11 @@ class TestPanwAirsLatestRoleMessageOnly: assert "First assistant reply" not in scanned @pytest.mark.asyncio - async def test_flag_true_anthropic_latest_only( - self, anthropic_request_data, anthropic_inputs - ): + async def test_flag_true_anthropic_latest_only(self, anthropic_request_data, anthropic_inputs): """Anthropic + flag true: latest-user-only applied.""" handler = make_handler(experimental_use_latest_role_message_only=True) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4539,10 +4342,11 @@ class TestPanwAirsLatestRoleMessageOnly: assert mock_api.call_args.kwargs["content"] == "Latest user message" @pytest.mark.asyncio - async def test_non_anthropic_any_flag_unchanged(self): - """Non-Anthropic + any flag state: existing role-filter behavior.""" - # Even with flag explicitly True, non-Anthropic should not change - handler = make_handler(experimental_use_latest_role_message_only=True) + @pytest.mark.parametrize("flag_value", [None, False]) + async def test_non_anthropic_flag_unset_or_false_full_scan(self, flag_value): + """Non-Anthropic + flag unset or False: existing role-filter behavior.""" + overrides = {} if flag_value is None else {"experimental_use_latest_role_message_only": flag_value} + handler = make_handler(**overrides) inputs: GenericGuardrailAPIInputs = { "texts": ["user prompt", "assistant reply", "system instruction"], @@ -4555,9 +4359,7 @@ class TestPanwAirsLatestRoleMessageOnly: # No proxy_server_request, no anthropic call_type → non-Anthropic request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4"} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4603,9 +4405,7 @@ class TestPanwAirsLatestRoleMessageOnly: }, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4646,9 +4446,7 @@ class TestPanwAirsLatestRoleMessageOnly: }, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await AnthropicMessagesHandler().process_input_messages( @@ -4681,9 +4479,7 @@ class TestPanwAirsLatestRoleMessageOnly: }, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -4750,9 +4546,7 @@ class TestPanwAirsLatestRoleMessageOnly: }, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4795,9 +4589,7 @@ class TestPanwAirsLatestRoleMessageOnly: }, } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4833,9 +4625,7 @@ class TestPanwAirsLatestRoleMessageOnly: "model": "gpt-4", } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4885,9 +4675,7 @@ class TestPanwAirsLatestRoleMessageOnly: ], } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4898,11 +4686,249 @@ class TestPanwAirsLatestRoleMessageOnly: # Only the developer message (latest human-authored) should be scanned assert mock_api.call_count == 1 - assert ( - mock_api.call_args.kwargs["content"] - == "Developer instruction after user" + assert mock_api.call_args.kwargs["content"] == "Developer instruction after user" + + +class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape: + LATEST: Final = "Latest user turn" + HISTORY: Final = ( + {"role": "user", "content": "First user turn"}, + {"role": "assistant", "content": "First assistant turn"}, + ) + ALLOW: Final[Mapping[str, object]] = {"action": "allow", "category": "benign"} + + def _scan( + self, handler: PanwPrismaAirsHandler, scan_result: Mapping[str, object] = ALLOW + ) -> tuple[AbstractContextManager[AsyncMock], AsyncMock]: + mock_api = AsyncMock(return_value=dict(scan_result)) + return patch.object(handler, "_call_panw_api", mock_api), mock_api + + def _responses_request(self, *input_items: Mapping[str, object], **extra: object) -> dict[str, object]: + return { + "litellm_call_id": "test-call-id", + "model": "gpt-4.1-mini", + "input": [*self.HISTORY, *input_items], + **extra, + } + + @pytest.mark.asyncio + async def test_flag_true_chat_completions_scans_latest_user_only(self): + from litellm.llms.openai.chat.guardrail_translation.handler import ( + OpenAIChatCompletionsHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = { + "litellm_call_id": "test-call-id", + "model": "gpt-4.1-mini", + "messages": [ + {"role": "system", "content": "You are terse"}, + *self.HISTORY, + {"role": "user", "content": self.LATEST}, + ], + } + patcher, mock_api = self._scan(handler) + with patcher: + await OpenAIChatCompletionsHandler().process_input_messages(data=request_data, guardrail_to_apply=handler) + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == [self.LATEST] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("history_tail", "instructions"), + [ + pytest.param((), None, id="plain"), + pytest.param((), "answer briefly", id="instructions"), + pytest.param( + ( + {"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": "tool result"}, + ), + None, + id="function_call_output", + ), + pytest.param( + ({"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]},), + None, + id="reasoning", + ), + ], + ) + async def test_flag_true_responses_scans_latest_user_only( + self, history_tail: Sequence[Mapping[str, object]], instructions: str | None + ): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = self._responses_request( + *history_tail, + {"role": "user", "content": self.LATEST}, + **({"instructions": instructions} if instructions is not None else {}), + ) + patcher, mock_api = self._scan( + handler, {"action": "allow", "category": "dlp", "prompt_masked_data": {"data": "[MASKED]"}} + ) + with patcher: + result = await OpenAIResponsesHandler().process_input_messages( + data=request_data, guardrail_to_apply=handler ) + assert [call.kwargs["content"] for call in mock_api.call_args_list] == [self.LATEST] + assert result["input"][-1]["content"] == "[MASKED]" + assert result["input"][0]["content"] == "First user turn" + + @pytest.mark.asyncio + async def test_flag_false_responses_scans_full_history(self): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=False) + request_data = self._responses_request({"role": "user", "content": self.LATEST}, instructions="answer briefly") + patcher, mock_api = self._scan(handler) + with patcher: + await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler) + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["First user turn", self.LATEST] + + @pytest.mark.asyncio + async def test_flag_true_unalignable_texts_fall_back_to_scanning_everything(self): + handler = make_handler(experimental_use_latest_role_message_only=True) + inputs: GenericGuardrailAPIInputs = { + "texts": ["First user turn", "not in any message", self.LATEST], + "structured_messages": [*self.HISTORY, {"role": "user", "content": self.LATEST}], + } + patcher, mock_api = self._scan(handler) + with patcher: + await handler.apply_guardrail(inputs=inputs, request_data={"litellm_call_id": "id"}, input_type="request") + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == list(inputs["texts"]) + + @pytest.mark.asyncio + async def test_flag_true_tool_output_equal_to_latest_user_text_still_scans_latest(self): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = self._responses_request( + {"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": self.LATEST}, + {"role": "user", "content": self.LATEST}, + ) + patcher, mock_api = self._scan(handler) + with patcher: + await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler) + + assert self.LATEST in [call.kwargs["content"] for call in mock_api.call_args_list] + + @pytest.mark.asyncio + async def test_flag_true_image_only_latest_turn_does_not_rescan_history_and_logs_why(self, caplog): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = self._responses_request( + {"role": "user", "content": [{"type": "input_image", "image_url": "https://example.test/cat.png"}]}, + ) + patcher, mock_api = self._scan(handler, {"action": "block", "category": "malicious"}) + with patcher, caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): + result = await OpenAIResponsesHandler().process_input_messages( + data=request_data, guardrail_to_apply=handler + ) + + assert mock_api.call_args_list == [] + assert result["input"] == request_data["input"] + skipped = [r.getMessage() for r in caplog.records if "leaves nothing to scan" in r.getMessage()] + assert skipped == [ + "PANW Prisma AIRS: latest user message has no text, so " + "experimental_use_latest_role_message_only leaves nothing to scan for call_id=test-call-id" + ], caplog.text + + @pytest.mark.asyncio + async def test_flag_true_trailing_reasoning_item_falls_back_to_scanning_history(self): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = self._responses_request( + {"role": "user", "content": self.LATEST}, + {"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]}, + ) + patcher, mock_api = self._scan(handler) + with patcher: + await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler) + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["First user turn", self.LATEST] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "tail", + [ + pytest.param((), id="trailing_reasoning"), + pytest.param( + ( + {"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": "tool result"}, + ), + id="tool_loop", + ), + ], + ) + async def test_flag_true_reasoning_content_after_latest_user_turn_still_scans_that_turn( + self, tail: Sequence[Mapping[str, object]] + ): + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + handler = make_handler(experimental_use_latest_role_message_only=True) + request_data = self._responses_request( + {"role": "user", "content": self.LATEST}, + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "thinking"}], + "content": [{"type": "reasoning_text", "text": "model chain of thought"}], + }, + *tail, + ) + patcher, mock_api = self._scan(handler) + with patcher: + await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler) + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == [self.LATEST] + + @pytest.mark.asyncio + async def test_flag_true_reasoning_content_not_accounted_for_in_texts_falls_back_to_scanning_history(self): + handler = make_handler(experimental_use_latest_role_message_only=True) + reasoning = {"type": "reasoning", "id": "rs_1", "content": [{"type": "reasoning_text", "text": "thinking"}]} + inputs: GenericGuardrailAPIInputs = { + "texts": ["First user turn", self.LATEST, "thinking"], + "structured_messages": [ + *self.HISTORY, + {"role": "user", "content": self.LATEST}, + {"role": "user", "content": [{"type": "text", "text": "thinking"}]}, + ], + } + request_data: dict[str, object] = { + "litellm_call_id": "test-call-id", + "input": [*self.HISTORY, {"role": "user", "content": self.LATEST}, reasoning, "not an input item"], + } + patcher, mock_api = self._scan(handler) + with patcher: + await handler.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + + assert [call.kwargs["content"] for call in mock_api.call_args_list] == [ + "First user turn", + self.LATEST, + "thinking", + ] + class TestPanwAirsMcpToolCallWithoutCallId: """Tests for MCP tool invocations flowing through apply_guardrail without @@ -4930,9 +4956,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: # NO litellm_call_id } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} # Should NOT raise HTTPException(500) @@ -4979,9 +5003,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: mock_logging_obj.model = "gpt-4" mock_logging_obj.model_call_details = {} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -4996,9 +5018,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: assert call_kwargs["call_id"] == "parent-call-id-123" @pytest.mark.asyncio - async def test_direct_apply_guardrail_empty_request_data_synthesizes_plain_uuid( - self, handler - ): + async def test_direct_apply_guardrail_empty_request_data_synthesizes_plain_uuid(self, handler): """Regression: /guardrails/apply_guardrail with empty request_data synthesizes a valid plain UUID.""" import uuid as uuid_mod @@ -5006,9 +5026,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: inputs: GenericGuardrailAPIInputs = {"texts": ["test prompt"]} request_data: dict = {} - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -5101,9 +5119,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: "litellm_call_id": None, # explicitly missing } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} result = await handler.apply_guardrail( @@ -5132,9 +5148,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: # NO mcp_tool_name, NO litellm_call_id } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -5161,9 +5175,7 @@ class TestPanwAirsMcpToolCallWithoutCallId: # no litellm_call_id } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} await handler.apply_guardrail( @@ -5192,24 +5204,18 @@ class TestPanwAirsStreamingFallbackFix: (not raise HTTPException) when _is_transient is set.""" assembled = ModelResponse( id="chatcmpl-123", - choices=[ - Choices(index=0, message=Message(role="assistant", content="hello")) - ], + choices=[Choices(index=0, message=Message(role="assistant", content="hello"))], model="gpt-4", ) request_data = _simple_data(litellm_call_id="test-call-id") - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "_is_transient": True, "action": "block", "category": "api_error", } - result = await handler._scan_and_process_streaming_response( - assembled, request_data, datetime.now() - ) + result = await handler._scan_and_process_streaming_response(assembled, request_data, datetime.now()) content_was_modified, response, scan_result = result assert content_was_modified is False assert scan_result.get("_is_transient") is True @@ -5220,24 +5226,18 @@ class TestPanwAirsStreamingFallbackFix: (not raise HTTPException) when _always_block is set.""" assembled = ModelResponse( id="chatcmpl-123", - choices=[ - Choices(index=0, message=Message(role="assistant", content="hello")) - ], + choices=[Choices(index=0, message=Message(role="assistant", content="hello"))], model="gpt-4", ) request_data = _simple_data(litellm_call_id="test-call-id") - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "_always_block": True, "action": "block", "category": "missing_call_id", } - result = await handler._scan_and_process_streaming_response( - assembled, request_data, datetime.now() - ) + result = await handler._scan_and_process_streaming_response(assembled, request_data, datetime.now()) content_was_modified, response, scan_result = result assert content_was_modified is False assert scan_result.get("_always_block") is True @@ -5267,16 +5267,12 @@ class TestPanwAirsMcpMasking: "litellm_call_id": "test-call-id", } - with patch.object( - handler_masking, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_masking, "_call_panw_api", new_callable=AsyncMock) as mock_api: # texts is empty, so only the MCP tool_event scan fires mock_api.return_value = { "action": "block", "category": "dlp", - "prompt_masked_data": { - "data": '{"path": "/etc/passwd", "secret": "****"}' - }, + "prompt_masked_data": {"data": '{"path": "/etc/passwd", "secret": "****"}'}, } await handler_masking.apply_guardrail( @@ -5308,9 +5304,7 @@ class TestPanwAirsMcpMasking: "litellm_call_id": "test-call-id", } - with patch.object( - handler_no_masking, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_no_masking, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -5338,9 +5332,7 @@ class TestPanwAirsMcpMasking: "litellm_call_id": "test-call-id", } - with patch.object( - handler_masking, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_masking, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -5358,9 +5350,7 @@ class TestPanwAirsMcpMasking: assert request_data["arguments"] == {"key": "****"} @pytest.mark.asyncio - async def test_mcp_structured_args_with_unparseable_masked_text_raises( - self, handler_masking - ): + async def test_mcp_structured_args_with_unparseable_masked_text_raises(self, handler_masking): """When original args are dict but masked text is not valid JSON, should block.""" inputs: GenericGuardrailAPIInputs = {"texts": []} request_data = { @@ -5371,9 +5361,7 @@ class TestPanwAirsMcpMasking: "litellm_call_id": "test-call-id", } - with patch.object( - handler_masking, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_masking, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -5402,9 +5390,7 @@ class TestPanwAirsMcpMasking: # No "arguments" or "mcp_arguments" keys } - with patch.object( - handler_masking, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler_masking, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -5439,9 +5425,7 @@ class TestPanwAirsResponseToolCallMasking: function=Function(name="search", arguments='{"query": "sensitive-data"}'), ) - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "block", "category": "dlp", @@ -5479,9 +5463,7 @@ class TestPanwAirsMcpMaskOnAllow: "litellm_call_id": "test-call-id", } - with patch.object( - handler, "_call_panw_api", new_callable=AsyncMock - ) as mock_api: + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: mock_api.return_value = { "action": "allow", "prompt_masked_data": {"data": '{"query": "my SSN is ****"}'}, @@ -5559,9 +5541,7 @@ class TestPanwAirsDualScanIndependence: } with ( - patch.object( - PanwPrismaAirsHandler, "_get_mcp_server_name", return_value="srv" - ), + patch.object(PanwPrismaAirsHandler, "_get_mcp_server_name", return_value="srv"), patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api, ): mock_api.return_value = {"action": "allow", "category": "benign"} @@ -5632,7 +5612,7 @@ class TestPanwAirsTimeoutCoercion: assert isinstance(params.timeout, float) def test_litellm_params_rejects_garbage_timeout(self): - with pytest.raises(ValueError, match='validation error for LitellmParams'): + with pytest.raises(ValueError, match="validation error for LitellmParams"): LitellmParams( guardrail="panw_prisma_airs", mode="pre_call", @@ -5859,6 +5839,8 @@ class TestPanwAirsScanIdExposure: assert "guardrail_scan_ids" in _UNTRUSTED_ROOT_CONTROL_FIELDS assert "guardrail_scan_metadata" in _UNTRUSTED_METADATA_CONTROL_FIELDS assert "guardrail_scan_metadata" in _UNTRUSTED_ROOT_CONTROL_FIELDS + + class TestPanwAirsBlockedErrorDetailPassthrough: """Regression tests for the full AIRS scan response on blocks. @@ -5897,9 +5879,7 @@ class TestPanwAirsBlockedErrorDetailPassthrough: @pytest.mark.asyncio @pytest.mark.parametrize("is_response", [False, True]) - async def test_block_returns_every_airs_field( - self, base_handler, user_api_key_dict, safe_prompt_data, is_response - ): + async def test_block_returns_every_airs_field(self, base_handler, user_api_key_dict, safe_prompt_data, is_response): response = ModelResponse( id="test_id", choices=[ @@ -5908,9 +5888,8 @@ class TestPanwAirsBlockedErrorDetailPassthrough: model="gpt-3.5-turbo", ) - with patch.object( - base_handler, "_call_panw_api", return_value=copy.deepcopy(self._FULL_BLOCK_RESPONSE) - ): + with patch.object(base_handler, "_call_panw_api", return_value=copy.deepcopy(self._FULL_BLOCK_RESPONSE)): + async def _call_hook(): if is_response: await base_handler.async_post_call_success_hook(