mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
e8e3172d7d
commit
ef3a3c16ae
8 changed files with 265 additions and 31 deletions
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}"}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue