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>
This commit is contained in:
devin-ai-integration[bot] 2026-08-15 11:49:03 -07:00 • committed by GitHub
parent fb3459d78c
commit fe9451c6cd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 223 additions and 1 deletions

View file

@ -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

View file

@ -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.

View file

@ -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":

View file

@ -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,

View file

@ -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,
):

View file

@ -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"])