mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(policy_endpoints.py): expose new endpoint for testing policies and guardrails
enables compliance playground to work as expected
This commit is contained in:
parent
c5394e7c3e
commit
7f6f6bc7ba
4 changed files with 539 additions and 91 deletions
|
|
@ -11,9 +11,10 @@ All /policy management endpoints
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, List, Literal, Optional, TypedDict, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -41,48 +42,87 @@ if TYPE_CHECKING:
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
class GuardrailApplyError(Exception):
|
||||
"""
|
||||
Raised when a guardrail's apply_guardrail fails during apply_policies.
|
||||
|
||||
Consumers (e.g. Compliance UI) can use guardrail_name and message to show
|
||||
which guardrail triggered and the error reason.
|
||||
"""
|
||||
|
||||
def __init__(self, guardrail_name: str, message: str) -> None:
|
||||
self.guardrail_name = guardrail_name
|
||||
self.message = message
|
||||
super().__init__(f"Guardrail '{guardrail_name}' failed: {message}")
|
||||
|
||||
|
||||
class GuardrailErrorEntry(TypedDict):
|
||||
"""One guardrail failure for ApplyPoliciesResult.guardrail_errors."""
|
||||
|
||||
guardrail_name: str
|
||||
message: str
|
||||
|
||||
|
||||
class ApplyPoliciesResult(TypedDict):
|
||||
"""Result of apply_policies: inputs plus any guardrail failures."""
|
||||
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
guardrail_errors: List[GuardrailErrorEntry]
|
||||
|
||||
|
||||
async def apply_policies(
|
||||
policy_names: Optional[list[str]],
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
proxy_logging_obj: "LiteLLMLoggingObj",
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
guardrail_names: Optional[list[str]] = None,
|
||||
) -> ApplyPoliciesResult:
|
||||
"""
|
||||
Resolve guardrails from the given policy names and apply them to inputs.
|
||||
Apply guardrails to inputs from policy names and/or a direct list of guardrail names.
|
||||
|
||||
Similar to add_guardrails_from_policy_engine + guardrail execution: resolves
|
||||
guardrails from the policy registry (with inheritance) and runs each
|
||||
guardrail's apply_guardrail on the inputs in order.
|
||||
Runs all guardrails in order; if one fails, the error is recorded and execution
|
||||
continues so that all inputs can complete testing and all guardrail failures are
|
||||
collected. No exception is raised; failures are returned in guardrail_errors.
|
||||
|
||||
Guardrails can be specified in two ways (both can be used together; names are merged):
|
||||
- policy_names: resolve guardrails from the policy registry (with inheritance).
|
||||
- guardrail_names: use this list of guardrail names directly (no policy registry needed).
|
||||
|
||||
Returns:
|
||||
ApplyPoliciesResult with "inputs" (final GenericGuardrailAPIInputs) and
|
||||
"guardrail_errors" (list of {"guardrail_name", "message"} for each failure).
|
||||
"""
|
||||
if not policy_names:
|
||||
return inputs
|
||||
guardrail_errors: List[GuardrailErrorEntry] = []
|
||||
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: policy engine not initialized, returning inputs unchanged"
|
||||
)
|
||||
return inputs
|
||||
guardrail_name_set: set[str] = set()
|
||||
|
||||
policies = registry.get_all_policies()
|
||||
guardrail_names: set[str] = set()
|
||||
if guardrail_names:
|
||||
guardrail_name_set.update(guardrail_names)
|
||||
|
||||
for policy_name in policy_names:
|
||||
resolved = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name,
|
||||
policies=policies,
|
||||
context=None,
|
||||
)
|
||||
guardrail_names.update(resolved.guardrails)
|
||||
if policy_names:
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: policy engine not initialized, skipping policy-resolved guardrails"
|
||||
)
|
||||
else:
|
||||
policies = registry.get_all_policies()
|
||||
for policy_name in policy_names:
|
||||
resolved = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name,
|
||||
policies=policies,
|
||||
context=None,
|
||||
)
|
||||
guardrail_name_set.update(resolved.guardrails)
|
||||
|
||||
if not guardrail_names:
|
||||
return inputs
|
||||
if not guardrail_name_set:
|
||||
return {"inputs": inputs, "guardrail_errors": guardrail_errors}
|
||||
|
||||
guardrail_registry = GuardrailRegistry()
|
||||
current_inputs = cast(GenericGuardrailAPIInputs, dict(inputs))
|
||||
|
||||
for guardrail_name in sorted(guardrail_names):
|
||||
for guardrail_name in sorted(guardrail_name_set):
|
||||
callback = guardrail_registry.get_initialized_guardrail_callback(
|
||||
guardrail_name=guardrail_name
|
||||
)
|
||||
|
|
@ -101,14 +141,78 @@ async def apply_policies(
|
|||
)
|
||||
continue
|
||||
|
||||
current_inputs = await callback.apply_guardrail(
|
||||
inputs=current_inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
logging_obj=proxy_logging_obj,
|
||||
)
|
||||
try:
|
||||
current_inputs = await callback.apply_guardrail(
|
||||
inputs=current_inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
error_reason = str(e)
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: guardrail '%s' failed: %s",
|
||||
guardrail_name,
|
||||
error_reason,
|
||||
)
|
||||
guardrail_errors.append(
|
||||
GuardrailErrorEntry(
|
||||
guardrail_name=guardrail_name,
|
||||
message=error_reason,
|
||||
)
|
||||
)
|
||||
# Continue to next guardrail; current_inputs unchanged for this failure
|
||||
|
||||
return current_inputs
|
||||
return {"inputs": current_inputs, "guardrail_errors": guardrail_errors}
|
||||
|
||||
|
||||
class TestPoliciesAndGuardrailsRequest(BaseModel):
|
||||
"""Request body for POST /utils/test_policies_and_guardrails."""
|
||||
|
||||
policy_names: Optional[List[str]] = Field(default=None, description="Policy names to resolve guardrails from")
|
||||
guardrail_names: Optional[List[str]] = Field(default=None, description="Guardrail names to apply directly")
|
||||
inputs: dict = Field(description="GenericGuardrailAPIInputs, e.g. { \"texts\": [\"...\"] }")
|
||||
request_data: dict = Field(default_factory=dict, description="Request context (model, user_id, etc.)")
|
||||
input_type: Literal["request", "response"] = Field(default="request", description="Whether inputs are request or response")
|
||||
|
||||
|
||||
@router.post(
|
||||
"/utils/test_policies_and_guardrails",
|
||||
tags=["utils"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def test_policies_and_guardrails(
|
||||
request: Request,
|
||||
data: TestPoliciesAndGuardrailsRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Apply policies and/or guardrails to inputs (for compliance UI testing).
|
||||
|
||||
Runs all guardrails in order; failures are collected and returned in guardrail_errors.
|
||||
Returns inputs (possibly modified) and any guardrail errors so the UI can show which
|
||||
guardrails failed and why.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
|
||||
try:
|
||||
inputs_typed = cast(GenericGuardrailAPIInputs, data.inputs)
|
||||
logging_obj = cast(LiteLLMLoggingObj, proxy_logging_obj)
|
||||
result = await apply_policies(
|
||||
policy_names=data.policy_names,
|
||||
inputs=inputs_typed,
|
||||
request_data=data.request_data,
|
||||
input_type=data.input_type,
|
||||
proxy_logging_obj=logging_obj,
|
||||
guardrail_names=data.guardrail_names,
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -58,7 +58,8 @@ class TestApplyPoliciesEarlyReturn:
|
|||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_inputs_unchanged_when_policy_names_empty(
|
||||
|
|
@ -71,7 +72,23 @@ class TestApplyPoliciesEarlyReturn:
|
|||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_inputs_unchanged_when_both_policy_and_guardrail_names_empty(
|
||||
self, sample_inputs, request_data, proxy_logging_obj
|
||||
):
|
||||
result = await apply_policies(
|
||||
policy_names=[],
|
||||
inputs=sample_inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
guardrail_names=[],
|
||||
)
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_inputs_unchanged_when_registry_not_initialized(
|
||||
|
|
@ -92,7 +109,8 @@ class TestApplyPoliciesEarlyReturn:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
mock_registry.is_initialized.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -120,7 +138,8 @@ class TestApplyPoliciesEarlyReturn:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
|
||||
class TestApplyPoliciesWithGuardrails:
|
||||
|
|
@ -165,7 +184,8 @@ class TestApplyPoliciesWithGuardrails:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == modified_inputs
|
||||
assert result["inputs"] == modified_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_applies_multiple_guardrails_in_order(
|
||||
|
|
@ -217,7 +237,8 @@ class TestApplyPoliciesWithGuardrails:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == second_output
|
||||
assert result["inputs"] == second_output
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_missing_guardrail_callback(
|
||||
|
|
@ -254,7 +275,58 @@ class TestApplyPoliciesWithGuardrails:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_records_guardrail_error_on_failure(
|
||||
self, sample_inputs, request_data, proxy_logging_obj
|
||||
):
|
||||
"""When a guardrail's apply_guardrail raises, error is recorded and inputs still returned."""
|
||||
from litellm.types.proxy.policy_engine import ResolvedPolicy
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.is_initialized.return_value = True
|
||||
mock_registry.get_all_policies.return_value = {}
|
||||
|
||||
callback = _FakeGuardrailWithApply(guardrail_name="failing_guardrail")
|
||||
|
||||
async def _raise(inputs, request_data, input_type, logging_obj=None):
|
||||
raise ValueError("Content blocked: PII detected")
|
||||
|
||||
callback.apply_guardrail = _raise
|
||||
|
||||
mock_guardrail_registry = MagicMock()
|
||||
mock_guardrail_registry.get_initialized_guardrail_callback.return_value = (
|
||||
callback
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.get_policy_registry",
|
||||
return_value=mock_registry,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.PolicyResolver.resolve_policy_guardrails",
|
||||
return_value=ResolvedPolicy(
|
||||
policy_name="p",
|
||||
guardrails=["failing_guardrail"],
|
||||
inheritance_chain=["p"],
|
||||
),
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.GuardrailRegistry",
|
||||
return_value=mock_guardrail_registry,
|
||||
):
|
||||
result = await apply_policies(
|
||||
policy_names=["my-policy"],
|
||||
inputs=sample_inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == [
|
||||
{"guardrail_name": "failing_guardrail", "message": "Content blocked: PII detected"}
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_callback_without_apply_guardrail(
|
||||
|
|
@ -301,7 +373,73 @@ class TestApplyPoliciesWithGuardrails:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == sample_inputs
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collects_all_guardrail_failures_when_multiple_fail(
|
||||
self, sample_inputs, request_data, proxy_logging_obj
|
||||
):
|
||||
"""When multiple guardrails raise, all failures are collected and inputs still returned."""
|
||||
from litellm.types.proxy.policy_engine import ResolvedPolicy
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.is_initialized.return_value = True
|
||||
mock_registry.get_all_policies.return_value = {}
|
||||
|
||||
callback_a = _FakeGuardrailWithApply(guardrail_name="guardrail_a")
|
||||
|
||||
async def _raise_a(inputs, request_data, input_type, logging_obj=None):
|
||||
raise ValueError("PII detected")
|
||||
|
||||
callback_a.apply_guardrail = _raise_a
|
||||
|
||||
callback_b = _FakeGuardrailWithApply(guardrail_name="guardrail_b")
|
||||
|
||||
async def _raise_b(inputs, request_data, input_type, logging_obj=None):
|
||||
raise RuntimeError("Toxicity detected")
|
||||
|
||||
callback_b.apply_guardrail = _raise_b
|
||||
|
||||
def get_callback(guardrail_name):
|
||||
if guardrail_name == "guardrail_a":
|
||||
return callback_a
|
||||
if guardrail_name == "guardrail_b":
|
||||
return callback_b
|
||||
return None
|
||||
|
||||
mock_guardrail_registry = MagicMock()
|
||||
mock_guardrail_registry.get_initialized_guardrail_callback.side_effect = (
|
||||
get_callback
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.get_policy_registry",
|
||||
return_value=mock_registry,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.PolicyResolver.resolve_policy_guardrails",
|
||||
return_value=ResolvedPolicy(
|
||||
policy_name="p",
|
||||
guardrails=["guardrail_a", "guardrail_b"],
|
||||
inheritance_chain=["p"],
|
||||
),
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.GuardrailRegistry",
|
||||
return_value=mock_guardrail_registry,
|
||||
):
|
||||
result = await apply_policies(
|
||||
policy_names=["my-policy"],
|
||||
inputs=sample_inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result["inputs"] == sample_inputs
|
||||
assert len(result["guardrail_errors"]) == 2
|
||||
by_name = {e["guardrail_name"]: e["message"] for e in result["guardrail_errors"]}
|
||||
assert by_name["guardrail_a"] == "PII detected"
|
||||
assert by_name["guardrail_b"] == "Toxicity detected"
|
||||
|
||||
|
||||
class TestApplyPoliciesMultiplePolicies:
|
||||
|
|
@ -355,4 +493,97 @@ class TestApplyPoliciesMultiplePolicies:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert result == final_inputs
|
||||
assert result["inputs"] == final_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
|
||||
class TestApplyPoliciesDirectGuardrailNames:
|
||||
"""Test apply_policies with direct guardrail_names (no policy registry)."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_applies_guardrails_from_direct_guardrail_names_only(
|
||||
self, sample_inputs, request_data, proxy_logging_obj
|
||||
):
|
||||
"""When only guardrail_names is passed, policy registry is not used."""
|
||||
modified_inputs: GenericGuardrailAPIInputs = {"texts": ["from direct guardrail"]}
|
||||
callback = _FakeGuardrailWithApply(guardrail_name="my_guardrail")
|
||||
callback.set_return(modified_inputs)
|
||||
|
||||
mock_guardrail_registry = MagicMock()
|
||||
mock_guardrail_registry.get_initialized_guardrail_callback.return_value = callback
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.GuardrailRegistry",
|
||||
return_value=mock_guardrail_registry,
|
||||
):
|
||||
result = await apply_policies(
|
||||
policy_names=None,
|
||||
inputs=sample_inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
guardrail_names=["my_guardrail"],
|
||||
)
|
||||
|
||||
assert result["inputs"] == modified_inputs
|
||||
assert result["guardrail_errors"] == []
|
||||
mock_guardrail_registry.get_initialized_guardrail_callback.assert_called_once_with(
|
||||
guardrail_name="my_guardrail"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_applies_guardrails_from_both_policy_names_and_guardrail_names(
|
||||
self, sample_inputs, request_data, proxy_logging_obj
|
||||
):
|
||||
"""Guardrails from policy_names and guardrail_names are merged and applied."""
|
||||
from litellm.types.proxy.policy_engine import ResolvedPolicy
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.is_initialized.return_value = True
|
||||
mock_registry.get_all_policies.return_value = {}
|
||||
|
||||
first_output: GenericGuardrailAPIInputs = {"texts": ["after first"]}
|
||||
second_output: GenericGuardrailAPIInputs = {"texts": ["after second"]}
|
||||
callback_from_policy = _FakeGuardrailWithApply(guardrail_name="from_policy")
|
||||
callback_from_policy.set_return(first_output)
|
||||
callback_direct = _FakeGuardrailWithApply(guardrail_name="direct_guardrail")
|
||||
callback_direct.set_return(second_output)
|
||||
|
||||
def get_callback(guardrail_name):
|
||||
if guardrail_name == "from_policy":
|
||||
return callback_from_policy
|
||||
if guardrail_name == "direct_guardrail":
|
||||
return callback_direct
|
||||
return None
|
||||
|
||||
mock_guardrail_registry = MagicMock()
|
||||
mock_guardrail_registry.get_initialized_guardrail_callback.side_effect = (
|
||||
get_callback
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.get_policy_registry",
|
||||
return_value=mock_registry,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.PolicyResolver.resolve_policy_guardrails",
|
||||
return_value=ResolvedPolicy(
|
||||
policy_name="p",
|
||||
guardrails=["from_policy"],
|
||||
inheritance_chain=["p"],
|
||||
),
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.policy_endpoints.GuardrailRegistry",
|
||||
return_value=mock_guardrail_registry,
|
||||
):
|
||||
result = await apply_policies(
|
||||
policy_names=["my-policy"],
|
||||
inputs=sample_inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
guardrail_names=["direct_guardrail"],
|
||||
)
|
||||
|
||||
# Sorted order: direct_guardrail then from_policy; final output is from_policy
|
||||
assert result["inputs"] == first_output
|
||||
assert result["guardrail_errors"] == []
|
||||
|
|
|
|||
|
|
@ -5438,6 +5438,68 @@ export const getPoliciesList = async (accessToken: string) => {
|
|||
}
|
||||
};
|
||||
|
||||
export interface TestPoliciesAndGuardrailsRequest {
|
||||
policy_names?: string[] | null;
|
||||
guardrail_names?: string[] | null;
|
||||
inputs: { texts?: string[]; images?: string[]; [key: string]: unknown };
|
||||
request_data?: Record<string, unknown>;
|
||||
input_type?: "request" | "response";
|
||||
}
|
||||
|
||||
export interface GuardrailErrorEntry {
|
||||
guardrail_name: string;
|
||||
message: string;
|
||||
}
|
||||
|
||||
export interface TestPoliciesAndGuardrailsResponse {
|
||||
inputs: Record<string, unknown>;
|
||||
guardrail_errors: GuardrailErrorEntry[];
|
||||
}
|
||||
|
||||
export const testPoliciesAndGuardrails = async (
|
||||
accessToken: string,
|
||||
body: TestPoliciesAndGuardrailsRequest
|
||||
): Promise<TestPoliciesAndGuardrailsResponse> => {
|
||||
try {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/utils/test_policies_and_guardrails`
|
||||
: `/utils/test_policies_and_guardrails`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
policy_names: body.policy_names ?? null,
|
||||
guardrail_names: body.guardrail_names ?? null,
|
||||
inputs: body.inputs,
|
||||
request_data: body.request_data ?? {},
|
||||
input_type: body.input_type ?? "request",
|
||||
}),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
let errorMessage = "Failed to test policies and guardrails";
|
||||
try {
|
||||
const errorJson = JSON.parse(errorData);
|
||||
if (errorJson.detail) errorMessage = typeof errorJson.detail === "string" ? errorJson.detail : JSON.stringify(errorJson.detail);
|
||||
else if (errorJson.message) errorMessage = errorJson.message;
|
||||
} catch {
|
||||
errorMessage = errorData || errorMessage;
|
||||
}
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return await response.json();
|
||||
} catch (error) {
|
||||
console.error("Failed to test policies and guardrails:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const getPolicyInfoWithGuardrails = async (accessToken: string, policyName: string) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/policy/info/${policyName}` : `/policy/info/${policyName}`;
|
||||
|
|
|
|||
|
|
@ -6,7 +6,11 @@ import {
|
|||
type ComplianceFramework,
|
||||
type CompliancePrompt,
|
||||
} from "@/data/compliancePrompts";
|
||||
import { getGuardrailsList, getPoliciesList } from "@/components/networking";
|
||||
import {
|
||||
getGuardrailsList,
|
||||
getPoliciesList,
|
||||
testPoliciesAndGuardrails,
|
||||
} from "@/components/networking";
|
||||
import {
|
||||
AlertTriangle,
|
||||
BarChart3,
|
||||
|
|
@ -292,43 +296,72 @@ export default function ComplianceUI({
|
|||
});
|
||||
};
|
||||
|
||||
const runQuickTest = useCallback(() => {
|
||||
if (!quickTestInput.trim()) return;
|
||||
const runQuickTest = useCallback(async () => {
|
||||
if (!quickTestInput.trim() || !accessToken) return;
|
||||
const text = quickTestInput.trim();
|
||||
const userMsg: QuickTestMessage = {
|
||||
id: `msg-${Date.now()}`,
|
||||
type: "user",
|
||||
text: quickTestInput.trim(),
|
||||
text,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
setQuickTestMessages((prev) => [...prev, userMsg]);
|
||||
setQuickTestInput("");
|
||||
setIsQuickTesting(true);
|
||||
setTimeout(() => {
|
||||
const rand = Math.random();
|
||||
const result: "blocked" | "allowed" = rand < 0.4 ? "blocked" : "allowed";
|
||||
try {
|
||||
const { inputs, guardrail_errors } = await testPoliciesAndGuardrails(
|
||||
accessToken,
|
||||
{
|
||||
policy_names:
|
||||
selectedPolicies.length > 0 ? selectedPolicies : undefined,
|
||||
guardrail_names:
|
||||
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
|
||||
inputs: { texts: [text] },
|
||||
request_data: {},
|
||||
input_type: "request",
|
||||
}
|
||||
);
|
||||
const result: "blocked" | "allowed" =
|
||||
guardrail_errors.length > 0 ? "blocked" : "allowed";
|
||||
const triggeredBy =
|
||||
result === "blocked"
|
||||
? selectedGuardrails.length > 0
|
||||
? guardrailOptions.find((g) => selectedGuardrails.includes(g.id))?.name
|
||||
: selectedPolicies.length > 0
|
||||
? policyOptions.find((p) => selectedPolicies.includes(p.id))?.name
|
||||
: "content-filter"
|
||||
guardrail_errors.length > 0
|
||||
? guardrail_errors
|
||||
.map((e) => `${e.guardrail_name}: ${e.message}`)
|
||||
.join("; ")
|
||||
: undefined;
|
||||
const displayText =
|
||||
result === "blocked"
|
||||
? `Blocked — ${triggeredBy ?? "content filter"}`
|
||||
: "Allowed — no policy or guardrail violations detected.";
|
||||
const sysMsg: QuickTestMessage = {
|
||||
id: `msg-${Date.now()}-sys`,
|
||||
type: "system",
|
||||
text:
|
||||
result === "blocked"
|
||||
? `Blocked — triggered by ${triggeredBy ?? "content filter"}`
|
||||
: "Allowed — no policy or guardrail violations detected.",
|
||||
text: displayText,
|
||||
result,
|
||||
triggeredBy,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
setQuickTestMessages((prev) => [...prev, sysMsg]);
|
||||
} catch (err) {
|
||||
const errorMessage = err instanceof Error ? err.message : String(err);
|
||||
const sysMsg: QuickTestMessage = {
|
||||
id: `msg-${Date.now()}-sys`,
|
||||
type: "system",
|
||||
text: `Error: ${errorMessage}`,
|
||||
result: "blocked",
|
||||
triggeredBy: errorMessage,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
setQuickTestMessages((prev) => [...prev, sysMsg]);
|
||||
} finally {
|
||||
setIsQuickTesting(false);
|
||||
}, 600 + Math.random() * 400);
|
||||
}, [quickTestInput, selectedPolicies, selectedGuardrails, policyOptions, guardrailOptions]);
|
||||
}
|
||||
}, [
|
||||
accessToken,
|
||||
quickTestInput,
|
||||
selectedPolicies,
|
||||
selectedGuardrails,
|
||||
]);
|
||||
|
||||
const handleQuickTestKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
|
|
@ -337,8 +370,8 @@ export default function ComplianceUI({
|
|||
}
|
||||
};
|
||||
|
||||
const runTests = useCallback(() => {
|
||||
if (selectedPromptIds.size === 0) return;
|
||||
const runTests = useCallback(async () => {
|
||||
if (selectedPromptIds.size === 0 || !accessToken) return;
|
||||
setIsRunning(true);
|
||||
setResultFilter("all");
|
||||
setRightTab("batch-results");
|
||||
|
|
@ -357,32 +390,38 @@ export default function ComplianceUI({
|
|||
status: "pending",
|
||||
}));
|
||||
setTestResults(pendingResults);
|
||||
pendingResults.forEach((result, index) => {
|
||||
setTimeout(() => {
|
||||
for (let index = 0; index < selected.length; index++) {
|
||||
const promptResult = selected[index];
|
||||
try {
|
||||
const { guardrail_errors } = await testPoliciesAndGuardrails(
|
||||
accessToken,
|
||||
{
|
||||
policy_names:
|
||||
selectedPolicies.length > 0 ? selectedPolicies : undefined,
|
||||
guardrail_names:
|
||||
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
|
||||
inputs: { texts: [promptResult.prompt] },
|
||||
request_data: {},
|
||||
input_type: "request",
|
||||
}
|
||||
);
|
||||
const actualResult: "blocked" | "allowed" =
|
||||
guardrail_errors.length > 0 ? "blocked" : "allowed";
|
||||
const triggeredBy =
|
||||
guardrail_errors.length > 0
|
||||
? guardrail_errors
|
||||
.map((e) => `${e.guardrail_name}: ${e.message}`)
|
||||
.join("; ")
|
||||
: undefined;
|
||||
const isMatch =
|
||||
(promptResult.expectedResult === "fail" &&
|
||||
actualResult === "blocked") ||
|
||||
(promptResult.expectedResult === "pass" &&
|
||||
actualResult === "allowed");
|
||||
setTestResults((prev) => {
|
||||
const updated = [...prev];
|
||||
const rand = Math.random();
|
||||
const actualResult: "blocked" | "allowed" =
|
||||
result.expectedResult === "fail"
|
||||
? rand < 0.85
|
||||
? "blocked"
|
||||
: "allowed"
|
||||
: rand < 0.9
|
||||
? "allowed"
|
||||
: "blocked";
|
||||
const isMatch =
|
||||
(result.expectedResult === "fail" && actualResult === "blocked") ||
|
||||
(result.expectedResult === "pass" && actualResult === "allowed");
|
||||
const triggeredBy =
|
||||
actualResult === "blocked"
|
||||
? selectedGuardrails.length > 0
|
||||
? guardrailOptions.find((g) => selectedGuardrails.includes(g.id))?.name
|
||||
: selectedPolicies.length > 0
|
||||
? policyOptions.find((p) => selectedPolicies.includes(p.id))?.name
|
||||
: "content-filter"
|
||||
: undefined;
|
||||
updated[index] = {
|
||||
...result,
|
||||
...pendingResults[index],
|
||||
actualResult,
|
||||
isMatch,
|
||||
triggeredBy,
|
||||
|
|
@ -390,16 +429,28 @@ export default function ComplianceUI({
|
|||
};
|
||||
return updated;
|
||||
});
|
||||
if (index === pendingResults.length - 1) setIsRunning(false);
|
||||
}, 300 + index * 120);
|
||||
});
|
||||
} catch (err) {
|
||||
const errorMessage = err instanceof Error ? err.message : String(err);
|
||||
setTestResults((prev) => {
|
||||
const updated = [...prev];
|
||||
updated[index] = {
|
||||
...pendingResults[index],
|
||||
actualResult: "blocked",
|
||||
isMatch: false,
|
||||
triggeredBy: `Error: ${errorMessage}`,
|
||||
status: "complete",
|
||||
};
|
||||
return updated;
|
||||
});
|
||||
}
|
||||
}
|
||||
setIsRunning(false);
|
||||
}, [
|
||||
accessToken,
|
||||
selectedPromptIds,
|
||||
selectedPolicies,
|
||||
selectedGuardrails,
|
||||
allFrameworks,
|
||||
policyOptions,
|
||||
guardrailOptions,
|
||||
]);
|
||||
|
||||
const completedResults = testResults.filter((r) => r.status === "complete");
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue