feat(guardrails): map each guardrail scan id to its guardrail, stage and provider (#40327)

* feat(guardrails): map each guardrail scan id to its guardrail, stage and provider

Adds the x-litellm-guardrail-scan-metadata response header, a JSON list of
{guardrail, stage, provider, scan_id} entries, next to the existing
comma-separated x-litellm-guardrail-scan-id header. Prisma AIRS records the
execution stage for every scan and OpenAI Moderation now records its
moderation id too. The new metadata key is internal: client-supplied values
are stripped and it is exposed through the UI CORS allow list.

Resolves LIT-6018

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(guardrails): cap the scan metadata response header at a configurable length

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(guardrails): hardcode the scan metadata header cap

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-08 23:32:31 -07:00 committed by GitHub
parent e8e3172d7d
commit ef3a3c16ae
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 265 additions and 31 deletions

View file

@ -143,6 +143,7 @@ DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD: Final = float(
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3)
)
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH: Final = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150))
MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH: Final = 2048
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS: Final = 2000
@ -197,6 +198,7 @@ LITELLM_UI_ALLOW_HEADERS: Final = [
"x-litellm-adaptive-router-model",
"x-litellm-applied-guardrails",
"x-litellm-guardrail-scan-id",
"x-litellm-guardrail-scan-metadata",
"x-litellm-cache-key",
]

View file

@ -1,10 +1,12 @@
import copy
import json
import os
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
from itertools import accumulate
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias
from typing_extensions import assert_never
from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm import get_secret
@ -12,6 +14,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH,
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
ROUTING_REQUEST_TAGS_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
@ -28,6 +31,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
@ -52,6 +56,15 @@ 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"
GUARDRAIL_SCAN_METADATA_METADATA_KEY: Final = "guardrail_scan_metadata"
class GuardrailScanMetadata(TypedDict):
guardrail: ReadOnly[str | None]
stage: ReadOnly[str]
provider: ReadOnly[str]
scan_id: ReadOnly[str]
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@ -450,6 +463,16 @@ def get_remaining_tokens_and_requests_from_request_data(data: dict) -> dict[str,
return headers
def _serialize_scan_metadata_header(entries: Iterable[object], *, max_length: int) -> str | None:
"""Compact JSON list of scan metadata entries, dropping trailing entries so the header fits in max_length."""
encoded: Final = tuple(json.dumps(entry, separators=(",", ":")) for entry in entries)
lengths: Final = tuple(accumulate(len(item) + 1 for item in encoded))
kept: Final = sum(1 for length in lengths if length + 1 <= max_length)
if kept == 0:
return None
return f"[{','.join(encoded[:kept])}]"
def get_logging_caching_headers(request_data: dict) -> dict | None:
_metadata: Final[dict] = {}
metadata_bucket: Final = request_data.get("metadata")
@ -468,6 +491,15 @@ def get_logging_caching_headers(request_data: dict) -> dict | None:
if scan_ids:
headers["x-litellm-guardrail-scan-id"] = ",".join(scan_ids)
scan_metadata: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
scan_metadata_header: Final = (
_serialize_scan_metadata_header(scan_metadata, max_length=MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH)
if isinstance(scan_metadata, (list, tuple))
else None
)
if scan_metadata_header:
headers["x-litellm-guardrail-scan-metadata"] = scan_metadata_header
if "applied_policies" in _metadata:
headers["x-litellm-applied-policies"] = ",".join(_metadata["applied_policies"])
@ -501,6 +533,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
"applied_policies",
"applied_guardrails",
GUARDRAIL_SCAN_IDS_METADATA_KEY,
GUARDRAIL_SCAN_METADATA_METADATA_KEY,
"policy_sources",
"guardrails",
"guardrail_config",
@ -565,21 +598,40 @@ 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:
def add_guardrail_scan_id(
request_data: dict[str, object],
scan_id: str | None,
*,
guardrail_name: str | None,
provider: str,
stage: GuardrailEventHooks,
) -> None:
"""
Record a provider scan id so it can be surfaced to the caller.
Record a provider scan id, keyed to the guardrail execution that produced it, 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.
audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header, and the
(guardrail, stage, provider, scan_id) entries become the x-litellm-guardrail-scan-metadata 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 ()
scan_ids: Final[tuple[object, ...]] = 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)
entry: Final[GuardrailScanMetadata] = {
"guardrail": guardrail_name,
"stage": stage.value,
"provider": provider,
"scan_id": scan_id,
}
existing_entries: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
entries: Final[tuple[object, ...]] = tuple(existing_entries) if isinstance(existing_entries, (list, tuple)) else ()
if entry not in entries:
_metadata[GUARDRAIL_SCAN_METADATA_METADATA_KEY] = (*entries, entry)
def add_policy_to_applied_policies_header(request_data: dict, policy_name: str | None):
"""

View file

@ -17,7 +17,8 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.proxy.common_utils.callback_utils import add_guardrail_scan_id
from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations
from litellm.types.utils import (
GenericGuardrailAPIInputs,
GuardrailStatus,
@ -218,6 +219,13 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
metadata: Final = request_data.get("metadata") or {}
request_data["metadata"] = metadata
metadata["_openai_moderation_response"] = moderation_response.model_dump()
add_guardrail_scan_id(
request_data=request_data,
scan_id=moderation_response.id,
guardrail_name=self.guardrail_name,
provider=SupportedGuardrailIntegrations.OPENAI_MODERATION.value,
stage=GuardrailEventHooks.post_call if input_type == "response" else GuardrailEventHooks.pre_call,
)
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)

View file

@ -721,10 +721,18 @@ class PanwPrismaAirsHandler(CustomGuardrail):
}
}
def _record_scan_id(self, request_data: dict[str, object], scan_result: Mapping[str, object]) -> None:
def _record_scan_id(
self, request_data: dict[str, object], scan_result: Mapping[str, object], stage: GuardrailEventHooks
) -> 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)
add_guardrail_scan_id(
request_data=request_data,
scan_id=str(scan_id) if scan_id else None,
guardrail_name=self.guardrail_name,
provider=self._PROVIDER_NAME,
stage=stage,
)
def _handle_api_error_with_logging(
self,
@ -948,7 +956,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)
self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool:
"""
@ -1078,7 +1086,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.pre_call,
)
self._record_scan_id(data, scan_result)
self._record_scan_id(data, scan_result, GuardrailEventHooks.pre_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@ -1199,7 +1207,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
self._record_scan_id(data, scan_result)
self._record_scan_id(data, scan_result, GuardrailEventHooks.post_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@ -1401,7 +1409,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
self._record_scan_id(request_data, scan_result)
self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
# Add guardrail to applied guardrails header for observability
add_guardrail_to_applied_guardrails_header(
@ -1475,7 +1483,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
continue
self._record_scan_id(request_data, scan_result)
self._record_scan_id(
request_data,
scan_result,
GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
)
action = scan_result.get("action", "block")
masked_args = self._masked_tool_call_arguments(
@ -1829,7 +1841,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
new_texts.append(text)
continue
self._record_scan_id(request_data, scan_result)
self._record_scan_id(
request_data,
scan_result,
GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
)
action = scan_result.get("action", "block")
masked_text = self._get_masked_text(scan_result, is_response=is_response)
@ -1901,7 +1917,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
# If we reach here, fallback_on_error="allow"
else:
self._record_scan_id(request_data, mcp_scan_result)
self._record_scan_id(request_data, mcp_scan_result, GuardrailEventHooks.pre_call)
action = mcp_scan_result.get("action", "block")
masked_text = self._get_masked_text(mcp_scan_result, is_response=False)
if action == "allow":

View file

@ -235,6 +235,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
"guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
"pillar_response_headers",
@ -291,6 +292,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
"guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,

View file

@ -1,30 +1,33 @@
import copy
import json
import sys
from types import ModuleType, SimpleNamespace
from typing import Final
from unittest.mock import patch
import pytest
import litellm
from litellm.caching.caching import DualCache
from litellm.constants import MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.callback_utils import (
_serialize_scan_metadata_header,
add_guardrail_scan_id,
add_policy_to_applied_policies_header,
decrypt_callback_vars,
encrypt_callback_vars,
get_logging_caching_headers,
initialize_callbacks_on_proxy,
get_remaining_tokens_and_requests_from_request_data,
initialize_callbacks_on_proxy,
normalize_callback_names,
process_callback,
sanitize_openai_provider_metadata,
strip_callback_config,
)
import litellm
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from unittest.mock import patch
from litellm.proxy.common_utils.callback_utils import process_callback
from litellm.types.guardrails import GuardrailEventHooks
def test_get_remaining_tokens_and_requests_from_request_data():
@ -189,20 +192,109 @@ def test_get_logging_caching_headers_merges_metadata_and_litellm_metadata():
assert headers["x-litellm-policy-sources"] == "global-baseline=team_default"
def _record(
request_data: dict[str, object],
scan_id: str | None,
guardrail_name: str = "airs",
provider: str = "panw_prisma_airs",
stage: GuardrailEventHooks = GuardrailEventHooks.pre_call,
) -> None:
add_guardrail_scan_id(
request_data=request_data, scan_id=scan_id, guardrail_name=guardrail_name, provider=provider, stage=stage
)
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)
_record(request_data, "scan-1")
_record(request_data, "scan-1")
_record(request_data, "scan-2")
_record(request_data, 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_scan_metadata_header_maps_each_id_to_its_guardrail_stage_and_provider():
request_data: Final[dict[str, object]] = {"litellm_metadata": {}}
_record(
request_data, "scan-1", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.pre_call
)
_record(
request_data, "mod-1", guardrail_name="mod", provider="openai_moderation", stage=GuardrailEventHooks.pre_call
)
_record(
request_data, "scan-2", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.post_call
)
_record(
request_data, "scan-2", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.post_call
)
_record(request_data, None, guardrail_name="mod", provider="openai_moderation", stage=GuardrailEventHooks.post_call)
headers: Final = get_logging_caching_headers(request_data)
assert headers is not None
assert headers["x-litellm-guardrail-scan-id"] == "scan-1,mod-1,scan-2"
assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
{"guardrail": "airs", "stage": "pre_call", "provider": "panw_prisma_airs", "scan_id": "scan-1"},
{"guardrail": "mod", "stage": "pre_call", "provider": "openai_moderation", "scan_id": "mod-1"},
{"guardrail": "airs", "stage": "post_call", "provider": "panw_prisma_airs", "scan_id": "scan-2"},
]
def test_scan_metadata_keeps_same_id_reused_across_stages():
request_data: Final[dict[str, object]] = {"metadata": {}}
_record(request_data, "scan-1", stage=GuardrailEventHooks.pre_call)
_record(request_data, "scan-1", stage=GuardrailEventHooks.post_call)
headers: Final = get_logging_caching_headers(request_data)
assert headers is not None
assert headers["x-litellm-guardrail-scan-id"] == "scan-1"
assert [entry["stage"] for entry in json.loads(headers["x-litellm-guardrail-scan-metadata"])] == [
"pre_call",
"post_call",
]
def test_scan_metadata_header_drops_trailing_entries_to_stay_within_length_limit():
request_data: Final[dict[str, object]] = {"litellm_metadata": {}}
scan_ids: Final = tuple(f"0f9c4b7e-3d2a-4c1b-9e8f-{index:012d}" for index in range(40))
for scan_id in scan_ids:
_record(request_data, scan_id, stage=GuardrailEventHooks.post_call)
headers: Final = get_logging_caching_headers(request_data)
assert headers is not None
assert headers["x-litellm-guardrail-scan-id"] == ",".join(scan_ids)
header: Final = headers["x-litellm-guardrail-scan-metadata"]
assert len(header) <= MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH
kept: Final = json.loads(header)
assert 1 < len(kept) < len(scan_ids)
assert [entry["scan_id"] for entry in kept] == list(scan_ids[: len(kept)])
def test_serialize_scan_metadata_header_keeps_exactly_the_entries_that_fit():
entries: Final = ({"scan_id": "a"}, {"scan_id": "b"}, {"scan_id": "c"})
two_entries: Final = '[{"scan_id":"a"},{"scan_id":"b"}]'
assert _serialize_scan_metadata_header(entries, max_length=len(two_entries)) == two_entries
assert _serialize_scan_metadata_header(entries, max_length=len(two_entries) - 1) == '[{"scan_id":"a"}]'
assert _serialize_scan_metadata_header(entries, max_length=len(two_entries) + 1) == two_entries
assert _serialize_scan_metadata_header(entries, max_length=1000) == json.dumps(entries, separators=(",", ":"))
assert _serialize_scan_metadata_header(entries, max_length=5) is None
assert _serialize_scan_metadata_header((), max_length=1000) is None
def test_scan_metadata_is_an_internal_metadata_key():
assert sanitize_openai_provider_metadata({"guardrail_scan_metadata": "x", "keep": "y"}) == {"keep": "y"}
def test_get_logging_caching_headers_omits_scan_headers_without_scans():
headers: Final = get_logging_caching_headers({"litellm_metadata": {}})
assert headers is not None
assert "x-litellm-guardrail-scan-id" not in headers
assert "x-litellm-guardrail-scan-metadata" not in headers
def test_initialize_callbacks_on_proxy_instantiates_compression_interception(

View file

@ -3,14 +3,19 @@
Test OpenAI Moderation Guardrail
"""
import json
import os
from typing import Final
from unittest.mock import MagicMock, patch
import httpx
import pytest
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
OpenAIModerationGuardrail,
)
@ -989,3 +994,30 @@ async def test_openai_moderation_initialize_guardrail_forwards_streaming_flags()
assert guardrail.streaming_sampling_rate == 2
finally:
litellm.logging_callback_manager._reset_all_callbacks()
@pytest.mark.asyncio
@pytest.mark.parametrize(("input_type", "stage"), [("request", "pre_call"), ("response", "post_call")])
async def test_openai_moderation_records_moderation_id_as_scan_metadata(input_type: str, stage: str):
"""Each moderation call's id is exposed with the guardrail name, stage and provider that produced it."""
payload: Final = {
"id": f"modr-{stage}",
"model": "omni-moderation-latest",
"results": [{"flagged": False, "categories": {}, "category_scores": {}, "category_applied_input_types": {}}],
}
http_client: Final = AsyncHTTPHandler()
http_client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _: httpx.Response(200, json=payload)))
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
guardrail: Final = OpenAIModerationGuardrail(guardrail_name="openai-mod")
guardrail.async_handler = http_client
request_data: Final[dict[str, object]] = {"metadata": {}}
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type=input_type)
headers: Final = get_logging_caching_headers(request_data)
assert headers is not None
assert headers["x-litellm-guardrail-scan-id"] == f"modr-{stage}"
assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
{"guardrail": "openai-mod", "stage": stage, "provider": "openai_moderation", "scan_id": f"modr-{stage}"}
]

View file

@ -12,6 +12,7 @@ This test file follows LiteLLM's testing patterns and covers:
import copy
import json
from datetime import datetime
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -5785,7 +5786,14 @@ class TestPanwAirsScanIdExposure:
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
assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
{
"guardrail": handler.guardrail_name,
"stage": "pre_call",
"provider": "panw_prisma_airs",
"scan_id": "scan-abc-123",
}
]
@pytest.mark.asyncio
async def test_request_and_response_scan_ids_are_both_exposed(self, user_api_key_dict):
@ -5809,6 +5817,26 @@ class TestPanwAirsScanIdExposure:
headers = get_logging_caching_headers(data)
assert headers["x-litellm-guardrail-scan-id"] == "scan-abc-123,scan-response-456"
assert [(e["stage"], e["scan_id"]) for e in json.loads(headers["x-litellm-guardrail-scan-metadata"])] == [
("pre_call", "scan-abc-123"),
("post_call", "scan-response-456"),
]
@pytest.mark.asyncio
async def test_apply_guardrail_response_scan_is_tagged_post_call(self):
from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers
handler: Final = self._handler(self.ALLOW_SCAN_RESULT)
request_data: Final[dict[str, object]] = {"litellm_call_id": "test-call-id", "model": "gpt-4", "metadata": {}}
await handler.apply_guardrail(
inputs={"texts": ["Hello world"]}, request_data=request_data, input_type="response"
)
headers: Final = get_logging_caching_headers(request_data)
assert headers is not None
entries: Final = json.loads(headers["x-litellm-guardrail-scan-metadata"])
assert [(e["stage"], e["provider"]) for e in entries] == [("post_call", "panw_prisma_airs")]
@pytest.mark.asyncio
async def test_repeated_scan_id_is_not_duplicated(self, user_api_key_dict):
@ -5850,6 +5878,8 @@ class TestPanwAirsScanIdExposure:
assert "guardrail_scan_ids" in _UNTRUSTED_METADATA_CONTROL_FIELDS
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.