mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: fix batch test endpoint for compliance playground
This commit is contained in:
parent
4349bdaa27
commit
11e030c7a1
6 changed files with 390 additions and 121 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue