mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge eb733bdd27 into ed4caebb65
This commit is contained in:
commit
956cbd5b04
23 changed files with 6485 additions and 78 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
1068
tests/integration/observability/_logging_only_scope_support.py
Normal file
1068
tests/integration/observability/_logging_only_scope_support.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
762
tests/integration/observability/test_logging_only_scope_chaos.py
Normal file
762
tests/integration/observability/test_logging_only_scope_chaos.py
Normal 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,
|
||||
)
|
||||
1516
tests/integration/observability/test_logging_only_scope_config.py
Normal file
1516
tests/integration/observability/test_logging_only_scope_config.py
Normal file
File diff suppressed because it is too large
Load diff
1579
tests/integration/observability/test_logging_only_scope_runtime.py
Normal file
1579
tests/integration/observability/test_logging_only_scope_runtime.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
</>
|
||||
);
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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"}>
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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(", ");
|
||||
|
|
|
|||
40
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
40
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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"];
|
||||
};
|
||||
};
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue