diff --git a/litellm/proxy/management_endpoints/policy_endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints.py index eeda25f6ba4..92af60313c0 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints.py @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/test_policy_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_policy_endpoints.py index 0abe6eccdc5..2fb2a6dcc8e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_policy_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_policy_endpoints.py @@ -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"] == [] diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 905514a69ba..bf2dd10dc80 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -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; + input_type?: "request" | "response"; +} + +export interface GuardrailErrorEntry { + guardrail_name: string; + message: string; +} + +export interface TestPoliciesAndGuardrailsResponse { + inputs: Record; + guardrail_errors: GuardrailErrorEntry[]; +} + +export const testPoliciesAndGuardrails = async ( + accessToken: string, + body: TestPoliciesAndGuardrailsRequest +): Promise => { + 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}`; diff --git a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx index 3e2eb2475ee..171908a7512 100644 --- a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx @@ -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) => { 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");