mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
fb3459d78c
commit
fe9451c6cd
6 changed files with 223 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue