fix: fix batch test endpoint for compliance playground

This commit is contained in:
Krrish Dholakia 2026-02-20 17:11:57 -08:00
parent 4349bdaa27
commit 11e030c7a1
6 changed files with 390 additions and 121 deletions

View file

@ -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:

View file

@ -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,
}

View file

@ -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] = []

View file

@ -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)

View file

@ -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<string, unknown>;
input_type?: "request" | "response";
}
@ -5452,8 +5461,10 @@ export interface GuardrailErrorEntry {
}
export interface TestPoliciesAndGuardrailsResponse {
inputs: Record<string, unknown>;
guardrail_errors: GuardrailErrorEntry[];
inputs?: Record<string, unknown>;
guardrail_errors?: GuardrailErrorEntry[];
/** Present when inputs_list was used; one result per input. */
results?: Array<{ inputs: Record<string, unknown>; 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",
}),

View file

@ -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({
<div className="flex items-center justify-between mb-2">
<h2 className="text-sm font-semibold text-gray-900">Results</h2>
{testResults.length > 0 && (
<div className="flex items-center gap-2.5 text-[11px]">
<div className="flex items-center gap-2">
<button
type="button"
onClick={exportBatchResults}
disabled={filteredResults.length === 0}
className="flex items-center gap-1 text-[11px] font-medium text-gray-600 hover:text-gray-900 hover:bg-gray-100 px-2 py-1 rounded transition-colors disabled:opacity-50 disabled:cursor-not-allowed disabled:hover:bg-transparent"
>
<Download className="w-3 h-3" /> Export CSV
</button>
<div className="flex items-center gap-2.5 text-[11px]">
<span className="flex items-center gap-1 text-green-600">
<CheckCircle2 className="w-3 h-3" />
{matchCount}
</span>
<span className="flex items-center gap-1 text-red-600">
<span
className="flex items-center gap-1 text-amber-600"
title="Allowed content that should have been blocked"
>
<AlertTriangle className="w-3 h-3" />
{falseNegativeCount} FN
</span>
<span
className="flex items-center gap-1 text-red-600"
title="Blocked content that should have been allowed"
>
<X className="w-3 h-3" />
{mismatchCount}
{falsePositiveCount} FP
</span>
{pendingCount > 0 && (
<span className="flex items-center gap-1 text-gray-500">
@ -1377,6 +1432,7 @@ export default function ComplianceUI({
</span>
)}
</div>
</div>
)}
</div>
{testResults.length > 0 && (
@ -1439,11 +1495,18 @@ export default function ComplianceUI({
<span className="text-gray-500">correct</span>
</span>
<div className="w-px h-4 bg-gray-200" />
<span>
<span className="font-semibold text-red-700">
{mismatchCount}
<span title="Allowed content that should have been blocked">
<span className="font-semibold text-amber-700">
{falseNegativeCount}
</span>{" "}
<span className="text-gray-500">gaps</span>
<span className="text-gray-500">false negative</span>
</span>
<div className="w-px h-4 bg-gray-200" />
<span title="Blocked content that should have been allowed">
<span className="font-semibold text-red-700">
{falsePositiveCount}
</span>{" "}
<span className="text-gray-500">false positive</span>
</span>
</div>
<div