This commit is contained in:
devin-ai-integration[bot] 2026-09-30 23:46:15 +00:00 • committed by GitHub
commit 956cbd5b04
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
23 changed files with 6485 additions and 78 deletions

View file

@ -21,6 +21,7 @@ from litellm.types.guardrails import (
DynamicGuardrailParams,
GuardrailEventHooks,
LitellmParams,
LoggingOnlyScope,
Mode,
)
from litellm.types.llms.openai import AllMessageValues
@ -175,6 +176,7 @@ class CustomGuardrail(CustomLogger):
use_native_lifecycle_hooks: ClassVar[bool] = False
records_own_guardrail_information: ClassVar[bool] = False
logging_only_scope: LoggingOnlyScope | None
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
super().__init_subclass__(**kwargs)
@ -246,6 +248,7 @@ class CustomGuardrail(CustomLogger):
self.run_in_parallel: bool = run_in_parallel
self.scan_raw_request: bool = scan_raw_request
self.only_scan_new_messages: bool = only_scan_new_messages
self.logging_only_scope = None
if supported_event_hooks:
## validate event_hook is in supported_event_hooks
@ -803,6 +806,13 @@ class CustomGuardrail(CustomLogger):
def uses_apply_guardrail_interface(self) -> bool:
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
@classmethod
def supports_logging_only_scope(cls) -> bool:
return (
cls.apply_guardrail is not CustomGuardrail.apply_guardrail
and cls.async_logging_hook is CustomGuardrail.async_logging_hook
)
def _deployment_hook_target(self) -> "CustomLogger":
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
return self
@ -934,7 +944,7 @@ class CustomGuardrail(CustomLogger):
result: object,
call_type: str,
) -> tuple[dict, object]: # mutable-ok: CustomLogger.async_logging_hook contract
"""logging_only: run apply_guardrail on copies of the logged request/response and record the verdict."""
"""logging_only: scan copies of the logged request and/or response according to logging_only_scope."""
from litellm.llms import get_guardrail_translation_mapping
if not self.uses_apply_guardrail_interface():
@ -981,6 +991,21 @@ class CustomGuardrail(CustomLogger):
"standard_logging_object": {**standard_logging_object, "guardrail_information": [*existing, *entries]},
}, result
def _copy_scratch_request_fields(
self,
kwargs: Mapping[str, object],
) -> tuple[object, object]:
optional_params: Final = kwargs.get("optional_params")
try:
return (
copy.deepcopy(kwargs.get("messages") or kwargs.get("input")),
copy.deepcopy(optional_params.get("tools") if isinstance(optional_params, Mapping) else None),
)
except Exception:
if self.logging_only_scope == "output":
return None, None
raise
async def _scan_logged_call(
self,
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
@ -989,22 +1014,28 @@ class CustomGuardrail(CustomLogger):
output_translation: "BaseTranslation",
scratch_metadata: dict, # mutable-ok: apply_guardrail records its verdict into request metadata
) -> None:
optional_params: Final = kwargs.get("optional_params") or {}
scratch_input: Final = copy.deepcopy(kwargs.get("messages") or kwargs.get("input"))
scratch_input, scratch_tools = self._copy_scratch_request_fields(kwargs)
scratch_request: Final = {
"model": kwargs.get("model"),
"messages": scratch_input,
"input": scratch_input,
"tools": copy.deepcopy(optional_params.get("tools")),
"tools": scratch_tools,
"litellm_call_id": kwargs.get("litellm_call_id"),
"metadata": scratch_metadata,
}
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
if response is None:
if self.logging_only_scope != "output":
try:
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
except Exception as e:
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
if response is None or self.logging_only_scope == "input":
return
await output_translation.process_output_response(
response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request
)
try:
await output_translation.process_output_response(
response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request
)
except Exception as e:
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
def supports_scan_only_tool_results(self) -> bool:
"""Whether this guardrail can scan tool-result content.

View file

@ -10432,6 +10432,23 @@
"description": "Google Cloud location/region (e.g., us-central1)",
"title": "Location"
},
"logging_only_scope": {
"anyOf": [
{
"enum": [
"input",
"output",
"both"
],
"type": "string"
},
{
"type": "null"
}
],
"description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.",
"title": "Logging Only Scope"
},
"mask_request_content": {
"anyOf": [
{
@ -11587,6 +11604,77 @@
"title": "GuardrailSubmissionSummary",
"type": "object"
},
"GuardrailUIAddGuardrailSettings": {
"properties": {
"content_filter_settings": {
"anyOf": [
{
"additionalProperties": true,
"type": "object"
},
{
"type": "null"
}
],
"title": "Content Filter Settings"
},
"pii_entity_categories": {
"items": {
"$ref": "#/components/schemas/PiiEntityCategoryMap"
},
"title": "Pii Entity Categories",
"type": "array"
},
"providers_without_directional_logging_only_scope": {
"items": {
"type": "string"
},
"title": "Providers Without Directional Logging Only Scope",
"type": "array"
},
"supported_actions": {
"items": {
"type": "string"
},
"title": "Supported Actions",
"type": "array"
},
"supported_entities": {
"items": {
"type": "string"
},
"title": "Supported Entities",
"type": "array"
},
"supported_modes": {
"items": {
"type": "string"
},
"title": "Supported Modes",
"type": "array"
},
"supported_modes_by_provider": {
"additionalProperties": {
"items": {
"type": "string"
},
"type": "array"
},
"title": "Supported Modes By Provider",
"type": "object"
}
},
"required": [
"supported_entities",
"supported_actions",
"supported_modes",
"supported_modes_by_provider",
"providers_without_directional_logging_only_scope",
"pii_entity_categories"
],
"title": "GuardrailUIAddGuardrailSettings",
"type": "object"
},
"HTTPValidationError": {
"properties": {
"detail": {
@ -12682,6 +12770,23 @@
"description": "Google Cloud location/region (e.g., us-central1)",
"title": "Location"
},
"logging_only_scope": {
"anyOf": [
{
"enum": [
"input",
"output",
"both"
],
"type": "string"
},
{
"type": "null"
}
],
"description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.",
"title": "Logging Only Scope"
},
"mask": {
"anyOf": [
{
@ -13742,6 +13847,27 @@
"title": "PiiAction",
"type": "string"
},
"PiiEntityCategoryMap": {
"properties": {
"category": {
"title": "Category",
"type": "string"
},
"entities": {
"items": {
"type": "string"
},
"title": "Entities",
"type": "array"
}
},
"required": [
"category",
"entities"
],
"title": "PiiEntityCategoryMap",
"type": "object"
},
"PiiEntityType": {
"enum": [
"CREDIT_CARD",
@ -15183,7 +15309,9 @@
"200": {
"content": {
"application/json": {
"schema": {}
"schema": {
"$ref": "#/components/schemas/GuardrailUIAddGuardrailSettings"
}
}
},
"description": "Successful Response"

View file

@ -33,7 +33,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
build_sandbox_globals,
compile_sandboxed,
)
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
from litellm.proxy.guardrails.guardrail_registry import (
GuardrailRegistry,
_configured_event_hooks,
parse_tolerant_litellm_params,
)
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.repositories.prisma_protocols import TableActions
@ -404,7 +408,11 @@ async def create_guardrail(
guardrail_id: Final = result.get("guardrail_id", "Unknown")
try:
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(guardrail=cast(Guardrail, result), source="db")
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
guardrail=cast(Guardrail, result),
source="db",
reject_invalid_logging_only_scope=True,
)
verbose_proxy_logger.info(
"Immediate sync: Successfully initialized guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
)
@ -527,7 +535,10 @@ async def update_guardrail(
guardrail_name: Final = result.get("guardrail_name", "Unknown")
try:
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=cast(Guardrail, result))
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, result),
reject_invalid_logging_only_scope=True,
)
verbose_proxy_logger.info(
"Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
)
@ -1206,19 +1217,33 @@ async def patch_guardrail(
# Update litellm_params if default_on is provided or pii_entities_config is provided
existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {})))
litellm_params = LitellmParams(**existing_litellm_params)
if request.litellm_params is not None:
requested_litellm_params: Final = request.litellm_params.model_dump(exclude_unset=True)
litellm_params_dict: Final = litellm_params.model_dump(exclude_unset=True)
litellm_params_dict.update(requested_litellm_params)
merged_litellm_params: Final = _as_str_object_mapping(litellm_params_dict)
try:
litellm_params = LitellmParams(**merged_litellm_params)
except ValidationError as validation_error:
raise HTTPException(
status_code=422,
detail=f"Invalid guardrail configuration, update rejected: {validation_error}",
) from validation_error
current_litellm_params: Final = parse_tolerant_litellm_params(
existing_litellm_params,
existing_guardrail.get("guardrail_name") or "Unknown",
)
requested_litellm_params: Final = (
request.litellm_params.model_dump(exclude_unset=True) if request.litellm_params is not None else {}
)
merged_litellm_params: Final = _as_str_object_mapping(
{**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params}
)
try:
parsed_litellm_params: Final = LitellmParams(**merged_litellm_params)
except ValidationError as validation_error:
raise HTTPException(
status_code=422,
detail=f"Invalid guardrail configuration, update rejected: {validation_error}",
) from validation_error
clear_stored_scope: Final = (
"logging_only_scope" not in requested_litellm_params
and parsed_litellm_params.logging_only_scope is not None
and GuardrailEventHooks.logging_only.value not in _configured_event_hooks(parsed_litellm_params.mode)
)
litellm_params: Final = (
LitellmParams(**{**merged_litellm_params, "logging_only_scope": None})
if clear_stored_scope
else parsed_litellm_params
)
# Update guardrail_info if provided
guardrail_info: Final = (
@ -1247,6 +1272,7 @@ async def patch_guardrail(
try:
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=guardrail,
reject_invalid_logging_only_scope="logging_only_scope" in requested_litellm_params,
)
verbose_proxy_logger.info(
"Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
@ -1260,15 +1286,7 @@ async def patch_guardrail(
# the caller instead of a misleading 200.
await GUARDRAIL_REGISTRY.update_guardrail_in_db(
guardrail_id=guardrail_id,
guardrail=Guardrail(
guardrail_id=guardrail_id,
guardrail_name=existing_guardrail.get("guardrail_name") or "",
litellm_params=LitellmParams(**existing_litellm_params),
guardrail_info=existing_guardrail.get(
"guardrail_info",
{}, # mutable-ok: Guardrail's own constructor takes a plain dict
),
),
guardrail=existing_guardrail,
prisma_client=prisma_client,
)
raise HTTPException(
@ -1391,7 +1409,7 @@ async def get_guardrail_info(guardrail_id: str):
tags=["Guardrails"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_guardrail_ui_settings():
async def get_guardrail_ui_settings() -> GuardrailUIAddGuardrailSettings:
"""
Get the UI settings for the guardrails
@ -1425,12 +1443,18 @@ async def get_guardrail_ui_settings():
# above; it only runs on pre_call.
{SupportedGuardrailIntegrations.HIDE_SECRETS.value: [GuardrailEventHooks.pre_call.value]}
)
providers_without_directional_logging_only_scope: Final = [
provider
for provider, guardrail_class in guardrail_class_registry.items()
if not guardrail_class.supports_logging_only_scope()
]
return GuardrailUIAddGuardrailSettings(
supported_entities=[entity.value for entity in PiiEntityType],
supported_actions=[action.value for action in PiiAction],
supported_modes=[mode.value for mode in GuardrailEventHooks],
supported_modes_by_provider=supported_modes_by_provider,
providers_without_directional_logging_only_scope=providers_without_directional_logging_only_scope,
pii_entity_categories=category_maps,
content_filter_settings={
"prebuilt_patterns": get_pattern_metadata(),

View file

@ -54,6 +54,7 @@ from .guardrail_hooks.llm_as_a_judge import (
initialize_guardrail as initialize_llm_as_a_judge,
)
from .guardrail_initializers import (
_configured_event_hooks,
initialize_bedrock,
initialize_hide_secrets,
initialize_lakera,
@ -436,9 +437,46 @@ def _as_callback_tuple(
return (initialized,)
def _configure_callback_scoping(
def _logging_only_scope_error(
custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams
) -> str | None:
logging_only_scope: Final = litellm_params.logging_only_scope
if logging_only_scope is not None and GuardrailEventHooks.logging_only.value not in _configured_event_hooks(
litellm_params.mode
):
return (
f"Guardrail {guardrail_name}: logging_only_scope is set, but mode does not include logging_only, "
"so it would never apply. Add logging_only to mode or remove logging_only_scope."
)
if logging_only_scope in ("input", "output") and not custom_guardrail_callback.supports_logging_only_scope():
return (
f"Guardrail {guardrail_name}: logging_only_scope={logging_only_scope!r} is not supported by this "
"guardrail, whose logging_only hook scans on its own. Remove logging_only_scope."
)
return None
def _configure_callback_scoping(
custom_guardrail_callback: CustomGuardrail,
guardrail_name: str,
litellm_params: LitellmParams,
*,
reject_invalid_logging_only_scope: bool = False,
) -> None:
logging_only_scope: Final = litellm_params.logging_only_scope
logging_only_scope_error: Final = _logging_only_scope_error(
custom_guardrail_callback, guardrail_name, litellm_params
)
if logging_only_scope_error is not None:
if reject_invalid_logging_only_scope:
raise ValueError(logging_only_scope_error)
verbose_proxy_logger.error(
"%s Ignoring logging_only_scope; the guardrail keeps its configured mode.",
logging_only_scope_error.replace("\r", "").replace("\n", ""),
)
custom_guardrail_callback.logging_only_scope = None
else:
custom_guardrail_callback.logging_only_scope = logging_only_scope
for scoping_param in (
"skip_system_message_in_guardrail",
"skip_tool_message_in_guardrail",
@ -461,6 +499,24 @@ def _configure_callback_scoping(
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
def parse_tolerant_litellm_params(
litellm_params_data: Mapping[str, object],
guardrail_name: str,
) -> LitellmParams:
try:
return LitellmParams(**litellm_params_data)
except ValidationError as validation_error:
if any(tuple(error["loc"]) != ("logging_only_scope",) for error in validation_error.errors()):
raise
verbose_proxy_logger.error(
"Guardrail %s: logging_only_scope=%r is not one of 'input', 'output' or 'both'. "
"Ignoring logging_only_scope; the guardrail keeps its configured mode.",
guardrail_name.replace("\r", "").replace("\n", ""),
str(litellm_params_data.get("logging_only_scope")).replace("\r", "").replace("\n", "")[:100],
)
return LitellmParams(**{**litellm_params_data, "logging_only_scope": None})
class InMemoryGuardrailHandler:
"""
Class that handles initializing guardrails and adding them to the CallbackManager
@ -497,6 +553,8 @@ class InMemoryGuardrailHandler:
config_file_path: str | None = None,
llm_router: Optional["Router"] = None,
source: Literal["db", "config"] = "config",
*,
reject_invalid_logging_only_scope: bool = False,
) -> Guardrail | None:
"""
Initialize a guardrail from a dictionary and add it to the litellm callback manager
@ -517,7 +575,10 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
if isinstance(litellm_params_data, dict):
litellm_params = LitellmParams(**litellm_params_data)
if reject_invalid_logging_only_scope:
litellm_params = LitellmParams(**litellm_params_data)
else:
litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"])
else:
litellm_params = litellm_params_data
@ -543,8 +604,18 @@ class InMemoryGuardrailHandler:
config_file_path=config_file_path,
llm_router=llm_router,
)
for custom_guardrail_callback in created_callbacks:
_configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params)
try:
for custom_guardrail_callback in created_callbacks:
_configure_callback_scoping(
custom_guardrail_callback,
guardrail["guardrail_name"],
litellm_params,
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
)
except Exception:
for custom_guardrail_callback in created_callbacks:
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
raise
parsed_guardrail: Final = Guardrail(
guardrail_id=guardrail.get("guardrail_id"),
@ -652,6 +723,8 @@ class InMemoryGuardrailHandler:
guardrail_id: str,
guardrail: Guardrail,
source: Literal["db", "config"] = "db",
*,
reject_invalid_logging_only_scope: bool = False,
) -> None:
"""
Update a guardrail in memory: a changed name or litellm_params rebuilds the
@ -659,8 +732,12 @@ class InMemoryGuardrailHandler:
previous instance and raises), anything else only refreshes the stored row
"""
updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id})
if self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
self.reinitialize_guardrail(guardrail=updated_guardrail, source=source)
if reject_invalid_logging_only_scope or self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
self.reinitialize_guardrail(
guardrail=updated_guardrail,
source=source,
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
)
return
self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail
self._sources[guardrail_id] = source
@ -747,6 +824,7 @@ class InMemoryGuardrailHandler:
@staticmethod
def _normalize_litellm_params_for_comparison(
params: LitellmParams | Mapping[str, object] | None,
guardrail_name: str,
) -> Mapping[str, object] | None:
"""
Render litellm_params to a canonical dict so an in-memory LitellmParams and
@ -763,7 +841,7 @@ class InMemoryGuardrailHandler:
return params.model_dump()
if isinstance(params, dict):
try:
return LitellmParams(**params).model_dump()
return parse_tolerant_litellm_params(params, guardrail_name).model_dump()
except ValidationError as e:
verbose_proxy_logger.warning(
"Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s",
@ -786,8 +864,12 @@ class InMemoryGuardrailHandler:
return True
# Compare litellm_params
existing_dict: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params"))
new_dict: Final = self._normalize_litellm_params_for_comparison(new_guardrail.get("litellm_params"))
existing_dict: Final = self._normalize_litellm_params_for_comparison(
existing.get("litellm_params"), existing.get("guardrail_name", "Unknown")
)
new_dict: Final = self._normalize_litellm_params_for_comparison(
new_guardrail.get("litellm_params"), new_guardrail.get("guardrail_name", "Unknown")
)
# Compare and identify specific differences
changed_fields = {}
@ -813,6 +895,8 @@ class InMemoryGuardrailHandler:
guardrail: Guardrail,
config_file_path: str | None = None,
source: Literal["db", "config"] = "config",
*,
reject_invalid_logging_only_scope: bool = False,
) -> Guardrail | None:
"""
Force re-initialization of a guardrail even if it exists in memory.
@ -842,7 +926,12 @@ class InMemoryGuardrailHandler:
# instance instead of leaving the guardrail silently removed: a guardrail
# that was enforcing must never fail open because an update was bad.
try:
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
return self.initialize_guardrail(
guardrail=guardrail,
config_file_path=config_file_path,
source=source,
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
)
except Exception as init_error:
if previous_guardrail is not None:
verbose_proxy_logger.exception(
@ -857,7 +946,13 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id)
raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
def sync_guardrail_from_db(
self,
guardrail: Guardrail,
config_file_path: str | None = None,
*,
reject_invalid_logging_only_scope: bool = False,
) -> Guardrail | None:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
This is the method to call during DB polling.
@ -867,7 +962,7 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.error("Cannot sync guardrail without guardrail_id")
return None
if self._has_guardrail_params_changed(guardrail_id, guardrail):
if reject_invalid_logging_only_scope or self._has_guardrail_params_changed(guardrail_id, guardrail):
guardrail_name: Final = guardrail.get("guardrail_name", "Unknown")
verbose_proxy_logger.info(
"Guardrail '%s' (ID: %s) params changed, re-initializing...", guardrail_name, guardrail_id
@ -876,6 +971,7 @@ class InMemoryGuardrailHandler:
guardrail=guardrail,
config_file_path=config_file_path,
source="db",
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
)
# Params unchanged but the entry is still DB-backed; make sure the

View file

@ -899,6 +899,8 @@ class ContentFilterConfigModel(BaseModel):
MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"})
LoggingOnlyScope = Literal["input", "output", "both"]
class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch update guardrails
api_key: str | None = Field(default=None, description="API key for the guardrail service")
@ -1140,6 +1142,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
),
)
logging_only_scope: LoggingOnlyScope | None = Field(
default=None,
description=(
"which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' "
"(default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking."
),
)
@field_validator(
"mode",
"default_action",
@ -1308,6 +1318,7 @@ class GuardrailUIAddGuardrailSettings(BaseModel):
supported_actions: list[str]
supported_modes: list[str]
supported_modes_by_provider: dict[str, list[str]]
providers_without_directional_logging_only_scope: list[str]
pii_entity_categories: list[PiiEntityCategoryMap]
content_filter_settings: dict[str, object] | None = None

File diff suppressed because it is too large Load diff

View file

@ -4,6 +4,7 @@ import signal
import socket
import uuid
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timezone
from pathlib import Path
from typing import Final
@ -1796,3 +1797,160 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway,
for response in responses:
assert response.status_code == 200, response.text
assert response.headers["content-type"].startswith("text/event-stream"), response.text
@pytest.mark.parametrize(
("logging_only_scope", "scanned_directions"),
(("input", ("request",)), ("output", ("response",)), ("both", ("request", "response"))),
)
def test_logging_only_scope_observes_only_the_configured_direction_without_blocking(
gateway: Gateway, tmp_path: Path, logging_only_scope: str, scanned_directions: tuple[str, ...]
) -> None:
identity: Final = "guardrail" + uuid.uuid4().hex
prompt: Final = "synthetic observed prompt " + identity
reply: Final = "synthetic observed reply " + identity
texts_by_direction: Final = {"request": [prompt], "response": [reply]}
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api"
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic observed denial"}).encode())
def provider(request: Request) -> Reply:
assert request.target == "/v1/chat/completions"
assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}]
return Reply(
body=json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14},
}
).encode()
)
with wire_server(guardrail) as policy, wire_server(provider) as upstream:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": identity,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "logging_only",
"logging_only_scope": logging_only_scope,
"default_on": True,
"api_base": policy.url,
"api_key": "synthetic-guardrail-key",
},
}
]
path: Final = tmp_path / "logging-only-scope.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model(api_base=upstream.url + "/v1")
response: Final = candidate.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]}
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == reply, response.text
assert len(upstream.drain()) == 1
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda values: len(values) == 1,
seconds=70,
)
scans: Final = tuple(json.loads(scan.body) for scan in policy.drain())
assert [(scan["input_type"], scan["texts"]) for scan in scans] == [
(direction, texts_by_direction[direction]) for direction in scanned_directions
], scans
entries: Final = object_value(rows[0]["metadata"])["guardrail_information"]
assert isinstance(entries, list), rows[0]
assert [
(entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"])
for entry in map(object_value, entries)
] == [(identity, "logging_only", "guardrail_intervened")] * len(scanned_directions), entries
today: Final = datetime.now(timezone.utc).date().isoformat()
guardrail_id: Final = next(
object_value(row)["guardrail_id"]
for row in candidate.get("/v2/guardrails/list")["guardrails"]
if object_value(row)["guardrail_name"] == identity
)
detail: Final = eventually(
lambda: candidate.request(
"GET",
f"/guardrails/usage/detail/{guardrail_id}",
params={"start_date": today, "end_date": today},
).json(),
lambda body: body["requestsEvaluated"] >= len(scanned_directions),
seconds=30,
return_last_on_timeout=True,
)
assert detail["requestsEvaluated"] == len(scanned_directions), detail
@pytest.mark.parametrize("logging_only_scope", ("input", "Input"))
def test_logging_only_scope_literal_or_mode_mismatch_is_ignored_at_load_and_keeps_blocking(
gateway: Gateway, tmp_path: Path, logging_only_scope: str
) -> None:
identity: Final = "guardrail" + uuid.uuid4().hex
prompt: Final = "synthetic invalid-scope prompt pineapple " + identity
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api"
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}).encode())
def provider(request: Request) -> Reply:
assert request.target == "/v1/chat/completions"
assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}]
return Reply(
body=json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "unchanged provider reply"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14},
}
).encode()
)
with wire_server(guardrail) as policy, wire_server(provider) as upstream:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": identity,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"logging_only_scope": logging_only_scope,
"default_on": True,
"blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}],
"api_base": policy.url,
"api_key": "synthetic-guardrail-key",
},
}
]
path: Final = tmp_path / "invalid-scope-pre-call.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model(api_base=upstream.url + "/v1")
response: Final = candidate.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]}
)
assert response.status_code == 400, response.text
assert "synthetic policy denial" in response.text, response.text
assert len(policy.drain()) == 1
assert len(upstream.drain()) == 0
guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"]
assert any(object_value(row)["guardrail_name"] == identity for row in guardrails), guardrails

View file

@ -0,0 +1,762 @@
from __future__ import annotations
import signal
import socket
import threading
import time
import uuid
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from itertools import accumulate, repeat
from pathlib import Path
from queue import SimpleQueue
from typing import Final
import httpx
import psutil
import pytest
from _logging_only_scope_support import (
JSON_OBJECT,
CallerResult,
ChaosCall,
_assert_response_id,
_call_client,
_chaos_calls,
_chaos_control_configuration,
_chaos_model_list,
_chaos_models,
_chaos_spend_minimums,
_configuration,
_direction,
_directions_for_audit_leg,
_drain_upstream,
_guardrail_entries,
_is_base_audit_leg,
_json_contains_exact_string,
_policy_call_id_matches,
_spend_rows,
_spend_rows_for_calls,
_spend_rows_matching_call,
wire_server,
)
from _logging_only_scope_support import (
_record_audit_properties as _record_audit_properties,
)
from anthropic import APIConnectionError as AnthropicAPIConnectionError
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy, owned_proxy_process
from integration._support.wire import Reply, Request
from openai import APIConnectionError as OpenAIAPIConnectionError
from pydantic import JsonValue
def test_K1_policy_edge_restart_mid_burst_keeps_output_observation_fail_open(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = f"logging-scope-k1-{uuid.uuid4().hex}"
marker: Final = uuid.uuid4().hex
def policy(_request: Request) -> Reply:
return Reply(body=b'{"action":"NONE"}')
with socket.socket() as reservation:
reservation.bind(("127.0.0.1", 0))
policy_port: Final = reservation.getsockname()[1]
with gateway.scenario() as scenario:
deployments: Final = _chaos_models(scenario, marker)
models: Final = tuple(deployment.model_name for deployment in deployments)
model_list: Final = _chaos_model_list(deployments)
chaos_calls: Final = _chaos_calls(deployments, marker)
expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output")
control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list)
with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy:
controls: Final = tuple(
_call_client(
call.client_kind,
call.endpoint,
control_proxy,
call.model,
call.prompt,
call.stream,
f"{call.call_id}-control",
)
for call in chaos_calls
)
control_upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(control_upstream) == 30, control_upstream
assert (
tuple(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream)
for call in chaos_calls
)
== (1,) * 30
), control_upstream
config: Final = _configuration(
tmp_path,
identity,
f"http://127.0.0.1:{policy_port}",
"output",
model_list=model_list,
)
with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate:
starts: Final = tuple(threading.Event() for _ in range(3))
def run(index: int) -> tuple[int, CallerResult]:
call: Final = chaos_calls[index]
phase: Final = index // 10
assert starts[phase].wait(timeout=90), (index, phase)
return index, _call_client(
call.client_kind,
call.endpoint,
candidate,
call.model,
call.prompt,
call.stream,
call.call_id,
)
with ThreadPoolExecutor(max_workers=30) as pool:
futures: Final = tuple(pool.submit(run, index) for index in range(30))
try:
with wire_server(policy, port=policy_port) as initial_edge:
starts[0].set()
first: Final = tuple(futures[index].result(timeout=90) for index in range(10))
eventually(
lambda: initial_edge.received.qsize(),
lambda count: count == 10 * len(expected_directions),
seconds=30,
)
tuple(
_spend_rows(model, minimum)
for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:10], 5))
)
starts[1].set()
middle: Final = tuple(futures[index].result(timeout=90) for index in range(10, 20))
tuple(
_spend_rows(model, minimum)
for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:20], 5))
)
with wire_server(policy, port=policy_port) as recovered_edge:
starts[2].set()
recovered: Final = tuple(futures[index].result(timeout=90) for index in range(20, 30))
eventually(
lambda: recovered_edge.received.qsize(),
lambda count: count == 10 * len(expected_directions),
seconds=30,
)
tuple(_spend_rows(model, 10) for model in models)
finally:
for start in starts:
start.set()
results: Final = first + middle + recovered
assert tuple(index for index, _ in results) == tuple(range(30)), results
assert all(result.status == controls[index].status for index, result in results), results
assert all(result.text == controls[index].text for index, result in results), results
candidate_ids: Final = tuple(result.response_id for _, result in results)
assert len(set(candidate_ids)) == 30, candidate_ids
observed_upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(observed_upstream) == 30, observed_upstream
assert (
tuple(
sum(
_json_contains_exact_string(observation["body"], call.prompt)
for observation in observed_upstream
)
for call in chaos_calls
)
== (1,) * 30
), observed_upstream
expected_success_ids: Final = frozenset(
call.call_id for call in chaos_calls if call.index < 10 or call.index >= 20
)
edge_payloads: Final = tuple(
JSON_OBJECT.validate_json(call.body) for call in initial_edge.drain() + recovered_edge.drain()
)
assert len(edge_payloads) == 20 * len(expected_directions), edge_payloads
successful_calls: Final = tuple(call for call in chaos_calls if call.call_id in expected_success_ids)
for call in successful_calls:
payloads_for_call: Final = tuple(
payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id)
)
assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple(
sorted(expected_directions)
), (
call.call_id,
payloads_for_call,
)
assert all(
payload["texts"]
== ([call.prompt] if _direction(payload) == "request" else [results[call.index][1].text])
for payload in payloads_for_call
), payloads_for_call
rows: Final = _spend_rows_for_calls(
models,
tuple(
(
chaos_calls[index].model,
response_id,
chaos_calls[index].call_id,
)
for index, response_id in enumerate(candidate_ids)
),
)
for index, call in enumerate(chaos_calls):
matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id)
assert len(matching_rows) == 1, (call.call_id, matching_rows)
row: Final = matching_rows[0]
expected_status: Final = "guardrail_failed_to_respond" if 10 <= index < 20 else "success"
entries: Final = _guardrail_entries(row)
assert tuple(
(entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries
) == tuple((identity, "logging_only", expected_status) for _ in expected_directions), (index, entries)
def test_K2_policy_edge_delay_does_not_delay_concurrent_callers(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = f"logging-scope-k2-{uuid.uuid4().hex}"
marker: Final = uuid.uuid4().hex
def policy(_request: Request) -> Reply:
time.sleep(2)
return Reply(body=b'{"action":"NONE"}')
with gateway.scenario() as scenario:
deployments: Final = _chaos_models(scenario, marker)
models: Final = tuple(deployment.model_name for deployment in deployments)
model_list: Final = _chaos_model_list(deployments)
chaos_calls: Final = _chaos_calls(deployments, marker)
expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output")
control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list)
with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy:
controls: Final = tuple(
_call_client(
call.client_kind,
call.endpoint,
control_proxy,
call.model,
call.prompt,
call.stream,
f"{call.call_id}-control",
)
for call in chaos_calls
)
control_upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(control_upstream) == 30, control_upstream
assert (
tuple(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream)
for call in chaos_calls
)
== (1,) * 30
), control_upstream
with wire_server(policy) as edge:
config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list)
with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate:
def run(call: ChaosCall) -> tuple[int, CallerResult, float]:
started: Final = time.monotonic()
result: Final = _call_client(
call.client_kind,
call.endpoint,
candidate,
call.model,
call.prompt,
call.stream,
call.call_id,
)
return call.index, result, time.monotonic() - started
with ThreadPoolExecutor(max_workers=30) as pool:
results: Final = tuple(pool.map(run, chaos_calls))
assert tuple(index for index, _, _ in results) == tuple(range(30)), results
assert all(result.status == controls[index].status for index, result, _ in results), results
assert all(result.text == controls[index].text for index, result, _ in results), results
assert all(duration < 2 for _, _, duration in results), results
eventually(
lambda: edge.received.qsize(),
lambda count: count == 30 * len(expected_directions),
seconds=30,
)
edge_calls: Final = edge.drain()
payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge_calls)
for call in chaos_calls:
payloads_for_call: Final = tuple(
payload for payload in payloads if _policy_call_id_matches(payload, call.call_id)
)
assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple(
sorted(expected_directions)
), (
call,
payloads_for_call,
)
assert all(
payload["texts"]
== ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text])
for payload in payloads_for_call
), payloads_for_call
response_ids: Final = tuple(result.response_id for _, result, _ in results)
assert len(set(response_ids)) == 30, response_ids
upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(upstream) == 30, upstream
assert (
tuple(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream)
for call in chaos_calls
)
== (1,) * 30
), upstream
rows: Final = _spend_rows_for_calls(
models,
tuple(
(call.model, result.response_id, call.call_id)
for call, (_, result, _) in zip(chaos_calls, results)
),
)
for call, (_, result, _) in zip(chaos_calls, results):
matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id)
assert len(matching_rows) == 1, (call, result.response_id, matching_rows)
row: Final = matching_rows[0]
entries: Final = _guardrail_entries(row)
assert tuple(
(entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"])
for entry in entries
) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries
@pytest.mark.timeout(180)
def test_K3_two_worker_sigkill_checks_post_kill_spend_rows(
gateway: Gateway,
tmp_path: Path,
record_property: Callable[[str, object], None],
) -> None:
identity: Final = f"logging-scope-k3-{uuid.uuid4().hex}"
marker: Final = uuid.uuid4().hex
scan_started: Final = threading.Event()
release_scans: Final = threading.Event()
def policy(_request: Request) -> Reply:
scan_started.set()
assert release_scans.wait(timeout=60), identity
return Reply(body=b'{"action":"NONE"}')
with gateway.scenario() as scenario:
deployments: Final = _chaos_models(scenario, marker)
models: Final = tuple(deployment.model_name for deployment in deployments)
model_list: Final = _chaos_model_list(deployments)
calls: Final = _chaos_calls(deployments, marker)
expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output")
expected_entries: Final = tuple((identity, "logging_only", "success") for _ in expected_directions)
control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list)
with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy:
controls: Final = tuple(
_call_client(
call.client_kind,
call.endpoint,
control_proxy,
call.model,
call.prompt,
call.stream,
f"{call.call_id}-control",
)
for call in calls
)
assert len(_drain_upstream(gateway.upstream_url)) == 30
with wire_server(policy) as edge:
config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
root: Final = psutil.Process(owned.process.pid)
workers: Final = eventually(
lambda: tuple(
child
for child in root.children(recursive=True)
if any("spawn_main" in part for part in child.cmdline())
),
lambda children: len(children) == 2,
seconds=30,
)
def run(call: ChaosCall) -> tuple[int, CallerResult | None, str | None]:
try:
result: Final = _call_client(
call.client_kind,
call.endpoint,
candidate,
call.model,
call.prompt,
call.stream,
call.call_id,
)
return call.index, result, None
except (OpenAIAPIConnectionError, AnthropicAPIConnectionError, httpx.RemoteProtocolError) as error:
return call.index, None, str(error)
with ThreadPoolExecutor(max_workers=30) as pool:
futures: Final = tuple(pool.submit(run, call) for call in calls)
try:
assert eventually(lambda: scan_started.is_set(), bool, seconds=30)
eventually(lambda: edge.received.qsize(), lambda count: count >= 5, seconds=30)
workers[0].send_signal(signal.SIGKILL)
killed_workers, surviving_workers = psutil.wait_procs((workers[0],), timeout=10)
assert len(killed_workers) == 1 and not surviving_workers, (
killed_workers,
surviving_workers,
)
finally:
release_scans.set()
outcomes: Final = tuple(future.result(timeout=90) for future in futures)
assert owned.process.poll() is None, "Proxy supervisor exited after a worker was killed"
successful: Final = tuple(
(calls[index], result) for index, result, error in outcomes if result is not None and error is None
)
assert successful, outcomes
assert all(
result.status == controls[call.index].status and result.text == controls[call.index].text
for call, result in successful
), successful
pre_kill_expected: Final = tuple(
(call.model, result.response_id, call.call_id) for call, result in successful
)
pre_kill_rows: Final = _spend_rows_for_calls(
models,
pre_kill_expected,
tolerate_missing=True,
)
pre_kill_rows_by_call: Final = tuple(
(call, _spend_rows_matching_call(pre_kill_rows, call.model, call.call_id)) for call in calls
)
assert all(len(rows) <= 1 for _, rows in pre_kill_rows_by_call), pre_kill_rows_by_call
pre_kill_missing_rows: Final = sum(not rows for _, rows in pre_kill_rows_by_call)
record_property("k3_pre_kill_missing_spend_rows", pre_kill_missing_rows)
for call, matching_rows in pre_kill_rows_by_call:
if not matching_rows:
continue
entries: Final = _guardrail_entries(matching_rows[0])
assert tuple(
sorted(
(
entry["guardrail_name"],
entry["guardrail_mode"],
entry["guardrail_status"],
)
for entry in entries
)
) == tuple(sorted(expected_entries)), (call, entries)
post_kill_templates: Final = calls[:6]
post_kill_calls: Final = tuple(
ChaosCall(
index=call.index,
endpoint=call.endpoint,
client_kind=call.client_kind,
model=call.model,
stream=call.stream,
prompt=f"synthetic K post-kill burst {marker}-{call.index}",
call_id=f"{marker}-k-post-kill-{call.index}",
)
for call in post_kill_templates
)
def run_post_kill(call: ChaosCall) -> CallerResult:
return _call_client(
call.client_kind,
call.endpoint,
candidate,
call.model,
call.prompt,
call.stream,
call.call_id,
)
with ThreadPoolExecutor(max_workers=len(post_kill_calls)) as pool:
post_kill_futures: Final = tuple(pool.submit(run_post_kill, call) for call in post_kill_calls)
post_kill_results: Final = tuple(future.result(timeout=90) for future in post_kill_futures)
assert all(
result.status == controls[call.index].status and result.text == controls[call.index].text
for call, result in zip(post_kill_calls, post_kill_results)
), post_kill_results
served: Final = successful + tuple(zip(post_kill_calls, post_kill_results))
response_ids: Final = tuple(result.response_id for _, result in served)
assert len(set(response_ids)) == len(response_ids), response_ids
served_calls: Final = tuple(call for call, _ in served)
requested_call_ids: Final = frozenset(call.call_id for call in calls + post_kill_calls)
upstream: Final = _drain_upstream(gateway.upstream_url)
assert all(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) == 1
for call in served_calls
), upstream
def accumulate_policy_payloads(
collected: tuple[dict[str, JsonValue], ...], _: None
) -> tuple[dict[str, JsonValue], ...]:
return collected + tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain())
payload_batches: Final = accumulate(repeat(None), accumulate_policy_payloads, initial=())
def has_expected_post_kill_scans(collected: tuple[dict[str, JsonValue], ...], call: ChaosCall) -> bool:
payloads_for_call: Final = tuple(
payload for payload in collected if _policy_call_id_matches(payload, call.call_id)
)
return all(
sum(_direction(payload) == direction for payload in payloads_for_call)
>= expected_directions.count(direction)
for direction in expected_directions
)
payloads: Final = eventually(
lambda: next(payload_batches),
lambda collected: all(has_expected_post_kill_scans(collected, call) for call in post_kill_calls),
seconds=30,
)
assert all(
any(_policy_call_id_matches(payload, call_id) for call_id in requested_call_ids)
and _direction(payload) in expected_directions
for payload in payloads
), payloads
for call, result in zip(post_kill_calls, post_kill_results):
payloads_for_call: Final = tuple(
payload for payload in payloads if _policy_call_id_matches(payload, call.call_id)
)
assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple(
sorted(expected_directions)
), (call, payloads_for_call)
assert all(
payload["texts"]
== ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text])
for payload in payloads_for_call
), payloads_for_call
post_kill_expected: Final = tuple(
(call.model, result.response_id, call.call_id)
for call, result in zip(post_kill_calls, post_kill_results)
)
rows: Final = _spend_rows_for_calls(
models,
post_kill_expected,
tolerate_missing=True,
)
all_candidate_calls: Final = calls + post_kill_calls
rows_by_call: Final = tuple(
(call, _spend_rows_matching_call(rows, call.model, call.call_id)) for call in all_candidate_calls
)
assert all(len(matching_rows) <= 1 for _, matching_rows in rows_by_call), rows_by_call
for call, matching_rows in rows_by_call:
if not matching_rows:
continue
entries: Final = _guardrail_entries(matching_rows[0])
assert tuple(
sorted(
(
entry["guardrail_name"],
entry["guardrail_mode"],
entry["guardrail_status"],
)
for entry in entries
)
) == tuple(sorted(expected_entries)), (call, entries)
for call, result in successful:
matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id)
if matching_rows:
_assert_response_id(
call.endpoint,
str(matching_rows[0]["request_id"]),
result.response_id,
marker if call.endpoint == "responses" else call.call_id,
)
for call, result in zip(post_kill_calls, post_kill_results):
matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id)
assert len(matching_rows) == 1, (result.response_id, matching_rows)
_assert_response_id(
call.endpoint,
str(matching_rows[0]["request_id"]),
result.response_id,
marker if call.endpoint == "responses" else call.call_id,
)
@pytest.mark.timeout(180)
def test_K4_proxy_restart_after_fifteen_responses_records_lost_ids(
gateway: Gateway,
tmp_path: Path,
record_property: Callable[[str, object], None],
) -> None:
identity: Final = f"logging-scope-k4-{uuid.uuid4().hex}"
marker: Final = uuid.uuid4().hex
def policy(_request: Request) -> Reply:
return Reply(body=b'{"action":"NONE"}')
with gateway.scenario() as scenario:
deployments: Final = _chaos_models(scenario, marker)
models: Final = tuple(deployment.model_name for deployment in deployments)
model_list: Final = _chaos_model_list(deployments)
calls: Final = _chaos_calls(deployments, marker)
expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output")
control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list)
with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy:
controls: Final = tuple(
_call_client(
call.client_kind,
call.endpoint,
control_proxy,
call.model,
call.prompt,
call.stream,
f"{call.call_id}-control",
)
for call in calls
)
control_upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(control_upstream) == 30, control_upstream
assert (
tuple(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream)
for call in calls
)
== (1,) * 30
), control_upstream
with wire_server(policy) as edge:
config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list)
restart_gate: Final = threading.Event()
second_wave_ready: Final = threading.Event()
second_wave_barrier: Final = threading.Barrier(15, action=second_wave_ready.set)
restarted_gateways: Final[SimpleQueue[Gateway]] = SimpleQueue()
def gateway_for_call(call: ChaosCall, first_gateway: Gateway) -> Gateway:
if call.index < 15:
return first_gateway
second_wave_barrier.wait(timeout=90)
assert restart_gate.wait(timeout=90)
return restarted_gateways.get()
def run(call: ChaosCall, first_gateway: Gateway) -> tuple[int, CallerResult]:
candidate: Final = gateway_for_call(call, first_gateway)
result: Final = _call_client(
call.client_kind,
call.endpoint,
candidate,
call.model,
call.prompt,
call.stream,
call.call_id,
)
return call.index, result
with ThreadPoolExecutor(max_workers=30) as pool:
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first_proxy:
first_proxy_port: Final = first_proxy.gateway.client.base_url.port
assert first_proxy_port is not None
futures: Final = tuple(pool.submit(run, call, first_proxy.gateway) for call in calls)
first_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15))
assert all(
result.status == controls[call.index].status and result.text == controls[call.index].text
for call, result in zip(calls[:15], first_results)
), first_results
assert eventually(lambda: second_wave_ready.is_set(), bool, seconds=30)
eventually(
lambda: edge.received.qsize(),
lambda count: count == 15 * len(expected_directions),
seconds=30,
)
first_wave_expected: Final = tuple(
(call.model, result.response_id, call.call_id)
for call, result in zip(calls[:15], first_results)
)
first_wave_rows: Final = _spend_rows_for_calls(
models,
first_wave_expected,
tolerate_missing=True,
)
assert all(
len(_spend_rows_matching_call(first_wave_rows, model, call_id)) <= 1
for model, _, call_id in first_wave_expected
), first_wave_rows
first_wave_present_response_ids: Final = frozenset(
response_id
for model, response_id, call_id in first_wave_expected
if len(_spend_rows_matching_call(first_wave_rows, model, call_id)) == 1
)
first_wave_lost_response_ids: Final = (
frozenset(response_id for _, response_id, _ in first_wave_expected)
- first_wave_present_response_ids
)
record_property(
"K4_PRE_RESTART_LOST_RESPONSE_IDS",
tuple(sorted(first_wave_lost_response_ids)),
)
record_property(
f"K4_PRE_RESTART_LOST_ROW_COUNT_{'base' if _is_base_audit_leg() else 'head'}",
len(first_wave_lost_response_ids),
)
assert first_proxy.process.poll() is not None, first_proxy.process.pid
eventually(
lambda: tuple(
connection
for connection in psutil.net_connections(kind="tcp")
if connection.status == psutil.CONN_LISTEN and connection.laddr.port == first_proxy_port
),
lambda listeners: not listeners,
seconds=30,
)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted_proxy:
for _ in range(15):
restarted_gateways.put(restarted_proxy.gateway)
restart_gate.set()
second_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15, 30))
assert all(
result.status == controls[call.index].status and result.text == controls[call.index].text
for call, result in zip(calls[15:], second_results)
), second_results
eventually(
lambda: edge.received.qsize(),
lambda count: count == 30 * len(expected_directions),
seconds=30,
)
results: Final = first_results + second_results
expected_response_ids: Final = frozenset(result.response_id for result in results)
assert len(expected_response_ids) == 30, results
upstream: Final = _drain_upstream(gateway.upstream_url)
assert len(upstream) == 30, upstream
assert (
tuple(
sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream)
for call in calls
)
== (1,) * 30
), upstream
edge_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain())
assert len(edge_payloads) == 30 * len(expected_directions), edge_payloads
for call, result in zip(calls, results):
payloads_for_call: Final = tuple(
payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id)
)
assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple(
sorted(expected_directions)
), (
call,
payloads_for_call,
)
assert all(
payload["texts"] == ([call.prompt] if _direction(payload) == "request" else [result.text])
for payload in payloads_for_call
), payloads_for_call
post_restart_expected: Final = tuple(
(call.model, result.response_id, call.call_id) for call, result in zip(calls[15:], second_results)
)
rows: Final = _spend_rows_for_calls(
models,
post_restart_expected,
tolerate_missing=True,
)
record_property("K4_RESPONSE_IDS", tuple(sorted(expected_response_ids)))
for call, result in zip(calls, results):
matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id)
assert len(matching_rows) <= 1, (result.response_id, matching_rows)
if call.index >= 15:
assert len(matching_rows) == 1, (call.call_id, result.response_id, matching_rows)
if matching_rows:
entries: Final = _guardrail_entries(matching_rows[0])
assert tuple(
(entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"])
for entry in entries
) == tuple((identity, "logging_only", "success") for _ in expected_directions), (
call.call_id,
entries,
)

File diff suppressed because it is too large Load diff

File diff suppressed because it is too large Load diff

View file

@ -1,5 +1,6 @@
import json
import time
from collections.abc import Callable
from datetime import datetime
from typing import Dict, List, Optional
from unittest.mock import AsyncMock
@ -10,6 +11,7 @@ import yaml
from fastapi import HTTPException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_endpoints import (
CreateGuardrailRequest,
@ -45,10 +47,12 @@ from litellm.proxy.guardrails.guardrail_registry import (
from litellm.types.guardrails import (
ApplyGuardrailRequest,
BaseLitellmParams,
GuardrailEventHooks,
Guardrail,
GuardrailInfoResponse,
LitellmParams,
)
from litellm.types.utils import GenericGuardrailAPIInputs
# Mock data for testing
MOCK_DB_GUARDRAIL = {
@ -88,6 +92,44 @@ MOCK_PATCH_REQUEST = PatchGuardrailRequest(
)
class _PatchScopeSupportedGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: str,
logging_obj: object | None = None,
) -> GenericGuardrailAPIInputs:
return inputs
class _PatchScopeUnsupportedGuardrail(_PatchScopeSupportedGuardrail):
async def async_logging_hook(
self,
kwargs: dict[str, object],
result: object,
call_type: str,
) -> tuple[dict[str, object], object]:
return kwargs, result
def _patch_scope_initializer(
callback_type: type[CustomGuardrail],
) -> Callable[[LitellmParams, Guardrail], CustomGuardrail]:
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
import litellm
callback = callback_type(
guardrail_name=guardrail["guardrail_name"],
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(callback)
return callback
return _initializer
@pytest.fixture
def mock_prisma_client(mocker):
"""Mock Prisma client for testing"""
@ -127,6 +169,37 @@ def mock_guardrail_registry(mocker):
return mock_registry
def _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
callback_type: type[CustomGuardrail],
guardrail_type: str,
litellm_params: dict[str, object],
) -> tuple[InMemoryGuardrailHandler, Guardrail]:
from litellm.proxy.guardrails import guardrail_registry as registry_module
guardrail: Guardrail = {
"guardrail_id": "patch-scope-test",
"guardrail_name": "Patch scope test",
"litellm_params": {"guardrail": guardrail_type, **litellm_params},
"guardrail_info": {},
}
mock_guardrail_registry.get_guardrail_by_id_from_db.return_value = guardrail
mock_guardrail_registry.update_guardrail_in_db.return_value = guardrail
monkeypatch.setitem(
registry_module.guardrail_initializer_registry,
guardrail_type,
_patch_scope_initializer(callback_type),
)
handler = InMemoryGuardrailHandler()
handler.initialize_guardrail(guardrail=guardrail, source="db")
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry)
mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", handler)
return handler, guardrail
@pytest.mark.asyncio
async def test_list_guardrails_v2_with_db_and_config(mocker, mock_prisma_client, mock_in_memory_handler):
"""Test listing guardrails from both DB and config"""
@ -1126,7 +1199,10 @@ async def test_update_guardrail_endpoint(
prisma_client=mocker.ANY,
)
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY)
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
guardrail=mocker.ANY,
reject_invalid_logging_only_scope=True,
)
if scenario == "success_sync_fails_unexpected_error":
assert mock_logger is not None
@ -1255,7 +1331,10 @@ async def test_patch_guardrail_endpoint(
mock_guardrail_registry.update_guardrail_in_db.assert_called_once()
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY)
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
guardrail=mocker.ANY,
reject_invalid_logging_only_scope=False,
)
if scenario == "success_sync_fails_unexpected_error":
assert mock_logger is not None
@ -1279,6 +1358,192 @@ async def test_patch_guardrail_rejects_mcp_only_on_violation_with_422(mocker, mo
mock_guardrail_registry.update_guardrail_in_db.assert_not_called()
@pytest.mark.asyncio
async def test_patch_guardrail_rejects_invalid_logging_only_scope_with_422(mocker, mock_guardrail_registry):
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
mocker.patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY",
mock_guardrail_registry,
)
mock_in_memory_handler = mocker.Mock(spec=InMemoryGuardrailHandler)
mock_in_memory_handler.sync_guardrail_from_db.side_effect = ValueError(
"Guardrail test-db-guardrail: logging_only_scope is set, but mode does not include logging_only"
)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode="pre_call", logging_only_scope="input"))
with pytest.raises(HTTPException) as exc_info:
await patch_guardrail("test-guardrail-id", request, user_api_key_dict=MOCK_ADMIN_USER)
assert exc_info.value.status_code == 422
assert "update rejected" in str(exc_info.value.detail)
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
guardrail=mocker.ANY,
reject_invalid_logging_only_scope=True,
)
@pytest.mark.asyncio
async def test_patch_guardrail_clears_scope_when_logging_only_mode_is_removed(
mocker, monkeypatch, mock_guardrail_registry
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeSupportedGuardrail,
"patch_scope_supported_test",
{
"mode": ["pre_call", "logging_only"],
"logging_only_scope": "output",
"default_on": True,
},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode=["pre_call"]))
try:
result = await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert result["guardrail_id"] == stored_guardrail["guardrail_id"]
persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"]
assert persisted_guardrail["litellm_params"].logging_only_scope is None
callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]]
assert callback.logging_only_scope is None
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_tolerates_stored_unsupported_scope_on_unrelated_update(
mocker, monkeypatch, mock_guardrail_registry, caplog
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_unsupported_test",
{"mode": "logging_only", "logging_only_scope": "output", "default_on": True},
)
caplog.clear()
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False))
try:
result = await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert result["guardrail_id"] == stored_guardrail["guardrail_id"]
callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]]
assert callback.logging_only_scope is None
assert any("Ignoring logging_only_scope" in record.getMessage() for record in caplog.records)
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_tolerates_invalid_stored_scope_on_unrelated_update(
mocker, monkeypatch, mock_guardrail_registry
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_invalid_literal_test",
{"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False))
try:
result = await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert result["guardrail_id"] == stored_guardrail["guardrail_id"]
persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"]
assert persisted_guardrail["litellm_params"].logging_only_scope is None
assert persisted_guardrail["litellm_params"].default_on is False
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_rejects_explicit_unsupported_scope_and_rolls_back(
mocker, monkeypatch, mock_guardrail_registry
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_unsupported_test",
{"mode": "logging_only", "logging_only_scope": "output", "default_on": True},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output"))
try:
with pytest.raises(HTTPException) as exc_info:
await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert exc_info.value.status_code == 422
assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2
restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"]
assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"]
callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]]
assert callback.logging_only_scope is None
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_rejected_update_restores_invalid_stored_scope_verbatim(
mocker, monkeypatch, mock_guardrail_registry
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_invalid_rollback_test",
{"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output"))
try:
with pytest.raises(HTTPException) as exc_info:
await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert exc_info.value.status_code == 422
assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2
restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"]
assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"]
callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]]
assert callback.logging_only_scope is None
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.parametrize(
"scenario,expected_result,expected_exception",
[
@ -2464,6 +2729,13 @@ async def test_ui_settings_map_matches_runtime_supported_event_hooks():
from litellm.proxy.guardrails.guardrail_registry import guardrail_class_registry
result = await get_guardrail_ui_settings()
expected_without_directional_scope = {
provider
for provider, guardrail_class in guardrail_class_registry.items()
if not guardrail_class.supports_logging_only_scope()
}
assert set(result.providers_without_directional_logging_only_scope) == expected_without_directional_scope
assert "xecguard" in result.providers_without_directional_logging_only_scope
for provider, guardrail_class in guardrail_class_registry.items():
declared = guardrail_class.get_supported_event_hooks()

View file

@ -1,15 +1,19 @@
from collections.abc import Iterable
from typing import ClassVar, Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from pydantic import ValidationError
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.guardrails.guardrail_registry import (
get_guardrail_initializer_from_hooks,
GuardrailRegistry,
InMemoryGuardrailHandler,
get_guardrail_initializer_from_hooks,
parse_tolerant_litellm_params,
)
from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, LoggingOnlyScope, Mode
from litellm.types.utils import GenericGuardrailAPIInputs
def test_get_guardrail_initializer_from_hooks():
@ -472,6 +476,24 @@ def test_unnormalizable_db_params_register_as_changed_without_raising():
assert handler._has_guardrail_params_changed(gid, new) is True
def test_invalid_scope_literal_db_params_compare_equal_after_normalization():
handler = InMemoryGuardrailHandler()
raw = _db_litellm_params()
gid = "77777777-7777-7777-7777-777777777777"
handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail(
guardrail_id=gid,
guardrail_name="cf",
litellm_params=LitellmParams(**{**raw, "logging_only_scope": None}),
)
new = Guardrail(
guardrail_id=gid,
guardrail_name="cf",
litellm_params={**raw, "logging_only_scope": "Input"},
)
assert handler._has_guardrail_params_changed(gid, new) is False
def _all_callback_lists():
import litellm
@ -932,6 +954,230 @@ class TestScanOnlyToolResultsInitRefusal:
)
class _LoggingOnlyScopeSupportedGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: str,
logging_obj: object | None = None,
) -> GenericGuardrailAPIInputs:
return inputs
class _LoggingOnlyScopeUnsupportedGuardrail(_LoggingOnlyScopeSupportedGuardrail):
async def async_logging_hook(
self,
kwargs: dict[str, object],
result: object,
call_type: str,
) -> tuple[dict[str, object], object]:
return kwargs, result
class _LoggingOnlyScopeNativeGuardrail(_LoggingOnlyScopeSupportedGuardrail):
use_native_lifecycle_hooks: ClassVar[bool] = True
def _invalid_scope_content_filter_guardrail() -> Guardrail:
return Guardrail(
guardrail_id="invalid-scope-content-filter-test",
guardrail_name="invalid-scope-content-filter",
litellm_params={
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}],
},
)
class TestLoggingOnlyScopeValidation:
def _initialize(
self,
mode: str | list[str] | Mode,
scope: LoggingOnlyScope | None,
callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail,
reject_invalid_logging_only_scope: bool = False,
assert_registered: bool = False,
) -> CustomGuardrail:
import litellm
from litellm.proxy.guardrails import guardrail_registry as registry_module
guardrail_type: Final = "logging_only_scope_test"
created_callbacks: Final[list[CustomGuardrail]] = []
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
supported_event_hooks: Final = (
[GuardrailEventHooks.logging_only] if callback_type.use_native_lifecycle_hooks else None
)
callback: Final = callback_type(
guardrail_name=guardrail["guardrail_name"],
event_hook=litellm_params.mode,
default_on=True,
supported_event_hooks=supported_event_hooks,
)
litellm.logging_callback_manager.add_litellm_callback(callback)
created_callbacks.append(callback)
return callback
registry_module.guardrail_initializer_registry[guardrail_type] = _initializer
lists: Final = _all_callback_lists()
snapshots: Final = [list(callback_list) for callback_list in lists]
try:
handler: Final = InMemoryGuardrailHandler()
result: Final = handler.initialize_guardrail(
guardrail={
"guardrail_name": "logging-only-scope-guardrail",
"litellm_params": {
"guardrail": guardrail_type,
"mode": mode,
"logging_only_scope": scope,
},
},
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
)
assert result is not None
callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
assert callback is not None
if assert_registered:
assert callback in lists[0]
return callback
except ValueError:
callback: Final = created_callbacks[0]
assert all(callback not in callback_list for callback_list in lists)
raise
finally:
for callback_list, snapshot in zip(lists, snapshots):
callback_list[:] = snapshot
registry_module.guardrail_initializer_registry.pop(guardrail_type, None)
def test_scope_without_logging_only_mode_is_ignored_at_load(self) -> None:
callback: Final = self._initialize(mode="pre_call", scope="input", assert_registered=True)
assert callback.logging_only_scope is None
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True
def test_scope_without_logging_only_mode_is_rejected_for_api_writes(self) -> None:
with pytest.raises(ValueError, match="logging_only_scope is set") as exc_info:
self._initialize(mode="pre_call", scope="input", reject_invalid_logging_only_scope=True)
assert str(exc_info.value) == (
"Guardrail logging-only-scope-guardrail: logging_only_scope is set, but mode does not include "
"logging_only, so it would never apply. Add logging_only to mode or remove logging_only_scope."
)
@pytest.mark.parametrize(
"mode",
(
"logging_only",
["pre_call", "logging_only"],
Mode(tags={"audit": "logging_only"}, default="pre_call"),
),
)
def test_scope_accepts_logging_only_in_supported_mode_forms(self, mode: str | list[str] | Mode) -> None:
callback: Final = self._initialize(mode=mode, scope="input")
assert callback.logging_only_scope == "input"
def test_directional_scope_is_ignored_at_load_when_guardrail_owns_logging_hook(self) -> None:
callback: Final = self._initialize(
mode="logging_only",
scope="input",
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
assert_registered=True,
)
assert callback.logging_only_scope is None
def test_directional_scope_rejected_for_api_writes_when_guardrail_owns_logging_hook(self) -> None:
with pytest.raises(ValueError, match="logging_only_scope='input' is not supported") as exc_info:
self._initialize(
mode="logging_only",
scope="input",
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
reject_invalid_logging_only_scope=True,
)
assert str(exc_info.value) == (
"Guardrail logging-only-scope-guardrail: logging_only_scope='input' is not supported by this "
"guardrail, whose logging_only hook scans on its own. Remove logging_only_scope."
)
def test_both_scope_accepted_when_guardrail_owns_logging_hook(self) -> None:
callback: Final = self._initialize(
mode="logging_only",
scope="both",
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
)
assert callback.logging_only_scope == "both"
def test_output_scope_accepted_for_native_lifecycle_guardrail(self) -> None:
callback: Final = self._initialize(
mode="logging_only",
scope="output",
callback_type=_LoggingOnlyScopeNativeGuardrail,
)
assert callback.logging_only_scope == "output"
def test_invalid_scope_fails_litellm_params_validation(self) -> None:
with pytest.raises(ValidationError):
LitellmParams(guardrail="test", mode="logging_only", logging_only_scope="request")
def test_invalid_scope_literal_keeps_content_filter_registered_and_blocking(self) -> None:
import litellm
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
handler: Final = InMemoryGuardrailHandler()
callback_lists: Final = _all_callback_lists()
callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists]
guardrail: Final = _invalid_scope_content_filter_guardrail()
try:
result: Final = handler.initialize_guardrail(guardrail=guardrail, source="config")
assert result is not None
callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
assert isinstance(callback, ContentFilterGuardrail)
assert callback in litellm.callbacks
assert callback.logging_only_scope is None
assert callback.event_hook == GuardrailEventHooks.pre_call
assert callback._check_blocked_words("pineapple") is not None
finally:
handler.delete_in_memory_guardrail(guardrail["guardrail_id"])
for callback_list, snapshot in zip(callback_lists, callback_snapshots):
callback_list[:] = snapshot
def test_invalid_scope_literal_is_rejected_for_strict_initialization_without_callback_leakage(self) -> None:
handler: Final = InMemoryGuardrailHandler()
callback_lists: Final = _all_callback_lists()
callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists]
with pytest.raises(ValueError, match="logging_only_scope"):
handler.initialize_guardrail(
guardrail=_invalid_scope_content_filter_guardrail(),
source="config",
reject_invalid_logging_only_scope=True,
)
assert all(callback_list == snapshot for callback_list, snapshot in zip(callback_lists, callback_snapshots))
def test_invalid_scope_literal_does_not_tolerate_other_litellm_params_errors(self) -> None:
with pytest.raises(ValidationError):
parse_tolerant_litellm_params(
{
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"default_on": "not-a-bool",
},
"invalid-scope-content-filter",
)
@pytest.mark.asyncio
async def test_update_guardrail_in_db_raises_when_row_missing():
prisma_client = MagicMock()

View file

@ -12,7 +12,7 @@ from litellm.integrations.custom_guardrail import (
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.guardrails import GuardrailEventHooks, LoggingOnlyScope, Mode
from litellm.types.utils import (
Choices,
GenericGuardrailAPIInputs,
@ -2577,6 +2577,32 @@ class TestLoggingOnlyApplyGuardrail:
assert "standard_logging_guardrail_information" not in kwargs["litellm_params"]["metadata"]
assert kwargs["standard_logging_object"] == {"guardrail_information": None}
@pytest.mark.parametrize(
"scope,expected_calls",
(
(None, [("request", ["hello there"]), ("response", ["general kenobi"])]),
("both", [("request", ["hello there"]), ("response", ["general kenobi"])]),
("input", [("request", ["hello there"])]),
("output", [("response", ["general kenobi"])]),
),
)
@pytest.mark.asyncio
async def test_logging_only_scope_scans_configured_directions(
self,
scope: LoggingOnlyScope | None,
expected_calls: list[tuple[str, list[str]]],
) -> None:
guardrail: Final = _ApplyOnlyObserver()
guardrail.logging_only_scope = scope
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == expected_calls
entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"]
assert len(entries) == len(expected_calls)
assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"] * len(expected_calls)
@pytest.mark.asyncio
async def test_appends_to_pre_call_verdicts_without_duplicating_them(self):
guardrail = _ApplyOnlyObserver()
@ -2604,6 +2630,23 @@ class TestLoggingOnlyApplyGuardrail:
assert out_kwargs is kwargs
assert out_response is response
@pytest.mark.asyncio
async def test_output_scope_scans_response_when_request_copy_fails(self):
import threading
guardrail: Final = _ApplyOnlyObserver()
guardrail.logging_only_scope = "output"
call: Final = _logged_call([{"role": "user", "content": "hello there", "lock": threading.Lock()}])
kwargs: Final = call[0]
response: Final = call[1]
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == [("response", ["general kenobi"])]
entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [entry["guardrail_name"] for entry in entries] == ["apply-only-observer"]
assert [entry["guardrail_status"] for entry in entries] == ["success"]
@pytest.mark.asyncio
async def test_block_verdict_is_recorded_without_raising(self):
guardrail = _ApplyOnlyObserver(block=True)
@ -2611,9 +2654,12 @@ class TestLoggingOnlyApplyGuardrail:
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == [("request", ["flagged content"])]
assert guardrail.calls == [
("request", ["flagged content"]),
("response", ["general kenobi"]),
]
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["guardrail_intervened"]
assert [entry["guardrail_status"] for entry in entries] == ["guardrail_intervened", "guardrail_intervened"]
@pytest.mark.asyncio
async def test_call_type_without_translation_is_skipped(self):
@ -2932,6 +2978,17 @@ class _NativeLifecycleLoggingGuardrail(CustomGuardrail):
return inputs
@pytest.mark.asyncio
async def test_native_lifecycle_guardrail_logging_only_scope_scans_only_input():
guardrail: Final = _NativeLifecycleLoggingGuardrail()
guardrail.logging_only_scope = "input"
kwargs, response = _logged_call([{"role": "user", "content": "native lifecycle input"}])
await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == [("request", ["native lifecycle input"])]
@pytest.mark.asyncio
async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response():
"""A use_native_lifecycle_hooks guardrail accepts mode logging_only and its
@ -2939,9 +2996,7 @@ async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(
from litellm.types.utils import Choices, Message, ModelResponse
guardrail = _NativeLifecycleLoggingGuardrail()
assembled = ModelResponse(
choices=[Choices(message=Message(role="assistant", content="assembled stream text"))]
)
assembled = ModelResponse(choices=[Choices(message=Message(role="assistant", content="assembled stream text"))])
sentinel_result = object()
kwargs = {
"model": "gpt-5.4-mini",

View file

@ -1,11 +1,17 @@
"use client";
import { CircleHelp } from "lucide-react";
import React, { useId } from "react";
import React, { useEffect, useId } from "react";
import { useController, type Control, type ControllerRenderProps, type RegisterOptions } from "react-hook-form";
import { Field, FieldDescription, FieldError, FieldLabel } from "@/components/ui/field";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip";
import {
getLoggingOnlyScopeOptions,
modeIncludesLoggingOnly,
normalizeLoggingOnlyScopeChoice,
type LoggingOnlyScopeChoice,
} from "./guardrail_info_helpers";
export interface GuardrailCriterion {
name: string;
@ -15,6 +21,7 @@ export interface GuardrailCriterion {
export interface GuardrailFormValues extends Record<string, unknown> {
criteria?: GuardrailCriterion[];
logging_only_scope_choice?: LoggingOnlyScopeChoice;
}
export type GuardrailFormControl = Control<GuardrailFormValues>;
export type GuardrailFieldRules = Pick<RegisterOptions<GuardrailFormValues, string>, "validate">;
@ -38,6 +45,9 @@ export const asText = (value: unknown): string => {
return "";
};
const isLoggingOnlyScopeChoice = (value: unknown): value is LoggingOnlyScopeChoice =>
value === "default" || value === "input" || value === "output" || value === "both";
export const asStringArray = (value: unknown): string[] => {
if (Array.isArray(value)) return value.filter((entry): entry is string => typeof entry === "string");
if (typeof value === "string" && value !== "") return [value];
@ -123,3 +133,55 @@ export const SkipMessageSelect: React.FC<{ control: GuardrailFieldControlProps }
</Select>
);
};
export const LoggingOnlyScopeSelect: React.FC<{
control: GuardrailFieldControlProps;
directionalScopeSupported: boolean;
}> = ({ control, directionalScopeSupported }) => {
const { id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy } = control;
const items = getLoggingOnlyScopeOptions(directionalScopeSupported);
useEffect(() => {
const currentChoice = isLoggingOnlyScopeChoice(value) ? value : "default";
const choice = normalizeLoggingOnlyScopeChoice(currentChoice, directionalScopeSupported);
if (choice !== value) onChange(choice);
}, [value, directionalScopeSupported, onChange]);
return (
<Select items={items} value={asText(value) || "default"} onValueChange={onChange}>
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy} className="w-full">
<SelectValue placeholder="Select an option" />
</SelectTrigger>
<SelectContent>
{items.map((item) => (
<SelectItem key={item.value} value={item.value}>
{item.label}
</SelectItem>
))}
</SelectContent>
</Select>
);
};
export const LoggingOnlyScopeField: React.FC<{
control: GuardrailFormControl;
mode: unknown;
directionalScopeSupported: boolean;
}> = ({ control, mode, directionalScopeSupported }) => {
if (!modeIncludesLoggingOnly(mode)) return null;
return (
<GuardrailField
control={control}
name="logging_only_scope_choice"
label={labelWithHint(
"Logging only scope",
"Which direction a logging_only scan observes. Observe-only scans never block; pre_call and post_call on this guardrail still block.",
)}
>
{(fieldControl) => (
<LoggingOnlyScopeSelect control={fieldControl} directionalScopeSupported={directionalScopeSupported} />
)}
</GuardrailField>
);
};

View file

@ -0,0 +1,43 @@
import React from "react";
import { Badge } from "@/components/ui/badge";
import { Card } from "@/components/ui/card";
import { formatGuardrailMode, formatLoggingOnlyScope, modeIncludesLoggingOnly } from "./guardrail_info_helpers";
type GuardrailModeParams = {
mode?: unknown;
default_on?: boolean;
logging_only_scope?: string | null;
};
export const GuardrailModeCard: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => (
<Card className="block p-6">
<p>Mode</p>
<div className="mt-2">
<h3 className="text-lg font-medium">{formatGuardrailMode(litellmParams.mode) || "-"}</h3>
<Badge variant={litellmParams.default_on ? "secondary" : "outline"}>
{litellmParams.default_on ? "Default On" : "Default Off"}
</Badge>
</div>
{modeIncludesLoggingOnly(litellmParams.mode) && (
<div className="mt-4">
<p>Logging only scope</p>
<h3 className="text-lg font-medium">{formatLoggingOnlyScope(litellmParams.logging_only_scope)}</h3>
</div>
)}
</Card>
);
export const GuardrailModeRows: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => (
<>
<div>
<p className="font-medium">Mode</p>
<div>{formatGuardrailMode(litellmParams.mode) || "-"}</div>
</div>
{modeIncludesLoggingOnly(litellmParams.mode) && (
<div>
<p className="font-medium">Logging only scope</p>
<div>{formatLoggingOnlyScope(litellmParams.logging_only_scope)}</div>
</div>
)}
</>
);

View file

@ -37,6 +37,7 @@ const uiSettings = {
supported_entities: [],
supported_actions: [],
supported_modes: ["pre_call", "post_call"],
providers_without_directional_logging_only_scope: [],
pii_entity_categories: [],
};
@ -105,6 +106,67 @@ describe("AddGuardrailForm create payload characterization", () => {
expect(payload()).toMatchObject({ litellm_params: { mode: ["pre_call", "post_call"] } });
});
it("sends the selected output logging-only scope", async () => {
vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({
...uiSettings,
supported_modes: ["pre_call", "logging_only"],
});
const user = userEvent.setup({ delay: null });
renderForm();
await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock");
await pickProvider(user, "Bedrock Guardrail");
await user.click(screen.getByLabelText("Mode"));
await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement);
await chooseSelectOption(user, await screen.findByLabelText("Logging only scope"), "Output only (response)");
await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123");
await user.click(screen.getByRole("button", { name: "Next" }));
await user.click(await screen.findByRole("button", { name: "Create Guardrail" }));
await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1));
expect(payload()?.litellm_params.logging_only_scope).toBe("output");
});
it("hides directional scope choices for providers that do not support them", async () => {
vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({
...uiSettings,
supported_modes: ["pre_call", "logging_only"],
providers_without_directional_logging_only_scope: ["xecguard"],
});
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({
...providerParams,
xecguard: { ui_friendly_name: "XecGuard" },
});
const user = userEvent.setup({ delay: null });
renderForm();
await pickProvider(user, "XecGuard");
await user.click(screen.getByLabelText("Mode"));
await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement);
await user.click(await screen.findByLabelText("Logging only scope"));
expect(screen.queryByRole("option", { name: "Input only (request)" })).not.toBeInTheDocument();
expect(screen.queryByRole("option", { name: "Output only (response)" })).not.toBeInTheDocument();
expect(screen.getByRole("option", { name: "Both (request and response)" })).toBeInTheDocument();
});
it("hides logging-only scope and omits it from a pre-call payload", async () => {
const user = userEvent.setup({ delay: null });
renderForm();
await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock");
await pickProvider(user, "Bedrock Guardrail");
expect(screen.queryByLabelText("Logging only scope")).not.toBeInTheDocument();
await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123");
await user.click(screen.getByRole("button", { name: "Next" }));
await user.click(await screen.findByRole("button", { name: "Create Guardrail" }));
await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1));
expect(payload()?.litellm_params).not.toHaveProperty("logging_only_scope");
});
it("blocks Next when the user deselects every mode", async () => {
const user = userEvent.setup({ delay: null });
renderForm();

View file

@ -1,5 +1,5 @@
import React, { useEffect, useMemo, useState } from "react";
import { useForm, type UseFormReturn } from "react-hook-form";
import { useForm, useWatch, type UseFormReturn } from "react-hook-form";
import { toast } from "@/lib/toast";
import {
createGuardrailCall,
@ -10,18 +10,23 @@ import {
import ContentFilterConfiguration from "./content_filter/ContentFilterConfiguration";
import { type CompetitorIntentConfig } from "./content_filter/CompetitorIntentConfiguration";
import {
choiceToLoggingOnlyScope,
choiceToSkipSystemForCreate,
choiceToSkipToolForCreate,
getGuardrailLogo,
getGuardrailProviders,
getSupportedModesForProvider,
guardrail_provider_map,
modeIncludesLoggingOnly,
populateGuardrailProviderMap,
populateGuardrailProviders,
shouldRenderContentFilterConfigSettings,
shouldRenderLLMJudgeFields,
shouldRenderPIIConfigSettings,
supportsDirectionalLoggingOnlyScope,
toModeArray,
type LoggingOnlyScope,
type LoggingOnlyScopeChoice,
} from "./guardrail_info_helpers";
import { Logo } from "@/components/molecules/logo/Logo";
import { MultiSelect } from "@/components/shared/MultiSelect";
@ -49,6 +54,7 @@ import {
requiredRule,
type GuardrailCriterion,
type GuardrailFormValues,
LoggingOnlyScopeField,
SkipMessageSelect,
} from "./GuardrailFormField";
import GuardrailOptionalParams from "./guardrail_optional_params";
@ -90,6 +96,7 @@ interface GuardrailSettings {
supported_actions: string[];
supported_modes: string[];
supported_modes_by_provider?: Record<string, string[]>;
providers_without_directional_logging_only_scope?: string[];
pii_entity_categories: Array<{
category: string;
entities: string[];
@ -160,6 +167,7 @@ type SkipMessageChoice = "inherit" | "yes" | "no";
const INITIAL_VALUES: GuardrailFormValues = {
mode: "pre_call",
default_on: false,
logging_only_scope_choice: "default",
skip_system_message_choice: "inherit",
skip_tool_message_choice: "inherit",
};
@ -199,6 +207,7 @@ interface ProviderParamsResponse {
const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, accessToken, onSuccess, preset }) => {
const form = useForm<GuardrailFormValues>({ defaultValues: INITIAL_VALUES });
const watchedMode = useWatch({ control: form.control, name: "mode" });
const [loading, setLoading] = useState(false);
const [selectedProvider, setSelectedProvider] = useState<string | null>(null);
const [guardrailSettings, setGuardrailSettings] = useState<GuardrailSettings | null>(null);
@ -234,6 +243,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
const providerValue = guardrail_provider_map[selectedProvider];
return (providerValue || "").toLowerCase() === "tool_permission";
}, [selectedProvider]);
const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, selectedProvider);
// Fetch guardrail UI settings + provider params on mount / accessToken change
useEffect(() => {
@ -277,6 +287,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
guardrail_name: preset.guardrailNameSuggestion,
mode: preset.mode,
default_on: preset.defaultOn,
logging_only_scope_choice: "default",
skip_system_message_choice: "inherit",
skip_tool_message_choice: "inherit",
};
@ -439,6 +450,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
guardrail_name: string;
litellm_params: {
guardrail: string;
logging_only_scope?: LoggingOnlyScope | null;
[key: string]: unknown; // Allow dynamic properties
};
guardrail_info: Record<string, unknown>;
@ -462,6 +474,13 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
guardrailData.litellm_params.skip_tool_message_in_guardrail = skipToolForCreate;
}
const loggingOnlyScope = choiceToLoggingOnlyScope(
values.logging_only_scope_choice as LoggingOnlyScopeChoice | undefined,
);
if (modeIncludesLoggingOnly(values.mode) && loggingOnlyScope !== null) {
guardrailData.litellm_params.logging_only_scope = loggingOnlyScope;
}
// For Presidio PII, add the entity and action configurations
if (providerKey === "PresidioPII" && selectedEntities.length > 0) {
const piiEntitiesConfig: { [key: string]: string } = {};
@ -796,6 +815,12 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
{(fieldControl) => <SkipMessageSelect control={fieldControl} />}
</GuardrailField>
<LoggingOnlyScopeField
control={form.control}
mode={watchedMode}
directionalScopeSupported={directionalScopeSupported}
/>
{/* Use the GuardrailProviderFields component to render provider-specific fields */}
{showProviderFields && (
<GuardrailProviderFields

View file

@ -2,6 +2,7 @@ import React, { useState, useRef, useEffect } from "react";
import { CheckCircle2, ChevronRight, Code, ExternalLink, PlayCircle, Save, Users, XCircle } from "lucide-react";
import { createGuardrailCall, updateGuardrailCall, testCustomCodeGuardrail } from "@/components/networking";
import { toast } from "@/lib/toast";
import { type LoggingOnlyScope } from "../guardrail_info_helpers";
import { Button } from "@/components/ui/button";
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
import {
@ -178,6 +179,7 @@ export interface EditGuardrailData {
mode?: string | string[];
default_on?: boolean;
custom_code?: string;
logging_only_scope?: LoggingOnlyScope | null;
[key: string]: any;
};
}

View file

@ -25,6 +25,7 @@ const uiSettings = {
supported_actions: [],
pii_entity_categories: [],
supported_modes: ["pre_call", "post_call"],
providers_without_directional_logging_only_scope: [],
};
const bedrockParams = {
@ -164,6 +165,62 @@ describe("GuardrailInfoView update payload characterization", () => {
expect(lastPayload()).toEqual({ litellm_params: { skip_system_message_in_guardrail: true } });
});
it("shows and updates the logging-only scope", async () => {
const guardrailParams = {
guardrailIdentifier: "gr-abc",
api_key: "sk-old",
mode: "logging_only",
logging_only_scope: "input",
};
vi.mocked(networking.getGuardrailInfo).mockResolvedValue(guardrail(guardrailParams));
const user = userEvent.setup({ delay: null });
renderView();
expect(await screen.findAllByText("Input only (request)")).toHaveLength(2);
await openEditor(user);
await chooseSelectOption(user, screen.getByLabelText("Logging only scope"), "Output only (response)");
await saveChanges(user);
await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1));
expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: "output" } });
});
it("clears the logging-only scope when the edit choice returns to default", async () => {
const guardrailParams = {
guardrailIdentifier: "gr-abc",
api_key: "sk-old",
mode: "logging_only",
logging_only_scope: "input",
};
vi.mocked(networking.getGuardrailInfo).mockResolvedValue(guardrail(guardrailParams));
const user = userEvent.setup({ delay: null });
renderView();
await openEditor(user);
await chooseSelectOption(user, screen.getByLabelText("Logging only scope"), "Default (request and response)");
await saveChanges(user);
await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1));
expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: null } });
});
it("clears a stored directional scope for a provider that does not support it", async () => {
vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({
...uiSettings,
providers_without_directional_logging_only_scope: ["bedrock"],
});
vi.mocked(networking.getGuardrailInfo).mockResolvedValue(
guardrail({ guardrailIdentifier: "gr-abc", mode: "logging_only", logging_only_scope: "output" }),
);
const user = userEvent.setup({ delay: null });
renderView();
await openEditor(user);
await saveChanges(user);
await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1));
expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: null } });
});
it("parses the guardrail information textarea into an object", async () => {
const user = userEvent.setup({ delay: null });
renderView();

View file

@ -28,16 +28,20 @@ import {
readRecord,
requiredRule,
type GuardrailFormValues,
LoggingOnlyScopeField,
SkipMessageSelect,
} from "./GuardrailFormField";
import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager";
import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal";
import { GuardrailModeCard, GuardrailModeRows } from "./GuardrailModeDisplay";
import {
formatGuardrailMode,
getLoggingOnlyScopeUpdate,
getGuardrailLogoAndName,
guardrail_provider_map,
loggingOnlyScopeToChoice,
skipSystemMessageToChoice,
skipToolMessageToChoice,
supportsDirectionalLoggingOnlyScope,
type SkipSystemMessageChoice,
type SkipToolMessageChoice,
} from "./guardrail_info_helpers";
@ -81,6 +85,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
entities: string[];
}>;
supported_modes: string[];
providers_without_directional_logging_only_scope?: string[];
content_filter_settings?: {
prebuilt_patterns: Array<{
name: string;
@ -109,6 +114,8 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
const [toolPermissionConfig, setToolPermissionConfig] = useState<ToolPermissionConfig>(emptyToolPermissionConfig);
const [toolPermissionDirty, setToolPermissionDirty] = useState(false);
const [customCodeModalVisible, setCustomCodeModalVisible] = useState(false);
const guardrailProvider = guardrailData?.litellm_params?.guardrail ?? null;
const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, guardrailProvider);
// Content Filter data ref (managed by ContentFilterManager)
const contentFilterDataRef = React.useRef<{
@ -219,6 +226,8 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
if (!guardrailData) return;
form.setValue("guardrail_name", guardrailData.guardrail_name);
form.setValue("default_on", guardrailData.litellm_params?.default_on);
const storedLoggingOnlyScope = guardrailData.litellm_params?.logging_only_scope;
form.setValue("logging_only_scope_choice", loggingOnlyScopeToChoice(storedLoggingOnlyScope));
form.setValue(
"skip_system_message_choice",
skipSystemMessageToChoice(guardrailData.litellm_params?.skip_system_message_in_guardrail),
@ -282,7 +291,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
// Prepare update data object - only include changed fields
const updateData: any = {
litellm_params: {},
litellm_params: getLoggingOnlyScopeUpdate(guardrailData.litellm_params, values.logging_only_scope_choice),
};
// Only include guardrail_name if it has changed
@ -556,17 +565,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
</div>
</Card>
<Card className="block p-6">
<p>Mode</p>
<div className="mt-2">
<h3 className="text-lg font-medium">
{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}
</h3>
<Badge variant={guardrailData.litellm_params?.default_on ? "secondary" : "outline"}>
{guardrailData.litellm_params?.default_on ? "Default On" : "Default Off"}
</Badge>
</div>
</Card>
<GuardrailModeCard litellmParams={guardrailData.litellm_params} />
<Card className="block p-6">
<p>Created At</p>
@ -745,6 +744,11 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
>
{(fieldControl) => <SkipMessageSelect control={fieldControl} />}
</GuardrailField>
<LoggingOnlyScopeField
control={form.control}
mode={guardrailData.litellm_params?.mode}
directionalScopeSupported={directionalScopeSupported}
/>
{guardrailData.litellm_params?.guardrail === "presidio" && (
<>
<SectionHeading>PII Protection</SectionHeading>
@ -856,10 +860,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
<p className="font-medium">Provider</p>
<div>{displayName}</div>
</div>
<div>
<p className="font-medium">Mode</p>
<div>{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}</div>
</div>
<GuardrailModeRows litellmParams={guardrailData.litellm_params} />
<div>
<p className="font-medium">Default On</p>
<Badge variant={guardrailData.litellm_params?.default_on ? "secondary" : "outline"}>

View file

@ -15,6 +15,14 @@ import {
skipToolMessageToChoice,
choiceToSkipToolForCreate,
formatGuardrailMode,
loggingOnlyScopeToChoice,
choiceToLoggingOnlyScope,
getLoggingOnlyScopeUpdate,
getLoggingOnlyScopeOptions,
formatLoggingOnlyScope,
modeIncludesLoggingOnly,
normalizeLoggingOnlyScopeChoice,
supportsDirectionalLoggingOnlyScope,
} from "./guardrail_info_helpers";
describe("guardrail_info_helpers", () => {
@ -27,6 +35,7 @@ describe("guardrail_info_helpers", () => {
"PresidioPII",
"Bedrock",
"Lakera",
"Xecguard",
"LitellmContentFilter",
"ToolPermission",
"BlockCodeExecution",
@ -239,6 +248,100 @@ describe("guardrail_info_helpers", () => {
});
});
describe("logging-only scope helpers", () => {
it("normalizes directional choices only when the provider does not support them", () => {
expect(normalizeLoggingOnlyScopeChoice("input", false)).toBe("default");
expect(normalizeLoggingOnlyScopeChoice("output", false)).toBe("default");
expect(normalizeLoggingOnlyScopeChoice("both", false)).toBe("both");
expect(normalizeLoggingOnlyScopeChoice("default", false)).toBe("default");
expect(normalizeLoggingOnlyScopeChoice("input", true)).toBe("input");
expect(normalizeLoggingOnlyScopeChoice("output", true)).toBe("output");
});
it("maps API scope values to choices and back", () => {
expect(loggingOnlyScopeToChoice("input")).toBe("input");
expect(loggingOnlyScopeToChoice("output")).toBe("output");
expect(loggingOnlyScopeToChoice("both")).toBe("both");
expect(loggingOnlyScopeToChoice(undefined)).toBe("default");
expect(loggingOnlyScopeToChoice(null)).toBe("default");
expect(loggingOnlyScopeToChoice("invalid")).toBe("default");
expect(choiceToLoggingOnlyScope("default")).toBeNull();
expect(choiceToLoggingOnlyScope(undefined)).toBeNull();
expect(choiceToLoggingOnlyScope("input")).toBe("input");
expect(choiceToLoggingOnlyScope("output")).toBe("output");
expect(choiceToLoggingOnlyScope("both")).toBe("both");
expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "input")).toEqual({});
expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "output")).toEqual({
logging_only_scope: "output",
});
expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "default")).toEqual({
logging_only_scope: null,
});
});
it("formats every scope and falls back to default for missing or unknown values", () => {
expect(formatLoggingOnlyScope("input")).toBe("Input only (request)");
expect(formatLoggingOnlyScope("output")).toBe("Output only (response)");
expect(formatLoggingOnlyScope("both")).toBe("Both (request and response)");
expect(formatLoggingOnlyScope(undefined)).toBe("Default (request and response)");
expect(formatLoggingOnlyScope(null)).toBe("Default (request and response)");
expect(formatLoggingOnlyScope("invalid")).toBe("Default (request and response)");
});
it("detects logging_only in string, array, and tagged mode values", () => {
expect(modeIncludesLoggingOnly("logging_only")).toBe(true);
expect(modeIncludesLoggingOnly(["pre_call", "logging_only"])).toBe(true);
expect(
modeIncludesLoggingOnly({
tags: { "Service-Type: internal-service": "logging_only" },
default: "pre_call",
}),
).toBe(true);
expect(modeIncludesLoggingOnly("pre_call")).toBe(false);
});
it("filters directional options for unsupported providers and keeps all options otherwise", () => {
expect(getLoggingOnlyScopeOptions(false).map((option) => option.value)).toEqual(["default", "both"]);
expect(getLoggingOnlyScopeOptions(true).map((option) => option.value)).toEqual([
"default",
"input",
"output",
"both",
]);
expect(
supportsDirectionalLoggingOnlyScope(
{ providers_without_directional_logging_only_scope: ["xecguard"] },
"Xecguard",
),
).toBe(false);
expect(
supportsDirectionalLoggingOnlyScope(
{ providers_without_directional_logging_only_scope: ["xecguard"] },
"xecguard",
),
).toBe(false);
expect(
supportsDirectionalLoggingOnlyScope(
{ providers_without_directional_logging_only_scope: ["xecguard"] },
"Bedrock",
),
).toBe(true);
expect(
supportsDirectionalLoggingOnlyScope(
{ providers_without_directional_logging_only_scope: ["xecguard"] },
"unknown-provider",
),
).toBe(true);
expect(supportsDirectionalLoggingOnlyScope(null, "Xecguard")).toBe(true);
expect(getLoggingOnlyScopeOptions(supportsDirectionalLoggingOnlyScope(null, "Xecguard"))).toEqual(
getLoggingOnlyScopeOptions(true),
);
});
});
describe("skipSystemMessageToChoice / choiceToSkipSystemForCreate", () => {
it("maps API values to form choices and back for create", () => {
expect(skipSystemMessageToChoice(undefined)).toBe("inherit");

View file

@ -114,6 +114,74 @@ export const toModeArray = (raw: unknown): string[] => {
return [];
};
export type LoggingOnlyScope = "input" | "output" | "both";
export type LoggingOnlyScopeChoice = "default" | LoggingOnlyScope;
export type LoggingOnlyScopeOption = { label: string; value: LoggingOnlyScopeChoice };
export const normalizeLoggingOnlyScopeChoice = (
choice: LoggingOnlyScopeChoice,
directionalScopeSupported: boolean,
): LoggingOnlyScopeChoice =>
directionalScopeSupported || choice === "default" || choice === "both" ? choice : "default";
const LOGGING_ONLY_SCOPE_OPTIONS: LoggingOnlyScopeOption[] = [
{ label: "Default (request and response)", value: "default" },
{ label: "Input only (request)", value: "input" },
{ label: "Output only (response)", value: "output" },
{ label: "Both (request and response)", value: "both" },
];
export const loggingOnlyScopeToChoice = (v: string | null | undefined): LoggingOnlyScopeChoice =>
v === "input" || v === "output" || v === "both" ? v : "default";
export const choiceToLoggingOnlyScope = (choice: LoggingOnlyScopeChoice | undefined): LoggingOnlyScope | null =>
choice === "input" || choice === "output" || choice === "both" ? choice : null;
export const getLoggingOnlyScopeUpdate = (
litellmParams: { logging_only_scope?: string | null } | null | undefined,
choice: LoggingOnlyScopeChoice | undefined,
): { logging_only_scope?: LoggingOnlyScope | null } => {
if (choice === undefined || choice === loggingOnlyScopeToChoice(litellmParams?.logging_only_scope)) return {};
return { logging_only_scope: choiceToLoggingOnlyScope(choice) };
};
export const formatLoggingOnlyScope = (v: string | null | undefined): string => {
if (v === "input") return "Input only (request)";
if (v === "output") return "Output only (response)";
if (v === "both") return "Both (request and response)";
return "Default (request and response)";
};
export const modeIncludesLoggingOnly = (raw: unknown): boolean => {
if (toModeArray(raw).includes("logging_only")) return true;
if (raw === null || typeof raw !== "object") return false;
const { tags, default: fallback } = raw as { tags?: Record<string, unknown>; default?: unknown };
const taggedModes =
tags && typeof tags === "object"
? Object.values(tags).some((mode) => toModeArray(mode).includes("logging_only"))
: false;
return toModeArray(fallback).includes("logging_only") || taggedModes;
};
export const getLoggingOnlyScopeOptions = (directionalScopeSupported: boolean): LoggingOnlyScopeOption[] =>
directionalScopeSupported
? LOGGING_ONLY_SCOPE_OPTIONS
: LOGGING_ONLY_SCOPE_OPTIONS.filter((option) => option.value === "default" || option.value === "both");
export const supportsDirectionalLoggingOnlyScope = (
settings: { providers_without_directional_logging_only_scope?: string[] } | null,
selectedProvider: string | null,
): boolean => {
const providerKey = selectedProvider
? (
guardrail_provider_map[selectedProvider] ??
Object.values(guardrail_provider_map).find((value) => value.toLowerCase() === selectedProvider.toLowerCase())
)?.toLowerCase()
: null;
return !providerKey || !settings?.providers_without_directional_logging_only_scope?.includes(providerKey);
};
export const formatGuardrailMode = (raw: unknown): string => {
const flat: string[] = toModeArray(raw);
if (flat.length > 0) return flat.join(", ");

View file

@ -25997,6 +25997,11 @@ export interface components {
* @description Google Cloud location/region (e.g., us-central1)
*/
location?: string | null;
/**
* Logging Only Scope
* @description which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.
*/
logging_only_scope?: ("input" | "output" | "both") | null;
/**
* Mask Request Content
* @description Will mask request content if guardrail makes any changes
@ -32032,6 +32037,27 @@ export interface components {
/** Output Text */
output_text: string;
};
/** GuardrailUIAddGuardrailSettings */
GuardrailUIAddGuardrailSettings: {
/** Content Filter Settings */
content_filter_settings?: {
[key: string]: unknown;
} | null;
/** Pii Entity Categories */
pii_entity_categories: components["schemas"]["PiiEntityCategoryMap"][];
/** Providers Without Directional Logging Only Scope */
providers_without_directional_logging_only_scope: string[];
/** Supported Actions */
supported_actions: string[];
/** Supported Entities */
supported_entities: string[];
/** Supported Modes */
supported_modes: string[];
/** Supported Modes By Provider */
supported_modes_by_provider: {
[key: string]: string[];
};
};
/**
* HTTPAuthSecurityScheme
* @description Defines a security scheme using HTTP authentication.
@ -35456,6 +35482,11 @@ export interface components {
* @description Google Cloud location/region (e.g., us-central1)
*/
location?: string | null;
/**
* Logging Only Scope
* @description which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.
*/
logging_only_scope?: ("input" | "output" | "both") | null;
/**
* Mask
* @description Enable content masking using Lasso classifix API
@ -38978,6 +39009,13 @@ export interface components {
* @enum {string}
*/
PiiAction: "BLOCK" | "MASK";
/** PiiEntityCategoryMap */
PiiEntityCategoryMap: {
/** Category */
category: string;
/** Entities */
entities: string[];
};
/**
* PiiEntityType
* @enum {string}
@ -58455,7 +58493,7 @@ export interface operations {
[name: string]: unknown;
};
content: {
"application/json": unknown;
"application/json": components["schemas"]["GuardrailUIAddGuardrailSettings"];
};
};
};