From 11e030c7a1e5c14f3c08ec8e5c0a9df0ddfa1265 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 20 Feb 2026 17:11:57 -0800 Subject: [PATCH] fix: fix batch test endpoint for compliance playground --- .../competitor_intent/airline.py | 23 +++ .../competitor_intent/base.py | 71 ++++++--- .../policy_endpoints/endpoints.py | 125 +++++++++------ .../content_filter/test_competitor_intent.py | 129 ++++++++++++++-- .../src/components/networking.tsx | 20 ++- .../playground/complianceUI/ComplianceUI.tsx | 143 +++++++++++++----- 6 files changed, 390 insertions(+), 121 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py index b8df4e36089..732da562de2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py @@ -78,6 +78,24 @@ AIRLINE_COMPETITOR_SIGNALS = [ r"\bprivilege\s+club\b", ] +# Operational-only: baggage, lounge, check-in, refund (no comparison language). +# When only these appear with ambiguous token → treat as product query (OTHER_MEANING). +AIRLINE_OPERATIONAL_SIGNALS = [ + r"\bbaggage\s+allowance\b", + r"\blounge\b", + r"\bcheck[- ]?in\b", + r"\brefund\b", + r"\bpremium\s+lounge\b", +] +# Comparison language: if present with competitor signals → COMPETITOR. +AIRLINE_COMPARISON_SIGNALS = [ + r"\bbetter\b", + r"\bbest\b", + r"\bvs\.?\b", + r"\bversus\b", + r"\bcompare\b", +] + # Explicit markers: strong override when present. AIRLINE_EXPLICIT_COMPETITOR_MARKER = r"\b(airways?|airline|carrier)\b" AIRLINE_EXPLICIT_OTHER_MEANING_MARKER = ( @@ -127,6 +145,11 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): text_lower ): return "OTHER_MEANING", 0.85 + # Operational-only: baggage/lounge/check-in/refund with no comparison → product query + has_comparison = _count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0 + operational_count = _count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS) + if not has_comparison and operational_count > 0: + return "OTHER_MEANING", 0.85 # Score: location/travel context vs airline context (no place-name list) other_count = _count_signals(text_lower, self._other_meaning_signals) if self._other_meaning_anchors: diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py index 4f4a0ececca..7a86e5542b2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py @@ -7,15 +7,15 @@ import unicodedata from typing import Any, Dict, List, Optional, Pattern, Set, Tuple, cast from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( - CompetitorActionHint, - CompetitorIntentEvidenceEntry, - CompetitorIntentResult, -) + CompetitorActionHint, CompetitorIntentEvidenceEntry, + CompetitorIntentResult) ZERO_WIDTH = re.compile(r"[\u200b-\u200d\u2060\ufeff]") LEET = {"@": "a", "4": "a", "0": "o", "3": "e", "1": "i", "5": "s", "7": "t"} -OTHER_MEANING_DEFAULT_THRESHOLD = 0.65 # Below this → treat as non-competitor (safe default). +OTHER_MEANING_DEFAULT_THRESHOLD = ( + 0.65 # Below this → treat as non-competitor (safe default). +) def normalize(text: str) -> str: @@ -75,7 +75,7 @@ class BaseCompetitorIntentChecker: for c in competitors: self._competitor_tokens.add(c) self.competitor_canonical[c] = c - for a in (aliases_map.get(c) or []): + for a in aliases_map.get(c) or []: a = a.lower().strip() if a: self._competitor_tokens.add(a) @@ -99,8 +99,21 @@ class BaseCompetitorIntentChecker: ) self._comparison_words: List[str] = list( config.get("comparison_words") - or ["better", "worse", "best", "vs", "versus", "compare", "alternative", "recommend"] + or [ + "better", + "worse", + "best", + "vs", + "versus", + "compare", + "alternative", + "recommend", + "ranked", + ] ) + self._domain_words: List[str] = [ + s.lower().strip() for s in (config.get("domain_words") or []) if s + ] def _classify_ambiguous(self, text: str, token: str) -> Tuple[str, float]: """ @@ -144,6 +157,34 @@ class BaseCompetitorIntentChecker: matches = self._find_matches(text) if not matches: + has_comparison = any( + re.search(r"\b" + re.escape(w) + r"\b", normalized) + for w in self._comparison_words + ) + has_domain = self._domain_words and any( + re.search(r"\b" + re.escape(w) + r"\b", normalized) + for w in self._domain_words + ) + if has_comparison and has_domain: + evidence.append( + { + "type": "signal", + "key": "category_ranking", + "match": "comparison + domain", + } + ) + action_hint = cast( + CompetitorActionHint, + self.policy.get("category_ranking", "reframe"), + ) + return { + "intent": "category_ranking", + "confidence": 0.65, + "entities": entities, + "signals": ["category_ranking"], + "action_hint": action_hint, + "evidence": evidence, + } return { "intent": "other", "confidence": 0.0, @@ -154,18 +195,7 @@ class BaseCompetitorIntentChecker: } competitor_resolved: List[str] = [] - for token, canonical, is_ambig in matches: - if not is_ambig: - competitor_resolved.append(canonical) - evidence.append( - { - "type": "entity", - "key": "competitor", - "value": canonical, - "match": token, - } - ) - continue + for token, canonical, _ in matches: label, conf = self._classify_ambiguous(normalized, token) if label == "OTHER_MEANING": evidence.append( @@ -235,7 +265,8 @@ class BaseCompetitorIntentChecker: "intent": intent, "confidence": round(confidence, 2), "entities": entities, - "signals": ["competitor_resolved"] + (["comparison"] if has_comparison else []), + "signals": ["competitor_resolved"] + + (["comparison"] if has_comparison else []), "action_hint": action_hint, "evidence": evidence, } diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index 36745a607cc..069603afdfd 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -12,14 +12,7 @@ All /policy management endpoints import copy import json import os -from typing import ( - TYPE_CHECKING, - AsyncIterator, - List, - Literal, - Optional, - cast, -) +from typing import TYPE_CHECKING, AsyncIterator, List, Literal, Optional, cast from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse @@ -27,11 +20,9 @@ from pydantic import BaseModel, Field from typing_extensions import TypedDict from litellm._logging import verbose_proxy_logger -from litellm.constants import ( - COMPETITOR_LLM_TEMPERATURE, - DEFAULT_COMPETITOR_DISCOVERY_MODEL, - MAX_COMPETITOR_NAMES, -) +from litellm.constants import (COMPETITOR_LLM_TEMPERATURE, + DEFAULT_COMPETITOR_DISCOVERY_MODEL, + MAX_COMPETITOR_NAMES) from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -39,21 +30,20 @@ from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.proxy.policy_engine.policy_resolver import PolicyResolver -from litellm.types.proxy.policy_engine import ( - PolicyGuardrailsResponse, - PolicyInfoResponse, - PolicyListResponse, - PolicyMatchContext, - PolicyScopeResponse, - PolicySummaryItem, - PolicyTestResponse, - PolicyValidateRequest, - PolicyValidationResponse, -) +from litellm.types.proxy.policy_engine import (PolicyGuardrailsResponse, + PolicyInfoResponse, + PolicyListResponse, + PolicyMatchContext, + PolicyScopeResponse, + PolicySummaryItem, + PolicyTestResponse, + PolicyValidateRequest, + PolicyValidationResponse) from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import \ + Logging as LiteLLMLoggingObj router = APIRouter() @@ -86,6 +76,19 @@ class ApplyPoliciesResult(TypedDict): guardrail_errors: List[GuardrailErrorEntry] +class ApplyPoliciesPerItemResult(TypedDict): + """Result for one input when using inputs_list.""" + + inputs: GenericGuardrailAPIInputs + guardrail_errors: List[GuardrailErrorEntry] + + +class ApplyPoliciesListResult(TypedDict): + """Result when using inputs_list: one result per input.""" + + results: List[ApplyPoliciesPerItemResult] + + async def apply_policies( policy_names: Optional[list[str]], inputs: GenericGuardrailAPIInputs, @@ -187,7 +190,14 @@ class TestPoliciesAndGuardrailsRequest(BaseModel): 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\": [\"...\"] }") + inputs: Optional[dict] = Field( + default=None, + description="GenericGuardrailAPIInputs, e.g. { \"texts\": [\"...\"] }. Use inputs_list for per-input processing.", + ) + inputs_list: Optional[List[dict]] = Field( + default=None, + description="List of GenericGuardrailAPIInputs; each item processed separately (for batch compliance testing).", + ) 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") @@ -206,26 +216,48 @@ async def test_policies_and_guardrails( """ 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. + Use inputs_list for batch testing: each input is processed as a separate call so + per-input block/allow and errors are returned. + + Use inputs for a single call (legacy). """ - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + 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 + if data.inputs_list is not None: + results: List[ApplyPoliciesPerItemResult] = [] + for inp in data.inputs_list: + inputs_typed = cast(GenericGuardrailAPIInputs, inp) + item_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, + ) + results.append( + ApplyPoliciesPerItemResult( + inputs=item_result["inputs"], + guardrail_errors=item_result["guardrail_errors"], + ) + ) + return ApplyPoliciesListResult(results=results) + if data.inputs is not None: + inputs_typed = cast(GenericGuardrailAPIInputs, data.inputs) + return 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, + ) + raise ValueError("Either inputs or inputs_list must be provided") except Exception as e: raise handle_exception_on_proxy(e) @@ -506,7 +538,8 @@ async def get_policy_templates( return _load_policy_templates_from_local_backup() try: - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.llms.custom_httpx.http_handler import \ + get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider async_client = get_async_httpx_client( @@ -981,9 +1014,8 @@ async def suggest_policy_templates( Calls an LLM with tool calling to match user requirements to available templates. """ - from litellm.proxy.management_endpoints.policy_endpoints.ai_policy_suggester import ( - AiPolicySuggester, - ) + from litellm.proxy.management_endpoints.policy_endpoints.ai_policy_suggester import \ + AiPolicySuggester templates = _load_policy_templates_from_local_backup() suggester = AiPolicySuggester() @@ -1053,9 +1085,8 @@ async def _test_guardrail_definitions( text: str, ) -> List[GuardrailTestResultEntry]: """Instantiate and run each guardrail definition against the text.""" - from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( - ContentFilterGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import \ + ContentFilterGuardrail results: List[GuardrailTestResultEntry] = [] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_competitor_intent.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_competitor_intent.py index dc5c225f3ac..e2deb3612c0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_competitor_intent.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_competitor_intent.py @@ -5,10 +5,7 @@ Tests for competitor intent detection (normalize, entity layer, scoring, policy) import pytest from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent import ( - AirlineCompetitorIntentChecker, - normalize, - text_for_entity_matching, -) + AirlineCompetitorIntentChecker, normalize, text_for_entity_matching) class TestNormalize: @@ -120,9 +117,8 @@ class TestContentFilterWithCompetitorIntent: @pytest.mark.asyncio async def test_apply_guardrail_with_competitor_intent_allow(self): - from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( - ContentFilterGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import \ + ContentFilterGuardrail guardrail = ContentFilterGuardrail( guardrail_name="test-competitor", @@ -143,9 +139,8 @@ class TestContentFilterWithCompetitorIntent: async def test_apply_guardrail_with_competitor_intent_refuse(self): from fastapi import HTTPException - from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( - ContentFilterGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import \ + ContentFilterGuardrail guardrail = ContentFilterGuardrail( guardrail_name="test-competitor", @@ -163,3 +158,117 @@ class TestContentFilterWithCompetitorIntent: inputs, request_data={}, input_type="request" ) assert exc_info.value.status_code == 403 + + +# Exact config from litellm/proxy/_new_secret_config.yaml (lines 27-53). +AIRLINE_PROXY_CONFIG = { + "brand_self": ["emirates", "ek"], + "competitors": ["qatar airways", "qatar", "etihad"], + "locations": ["qatar", "doha", "doh"], + "competitor_aliases": { + "qatar airways": ["qr", "doha airline"], + "qatar": ["qr"], + }, + "policy": { + "competitor_comparison": "refuse", + "possible_competitor_comparison": "reframe", + }, + "threshold_high": 0.70, + "threshold_medium": 0.45, + "threshold_low": 0.30, +} + +# Config matching proxy + domain_words, category_ranking, and competitors for full dataset. +AIRLINE_COMPLIANCE_CONFIG = { + "brand_self": ["emirates", "ek"], + "competitors": [ + "qatar airways", + "qatar", + "etihad", + "turkish airlines", + "lufthansa", + "singapore airlines", + ], + "locations": ["qatar", "doha", "doh"], + "competitor_aliases": { + "qatar airways": ["qr", "doha airline"], + "qatar": ["qr"], + }, + "domain_words": ["airline", "airlines", "carrier"], + "policy": { + "competitor_comparison": "refuse", + "possible_competitor_comparison": "reframe", + "category_ranking": "reframe", + }, + "threshold_high": 0.70, + "threshold_medium": 0.45, + "threshold_low": 0.30, +} + +# Dataset from ui/litellm-dashboard/src/data/compliancePrompts.ts airlineCompliancePrompts +AIRLINE_COMPLIANCE_DATASET = [ + # Destination Intent (should PASS) + ("airline_001", "Do you have flights to Qatar?", "pass"), + ("airline_002", "How long is the flight to Doha?", "pass"), + ("airline_003", "Do I need a visa to transit through Qatar?", "pass"), + ("airline_004", "Are there lounges at Hamad Airport?", "pass"), + ("airline_005", "What documents are required to enter Singapore?", "pass"), + ("airline_006", "What airlines fly to Turkey?", "pass"), + ("airline_007", "What is the airport code for Doha?", "pass"), + ("airline_008", "Is Doha a good place for a long layover?", "pass"), + ("airline_009", "What is the weather like in Qatar in March?", "pass"), + ("airline_010", "How much connection time do I need in Doha?", "pass"), + # Competitor Comparison (should FAIL) + ("airline_011", "Is Qatar Airways better than Emirates?", "fail"), + ("airline_012", "Which airline has better business class, Qatar or Singapore Airlines?", "fail"), + ("airline_013", "Which airline is the best for long haul flights?", "fail"), + ("airline_014", "Should I choose Qatar Airways or another airline?", "fail"), + ("airline_015", "Compare Turkish Airlines vs Lufthansa.", "fail"), + ("airline_016", "Which airline has better lounges?", "fail"), + ("airline_017", "Which airline has the best customer satisfaction?", "fail"), + ("airline_018", "Is the Doha airline better than other carriers?", "fail"), + ("airline_019", "Should I switch to another airline for a better experience?", "fail"), + ("airline_020", "Which airline is ranked number one worldwide?", "fail"), + # Ambiguous Entity (should PASS) + ("airline_021", "Qatar baggage allowance", "pass"), + ("airline_022", "Qatar lounge access rules", "pass"), + ("airline_023", "Qatar check in time", "pass"), + ("airline_024", "Doha premium lounge access", "pass"), + ("airline_025", "Qatar refund policy", "pass"), +] + + +class TestAirlineComplianceDataset: + """Run full airline compliance dataset with proxy config; all cases must match expected outcome.""" + + def test_airline_001_passes_with_exact_proxy_config(self): + """With exact proxy config, first compliance case (flights to Qatar) must pass (allow).""" + checker = AirlineCompetitorIntentChecker(AIRLINE_PROXY_CONFIG) + result = checker.run("Do you have flights to Qatar?") + assert result["intent"] == "other" + assert result["action_hint"] == "allow" + + def test_airline_compliance_dataset_with_proxy_config(self): + """Every prompt must get intent/action consistent with expectedResult (pass=allow, fail=refuse/reframe).""" + checker = AirlineCompetitorIntentChecker(AIRLINE_COMPLIANCE_CONFIG) + failures = [] + for prompt_id, prompt_text, expected in AIRLINE_COMPLIANCE_DATASET: + result = checker.run(prompt_text) + intent = result.get("intent", "other") + action_hint = result.get("action_hint", "allow") + if expected == "pass": + allowed = intent == "other" and action_hint == "allow" + if not allowed: + failures.append( + f"{prompt_id}: expected pass, got intent={intent!r} action_hint={action_hint!r} for {prompt_text!r}" + ) + else: + blocked = ( + intent != "other" + and action_hint in ("refuse", "reframe") + ) + if not blocked: + failures.append( + f"{prompt_id}: expected fail, got intent={intent!r} action_hint={action_hint!r} for {prompt_text!r}" + ) + assert not failures, f"Airline compliance dataset failures:\n" + "\n".join(failures) diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 8b624f7caf6..b7b66122677 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5438,10 +5438,19 @@ export const getPoliciesList = async (accessToken: string) => { } }; +export interface GuardrailInputs { + texts?: string[]; + images?: string[]; + [key: string]: unknown; +} + export interface TestPoliciesAndGuardrailsRequest { policy_names?: string[] | null; guardrail_names?: string[] | null; - inputs: { texts?: string[]; images?: string[]; [key: string]: unknown }; + /** Single input (legacy). Use inputs_list for per-input batch processing. */ + inputs?: GuardrailInputs | null; + /** List of inputs; each processed separately for batch compliance testing. */ + inputs_list?: GuardrailInputs[] | null; request_data?: Record; input_type?: "request" | "response"; } @@ -5452,8 +5461,10 @@ export interface GuardrailErrorEntry { } export interface TestPoliciesAndGuardrailsResponse { - inputs: Record; - guardrail_errors: GuardrailErrorEntry[]; + inputs?: Record; + guardrail_errors?: GuardrailErrorEntry[]; + /** Present when inputs_list was used; one result per input. */ + results?: Array<{ inputs: Record; guardrail_errors: GuardrailErrorEntry[] }>; } export const testPoliciesAndGuardrails = async ( @@ -5473,7 +5484,8 @@ export const testPoliciesAndGuardrails = async ( body: JSON.stringify({ policy_names: body.policy_names ?? null, guardrail_names: body.guardrail_names ?? null, - inputs: body.inputs, + inputs: body.inputs ?? null, + inputs_list: body.inputs_list ?? null, request_data: body.request_data ?? {}, input_type: body.input_type ?? "request", }), diff --git a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx index ccd3dcfe8d5..e69b34b1861 100644 --- a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx @@ -445,7 +445,7 @@ export default function ComplianceUI({ setQuickTestInput(""); setIsQuickTesting(true); try { - const { inputs, guardrail_errors } = await testPoliciesAndGuardrails( + const { inputs, guardrail_errors = [] } = await testPoliciesAndGuardrails( accessToken, { policy_names: @@ -533,39 +533,44 @@ export default function ComplianceUI({ })); setTestResults(pendingResults); try { - const { inputs, guardrail_errors } = await testPoliciesAndGuardrails( - accessToken, - { - policy_names: - selectedPolicies.length > 0 ? selectedPolicies : undefined, - guardrail_names: - selectedGuardrails.length > 0 ? selectedGuardrails : undefined, - inputs: { texts: allTexts }, - 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 returnedTexts: (string | undefined)[] = - Array.isArray(inputs?.texts) ? inputs.texts : []; + const inputsList = allTexts.map((text) => ({ texts: [text] })); + const response = await testPoliciesAndGuardrails(accessToken, { + policy_names: + selectedPolicies.length > 0 ? selectedPolicies : undefined, + guardrail_names: + selectedGuardrails.length > 0 ? selectedGuardrails : undefined, + inputs_list: inputsList, + request_data: {}, + input_type: "request", + }); + const results = response.results ?? []; setTestResults( - pendingResults.map((row, index) => ({ - ...row, - actualResult, - isMatch: - (row.expectedResult === "fail" && actualResult === "blocked") || - (row.expectedResult === "pass" && actualResult === "allowed"), - triggeredBy, - returnedText: returnedTexts[index], - status: "complete" as const, - })) + pendingResults.map((row, index) => { + const item = results[index]; + const guardrail_errors = item?.guardrail_errors ?? []; + 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 returnedText = + Array.isArray(item?.inputs?.texts) && item.inputs.texts.length > 0 + ? item.inputs.texts[0] + : undefined; + return { + ...row, + actualResult, + isMatch: + (row.expectedResult === "fail" && actualResult === "blocked") || + (row.expectedResult === "pass" && actualResult === "allowed"), + triggeredBy, + returnedText, + status: "complete" as const, + }; + }) ); } catch (err) { const errorMessage = err instanceof Error ? err.message : String(err); @@ -592,6 +597,12 @@ export default function ComplianceUI({ const completedResults = testResults.filter((r) => r.status === "complete"); const matchCount = completedResults.filter((r) => r.isMatch).length; const mismatchCount = completedResults.filter((r) => !r.isMatch).length; + const falsePositiveCount = completedResults.filter( + (r) => r.expectedResult === "pass" && r.actualResult === "blocked" + ).length; + const falseNegativeCount = completedResults.filter( + (r) => r.expectedResult === "fail" && r.actualResult === "allowed" + ).length; const pendingCount = testResults.filter((r) => r.status !== "complete").length; const filteredResults = testResults.filter((r) => { if (resultFilter === "matches") return r.status === "complete" && r.isMatch; @@ -600,6 +611,31 @@ export default function ComplianceUI({ return true; }); + const exportBatchResults = () => { + if (filteredResults.length === 0) return; + const rows = filteredResults.map((r) => ({ + prompt_id: r.promptId, + prompt: r.prompt, + category: r.category, + expected_result: r.expectedResult, + actual_result: r.actualResult, + is_match: r.isMatch ? "yes" : "no", + status: r.status, + triggered_by: r.triggeredBy ?? "", + returned_text: r.returnedText ?? "", + })); + const csv = Papa.unparse(rows); + const blob = new Blob([csv], { type: "text/csv" }); + const url = window.URL.createObjectURL(blob); + const a = document.createElement("a"); + a.href = url; + a.download = `compliance_batch_results_${new Date().toISOString().slice(0, 10)}.csv`; + document.body.appendChild(a); + a.click(); + document.body.removeChild(a); + window.URL.revokeObjectURL(url); + }; + const filteredFrameworks = allFrameworks .map((fw) => ({ ...fw, @@ -1361,14 +1397,33 @@ export default function ComplianceUI({

Results

{testResults.length > 0 && ( -
+
+ +
{matchCount} - + + + {falseNegativeCount} FN + + - {mismatchCount} + {falsePositiveCount} FP {pendingCount > 0 && ( @@ -1377,6 +1432,7 @@ export default function ComplianceUI({ )}
+
)}
{testResults.length > 0 && ( @@ -1439,11 +1495,18 @@ export default function ComplianceUI({ correct
- - - {mismatchCount} + + + {falseNegativeCount} {" "} - gaps + false negative + +
+ + + {falsePositiveCount} + {" "} + false positive