mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): support directional scope in PATCH and provider UI
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
089608554d
commit
6cbb57a797
13 changed files with 348 additions and 45 deletions
|
|
@ -806,10 +806,11 @@ class CustomGuardrail(CustomLogger):
|
|||
def uses_apply_guardrail_interface(self) -> bool:
|
||||
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
|
||||
|
||||
def supports_logging_only_scope(self) -> bool:
|
||||
@classmethod
|
||||
def supports_logging_only_scope(cls) -> bool:
|
||||
return (
|
||||
self.uses_apply_guardrail_interface()
|
||||
and type(self).async_logging_hook is CustomGuardrail.async_logging_hook
|
||||
cls.apply_guardrail is not CustomGuardrail.apply_guardrail
|
||||
and cls.async_logging_hook is CustomGuardrail.async_logging_hook
|
||||
)
|
||||
|
||||
def _deployment_hook_target(self) -> "CustomLogger":
|
||||
|
|
@ -992,13 +993,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
def _copy_scratch_request_fields(
|
||||
self,
|
||||
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
|
||||
kwargs: Mapping[str, object],
|
||||
) -> tuple[object, object]:
|
||||
optional_params: Final = kwargs.get("optional_params") or {}
|
||||
optional_params: Final = kwargs.get("optional_params")
|
||||
try:
|
||||
return (
|
||||
copy.deepcopy(kwargs.get("messages") or kwargs.get("input")),
|
||||
copy.deepcopy(optional_params.get("tools")),
|
||||
copy.deepcopy(optional_params.get("tools") if isinstance(optional_params, Mapping) else None),
|
||||
)
|
||||
except Exception:
|
||||
if self.logging_only_scope == "output":
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ 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
|
||||
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
|
||||
|
|
@ -1212,19 +1212,30 @@ 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 = LitellmParams(**existing_litellm_params)
|
||||
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 = (
|
||||
|
|
@ -1253,7 +1264,7 @@ async def patch_guardrail(
|
|||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=guardrail,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
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
|
||||
|
|
@ -1398,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
|
||||
|
||||
|
|
@ -1432,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(),
|
||||
|
|
|
|||
|
|
@ -711,7 +711,7 @@ 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):
|
||||
if reject_invalid_logging_only_scope or self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
|
||||
self.reinitialize_guardrail(
|
||||
guardrail=updated_guardrail,
|
||||
source=source,
|
||||
|
|
@ -936,7 +936,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
|
||||
|
|
|
|||
|
|
@ -1314,6 +1314,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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -9,6 +10,7 @@ import pytest
|
|||
|
||||
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,
|
||||
|
|
@ -42,10 +44,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 = {
|
||||
|
|
@ -85,6 +89,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"""
|
||||
|
|
@ -124,6 +166,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"""
|
||||
|
|
@ -1257,7 +1330,7 @@ async def test_patch_guardrail_endpoint(
|
|||
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
|
||||
guardrail=mocker.ANY,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
reject_invalid_logging_only_scope=False,
|
||||
)
|
||||
|
||||
if scenario == "success_sync_fails_unexpected_error":
|
||||
|
|
@ -1310,6 +1383,103 @@ async def test_patch_guardrail_rejects_invalid_logging_only_scope_with_422(mocke
|
|||
)
|
||||
|
||||
|
||||
@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_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"].logging_only_scope == "output"
|
||||
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",
|
||||
[
|
||||
|
|
@ -2495,6 +2665,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()
|
||||
|
|
|
|||
|
|
@ -2996,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",
|
||||
|
|
|
|||
|
|
@ -6,7 +6,11 @@ import { useController, type Control, type ControllerRenderProps, type RegisterO
|
|||
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 { modeIncludesLoggingOnly, type LoggingOnlyScopeChoice } from "./guardrail_info_helpers";
|
||||
import {
|
||||
getLoggingOnlyScopeOptions,
|
||||
modeIncludesLoggingOnly,
|
||||
type LoggingOnlyScopeChoice,
|
||||
} from "./guardrail_info_helpers";
|
||||
|
||||
export interface GuardrailCriterion {
|
||||
name: string;
|
||||
|
|
@ -126,23 +130,20 @@ export const SkipMessageSelect: React.FC<{ control: GuardrailFieldControlProps }
|
|||
);
|
||||
};
|
||||
|
||||
const LOGGING_ONLY_SCOPE_ITEMS: Array<{ label: string; value: LoggingOnlyScopeChoice }> = [
|
||||
{ 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 LoggingOnlyScopeSelect: React.FC<{ control: GuardrailFieldControlProps }> = ({ control }) => {
|
||||
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);
|
||||
|
||||
return (
|
||||
<Select items={LOGGING_ONLY_SCOPE_ITEMS} value={asText(value) || "default"} onValueChange={onChange}>
|
||||
<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>
|
||||
{LOGGING_ONLY_SCOPE_ITEMS.map((item) => (
|
||||
{items.map((item) => (
|
||||
<SelectItem key={item.value} value={item.value}>
|
||||
{item.label}
|
||||
</SelectItem>
|
||||
|
|
@ -152,10 +153,11 @@ export const LoggingOnlyScopeSelect: React.FC<{ control: GuardrailFieldControlPr
|
|||
);
|
||||
};
|
||||
|
||||
export const LoggingOnlyScopeField: React.FC<{ control: GuardrailFormControl; mode: unknown }> = ({
|
||||
control,
|
||||
mode,
|
||||
}) => {
|
||||
export const LoggingOnlyScopeField: React.FC<{
|
||||
control: GuardrailFormControl;
|
||||
mode: unknown;
|
||||
directionalScopeSupported: boolean;
|
||||
}> = ({ control, mode, directionalScopeSupported }) => {
|
||||
if (!modeIncludesLoggingOnly(mode)) return null;
|
||||
|
||||
return (
|
||||
|
|
@ -167,7 +169,9 @@ export const LoggingOnlyScopeField: React.FC<{ control: GuardrailFormControl; mo
|
|||
"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} />}
|
||||
{(fieldControl) => (
|
||||
<LoggingOnlyScopeSelect control={fieldControl} directionalScopeSupported={directionalScopeSupported} />
|
||||
)}
|
||||
</GuardrailField>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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: [],
|
||||
};
|
||||
|
||||
|
|
@ -131,6 +132,29 @@ describe("AddGuardrailForm create payload characterization", () => {
|
|||
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();
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ import {
|
|||
shouldRenderContentFilterConfigSettings,
|
||||
shouldRenderLLMJudgeFields,
|
||||
shouldRenderPIIConfigSettings,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
toModeArray,
|
||||
type LoggingOnlyScope,
|
||||
type LoggingOnlyScopeChoice,
|
||||
|
|
@ -95,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[];
|
||||
|
|
@ -688,6 +690,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
|
|||
const providerLabels: Record<string, string> = getGuardrailProviders();
|
||||
const providerKeys = Object.keys(providerLabels);
|
||||
const supportedModes = getSupportedModesForProvider(guardrailSettings, selectedProvider) ?? DEFAULT_MODES;
|
||||
const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, selectedProvider);
|
||||
return (
|
||||
<FieldGroup>
|
||||
<GuardrailField
|
||||
|
|
@ -812,7 +815,11 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
|
|||
{(fieldControl) => <SkipMessageSelect control={fieldControl} />}
|
||||
</GuardrailField>
|
||||
|
||||
<LoggingOnlyScopeField control={form.control} mode={watchedMode} />
|
||||
<LoggingOnlyScopeField
|
||||
control={form.control}
|
||||
mode={watchedMode}
|
||||
directionalScopeSupported={directionalScopeSupported}
|
||||
/>
|
||||
|
||||
{/* Use the GuardrailProviderFields component to render provider-specific fields */}
|
||||
{showProviderFields && (
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ import {
|
|||
loggingOnlyScopeToChoice,
|
||||
skipSystemMessageToChoice,
|
||||
skipToolMessageToChoice,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
type SkipSystemMessageChoice,
|
||||
type SkipToolMessageChoice,
|
||||
} from "./guardrail_info_helpers";
|
||||
|
|
@ -84,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;
|
||||
|
|
@ -112,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<{
|
||||
|
|
@ -742,6 +746,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
<LoggingOnlyScopeField
|
||||
control={form.control}
|
||||
mode={guardrailData.litellm_params?.mode}
|
||||
directionalScopeSupported={directionalScopeSupported}
|
||||
/>
|
||||
{guardrailData.litellm_params?.guardrail === "presidio" && (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -18,8 +18,10 @@ import {
|
|||
loggingOnlyScopeToChoice,
|
||||
choiceToLoggingOnlyScope,
|
||||
getLoggingOnlyScopeUpdate,
|
||||
getLoggingOnlyScopeOptions,
|
||||
formatLoggingOnlyScope,
|
||||
modeIncludesLoggingOnly,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
} from "./guardrail_info_helpers";
|
||||
|
||||
describe("guardrail_info_helpers", () => {
|
||||
|
|
@ -32,6 +34,7 @@ describe("guardrail_info_helpers", () => {
|
|||
"PresidioPII",
|
||||
"Bedrock",
|
||||
"Lakera",
|
||||
"Xecguard",
|
||||
"LitellmContentFilter",
|
||||
"ToolPermission",
|
||||
"BlockCodeExecution",
|
||||
|
|
@ -288,6 +291,45 @@ describe("guardrail_info_helpers", () => {
|
|||
).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", () => {
|
||||
|
|
|
|||
|
|
@ -116,6 +116,14 @@ export const toModeArray = (raw: unknown): string[] => {
|
|||
|
||||
export type LoggingOnlyScope = "input" | "output" | "both";
|
||||
export type LoggingOnlyScopeChoice = "default" | LoggingOnlyScope;
|
||||
export type LoggingOnlyScopeOption = { label: string; value: LoggingOnlyScopeChoice };
|
||||
|
||||
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";
|
||||
|
|
@ -150,6 +158,24 @@ export const modeIncludesLoggingOnly = (raw: unknown): boolean => {
|
|||
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(", ");
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue