From fe9451c6cd7fe486f45f72f704033f9238de1498 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 15 Aug 2026 11:49:03 -0700 Subject: [PATCH] fix(panw_prisma_airs): surface scan_id on allowed requests (#37037) * fix(panw_prisma_airs): surface scan_id and scan metadata on allowed requests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style: ruff format panw guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(panw_prisma_airs): expose scan id header only Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(panw_prisma_airs): inject http client instead of patching private api Adds an http_client seam so the scan-id tests drive the real AIRS request/parse path through a mock transport, plus direct coverage for the scan-id header helper. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): expose guardrail scan id header to browser clients Keeps the panw optional_fields block untouched to avoid a needless conflict with a sibling PR that deletes it. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 2 + litellm/proxy/common_utils/callback_utils.py | 23 +++ .../panw_prisma_airs/panw_prisma_airs.py | 22 ++- litellm/proxy/litellm_pre_call_utils.py | 2 + .../proxy/common_utils/test_callback_utils.py | 17 ++ .../guardrail_hooks/test_panw_prisma_airs.py | 158 ++++++++++++++++++ 6 files changed, 223 insertions(+), 1 deletion(-) diff --git a/litellm/constants.py b/litellm/constants.py index 14ef572888f..8f236eba327 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -141,6 +141,8 @@ LITELLM_UI_ALLOW_HEADERS: Final = [ "x-litellm-semantic-filter", "x-litellm-semantic-filter-tools", "x-litellm-adaptive-router-model", + "x-litellm-applied-guardrails", + "x-litellm-guardrail-scan-id", ] # Gemini model-specific minimal thinking budget constants diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 60a03689804..4afd7c76a35 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -49,6 +49,8 @@ reset_color_code: Final = "\033[0m" TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY: Final = "_pillar_response_headers_trusted" +GUARDRAIL_SCAN_IDS_METADATA_KEY: Final = "guardrail_scan_ids" + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -460,6 +462,10 @@ def get_logging_caching_headers(request_data: dict) -> dict | None: if "applied_guardrails" in _metadata: headers["x-litellm-applied-guardrails"] = ",".join(_metadata["applied_guardrails"]) + scan_ids: Final = _metadata.get(GUARDRAIL_SCAN_IDS_METADATA_KEY) + if scan_ids: + headers["x-litellm-guardrail-scan-id"] = ",".join(scan_ids) + if "applied_policies" in _metadata: headers["x-litellm-applied-policies"] = ",".join(_metadata["applied_policies"]) @@ -492,6 +498,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset( { "applied_policies", "applied_guardrails", + GUARDRAIL_SCAN_IDS_METADATA_KEY, "policy_sources", "guardrails", "guardrail_config", @@ -554,6 +561,22 @@ def add_guardrail_to_applied_guardrails_header(request_data: dict, guardrail_nam _metadata["applied_guardrails"] = [guardrail_name] +def add_guardrail_scan_id(request_data: dict, scan_id: str | None) -> None: + """ + Record a provider scan id so it can be surfaced to the caller. + + Guardrails only return scan details to the client when they block, so allowed requests carry no + audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header. + """ + if not scan_id: + return + _, _metadata = get_or_create_metadata_bucket(request_data) + existing: Final = _metadata.get(GUARDRAIL_SCAN_IDS_METADATA_KEY) + scan_ids: Final = tuple(existing) if isinstance(existing, (list, tuple)) else () + if scan_id not in scan_ids: + _metadata[GUARDRAIL_SCAN_IDS_METADATA_KEY] = (*scan_ids, scan_id) + + def add_policy_to_applied_policies_header(request_data: dict, policy_name: str | None): """ Add a policy name to the applied_policies list in request metadata. 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 dcd86f98ee4..3765771247d 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 @@ -27,11 +27,13 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( effective_scan_only_tool_results_for_guardrail, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_scan_id, add_guardrail_to_applied_guardrails_header, ) from litellm.types.guardrails import GuardrailEventHooks @@ -83,6 +85,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): fallback_on_error: Literal["block", "allow"] = "block", timeout: float = 10.0, violation_message_template: str | None = None, + http_client: AsyncHTTPHandler | None = None, **kwargs, ): """Initialize PANW Prisma AIRS guardrail handler.""" @@ -130,6 +133,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): guardrail_name, ) + self.http_client = http_client self.fallback_on_error = fallback_on_error # Coerce defensively. The dashboard UI persists this field as a JSON # string, and Pydantic extras (the path that splats model_dump into @@ -344,7 +348,9 @@ class PanwPrismaAirsHandler(CustomGuardrail): try: # Use LiteLLM's async HTTP client - async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + async_client: Final = self.http_client or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) # Bypass wrapper to access follow_redirects parameter response: Final = await async_client.client.post( @@ -675,6 +681,11 @@ class PanwPrismaAirsHandler(CustomGuardrail): return error_detail + def _record_scan_id(self, request_data: dict[str, Any], scan_result: Mapping[str, object]) -> None: + """Surface the AIRS scan id on the response, so allowed calls are auditable too.""" + scan_id: Final = scan_result.get("scan_id") + add_guardrail_scan_id(request_data=request_data, scan_id=str(scan_id) if scan_id else None) + def _handle_api_error_with_logging( self, scan_result: dict[str, object], @@ -897,6 +908,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): event_type=GuardrailEventHooks.post_call, ) add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) + self._record_scan_id(request_data, scan_result) def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool: """ @@ -1026,6 +1038,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): duration=(end_time - start_time).total_seconds(), event_type=GuardrailEventHooks.pre_call, ) + self._record_scan_id(data, scan_result) action: Final = scan_result.get("action", "block") category: Final = scan_result.get("category", "unknown") @@ -1146,6 +1159,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): duration=(end_time - start_time).total_seconds(), event_type=GuardrailEventHooks.post_call, ) + self._record_scan_id(data, scan_result) action: Final = scan_result.get("action", "block") category: Final = scan_result.get("category", "unknown") @@ -1347,6 +1361,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): duration=(end_time - start_time).total_seconds(), event_type=GuardrailEventHooks.post_call, ) + self._record_scan_id(request_data, scan_result) # Add guardrail to applied guardrails header for observability add_guardrail_to_applied_guardrails_header( @@ -1450,6 +1465,8 @@ class PanwPrismaAirsHandler(CustomGuardrail): ) continue # fallback_on_error="allow" — leave args unchanged + self._record_scan_id(request_data, scan_result) + 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 @@ -1768,6 +1785,8 @@ class PanwPrismaAirsHandler(CustomGuardrail): new_texts.append(text) continue + self._record_scan_id(request_data, scan_result) + action = scan_result.get("action", "block") masked_text = self._get_masked_text(scan_result, is_response=is_response) @@ -1838,6 +1857,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): ) # If we reach here, fallback_on_error="allow" else: + self._record_scan_id(request_data, mcp_scan_result) action = mcp_scan_result.get("action", "block") masked_text = self._get_masked_text(mcp_scan_result, is_response=False) if action == "allow": diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 0a5626ba0a7..c4a350fb285 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -207,6 +207,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = ( "applied_guardrails", "applied_policies", "policy_sources", + "guardrail_scan_ids", "routing_decision", "pillar_response_headers", "_guardrail_pipelines", @@ -260,6 +261,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = ( "applied_guardrails", "applied_policies", "policy_sources", + "guardrail_scan_ids", "routing_decision", SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, CONSUMED_REQUEST_TAGS_METADATA_KEY, diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py index 0c73c3fcf22..59963bd3707 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py @@ -10,6 +10,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_scan_id, add_policy_to_applied_policies_header, decrypt_callback_vars, encrypt_callback_vars, @@ -192,6 +193,22 @@ def test_get_logging_caching_headers_merges_metadata_and_litellm_metadata(): assert headers["x-litellm-policy-sources"] == "global-baseline=team_default" +def test_add_guardrail_scan_id_dedupes_and_becomes_response_header(): + request_data = {"litellm_metadata": {}} + + add_guardrail_scan_id(request_data=request_data, scan_id="scan-1") + add_guardrail_scan_id(request_data=request_data, scan_id="scan-1") + add_guardrail_scan_id(request_data=request_data, scan_id="scan-2") + add_guardrail_scan_id(request_data=request_data, scan_id=None) + + assert request_data["litellm_metadata"]["guardrail_scan_ids"] == ("scan-1", "scan-2") + assert get_logging_caching_headers(request_data)["x-litellm-guardrail-scan-id"] == "scan-1,scan-2" + + +def test_get_logging_caching_headers_omits_scan_id_header_without_scans(): + assert "x-litellm-guardrail-scan-id" not in get_logging_caching_headers({"litellm_metadata": {}}) + + def test_initialize_callbacks_on_proxy_instantiates_compression_interception( monkeypatch, ): 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 2f0fd51539d..17d4a3e304a 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 @@ -19,6 +19,7 @@ import pytest from fastapi import HTTPException from litellm.caching import DualCache +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( PanwPrismaAirsHandler, @@ -5491,5 +5492,162 @@ class TestPanwAirsTimeoutCoercion: assert handler.timeout == 10.0 +class TestPanwAirsScanIdExposure: + """Allowed scans must expose the AIRS scan id to the caller (LIT-5278).""" + + ALLOW_SCAN_RESULT = { + "action": "allow", + "category": "benign", + "scan_id": "scan-abc-123", + "report_id": "report-abc-123", + "profile_name": "test_profile", + "profile_id": "profile-1", + "tr_id": "tr-9", + } + + @staticmethod + def _handler(*scan_results) -> PanwPrismaAirsHandler: + """Handler wired to a stubbed AIRS endpoint, one queued scan result per call.""" + pending = list(scan_results) + + def respond(request: httpx.Request) -> httpx.Response: + payload = pending.pop(0) if len(pending) > 1 else pending[0] + return httpx.Response(200, json=payload) + + http_client = AsyncHTTPHandler() + http_client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + return make_handler(http_client=http_client) + + @staticmethod + def _recorded_scan_ids(request_data): + metadata = {**request_data.get("metadata", {}), **request_data.get("litellm_metadata", {})} + return metadata.get("guardrail_scan_ids", ()) + + @staticmethod + def _response() -> ModelResponse: + return ModelResponse( + id="test_id", + choices=[Choices(index=0, message=Message(role="assistant", content="hi"))], + model="gpt-4", + ) + + @pytest.mark.asyncio + async def test_pre_call_allow_records_scan_id(self, user_api_key_dict): + handler = self._handler(self.ALLOW_SCAN_RESULT) + data = _simple_data(litellm_call_id="test-call-id", metadata={}) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + + assert self._recorded_scan_ids(data) == ("scan-abc-123",) + + @pytest.mark.asyncio + async def test_post_call_allow_records_response_scan_id(self, user_api_key_dict): + handler = self._handler(self.ALLOW_SCAN_RESULT) + data = {"model": "gpt-4", "litellm_call_id": "test-call-id", "metadata": {}} + + await handler.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=self._response() + ) + + assert self._recorded_scan_ids(data) == ("scan-abc-123",) + + @pytest.mark.asyncio + async def test_apply_guardrail_allow_records_scan_id(self): + handler = self._handler(self.ALLOW_SCAN_RESULT) + inputs: GenericGuardrailAPIInputs = {"texts": ["Hello world"]} + request_data = {"litellm_call_id": "test-call-id", "model": "gpt-4", "metadata": {}} + + await handler.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + + assert self._recorded_scan_ids(request_data) == ("scan-abc-123",) + + @pytest.mark.asyncio + async def test_allowed_scan_id_becomes_response_header(self, user_api_key_dict): + from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers + + handler = self._handler(self.ALLOW_SCAN_RESULT) + data = _simple_data(litellm_call_id="test-call-id", metadata={}) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + + headers = get_logging_caching_headers(data) + assert headers["x-litellm-guardrail-scan-id"] == "scan-abc-123" + assert "x-litellm-guardrail-scan-metadata" not in headers + + @pytest.mark.asyncio + async def test_request_and_response_scan_ids_are_both_exposed(self, user_api_key_dict): + from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers + + handler = self._handler( + self.ALLOW_SCAN_RESULT, + {**self.ALLOW_SCAN_RESULT, "scan_id": "scan-response-456"}, + ) + data = _simple_data(litellm_call_id="test-call-id", metadata={}) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + await handler.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=self._response() + ) + + headers = get_logging_caching_headers(data) + assert headers["x-litellm-guardrail-scan-id"] == "scan-abc-123,scan-response-456" + + @pytest.mark.asyncio + async def test_repeated_scan_id_is_not_duplicated(self, user_api_key_dict): + handler = self._handler(self.ALLOW_SCAN_RESULT) + data = _simple_data(litellm_call_id="test-call-id", metadata={}) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + await handler.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=self._response() + ) + + assert self._recorded_scan_ids(data) == ("scan-abc-123",) + + @pytest.mark.asyncio + async def test_blocked_scan_still_returns_scan_id_in_error(self, user_api_key_dict): + handler = self._handler({**self.ALLOW_SCAN_RESULT, "action": "block", "category": "malicious"}) + data = _simple_data(litellm_call_id="test-call-id", metadata={}) + + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + + assert exc_info.value.detail["error"]["scan_id"] == "scan-abc-123" + + def test_client_supplied_scan_ids_are_stripped(self): + from litellm.proxy.litellm_pre_call_utils import ( + _UNTRUSTED_METADATA_CONTROL_FIELDS, + _UNTRUSTED_ROOT_CONTROL_FIELDS, + ) + + assert "guardrail_scan_ids" in _UNTRUSTED_METADATA_CONTROL_FIELDS + assert "guardrail_scan_ids" in _UNTRUSTED_ROOT_CONTROL_FIELDS + + if __name__ == "__main__": pytest.main([__file__, "-v"])