"""PII masking evaluation helpers for local regression fixtures.""" from __future__ import annotations import json from pathlib import Path from typing import Any, Callable, Iterable, Mapping from . import guardrail MaskFunc = Callable[[str], guardrail.MaskResult] REPORT_SCHEMA_VERSION = "vignette.pii_masking_eval_report.v1" INPUT_SCHEMA_VERSION = "vignette.pii_masking_eval_input.v1" DEFAULT_CASE_META = { "locale": "ko-KR", "source": "synthetic", "category": "unspecified", "severity": "medium", } def load_cases(path: Path) -> list[dict[str, Any]]: data = json.loads(path.read_text(encoding="utf-8")) if not isinstance(data, list): raise ValueError("PII masking fixture must be a list") return [dict(item) for item in data] def evaluate_case( case: Mapping[str, Any], *, mask_func: MaskFunc = guardrail.mask_pii, include_evidence_text: bool = False, ) -> dict[str, Any]: case_id = str(case.get("id") or "") text = str(case.get("text") or "") locale = str(case.get("locale") or DEFAULT_CASE_META["locale"]) source = str(case.get("source") or DEFAULT_CASE_META["source"]) category = str(case.get("category") or DEFAULT_CASE_META["category"]) severity = str(case.get("severity") or DEFAULT_CASE_META["severity"]) result = mask_func(text) entities = set(result.entities) expected_entities = {str(item) for item in case.get("expected_entities") or []} unexpected_entities = {str(item) for item in case.get("unexpected_entities") or []} forbidden_substrings = [str(item) for item in case.get("forbidden_substrings") or []] required_substrings = [str(item) for item in case.get("required_substrings") or []] missing_entities = sorted(expected_entities - entities) unexpected_detected = sorted(unexpected_entities & entities) forbidden_remaining = [item for item in forbidden_substrings if item and item in result.text_masked] required_missing = [item for item in required_substrings if item and item not in result.text_masked] passed = not (missing_entities or unexpected_detected or forbidden_remaining or required_missing) report = { "id": case_id, "locale": locale, "source": source, "category": category, "severity": severity, "passed": passed, "entities": sorted(entities), "missing_entities": missing_entities, "unexpected_entities": unexpected_detected, "forbidden_remaining_count": len(forbidden_remaining), "required_missing": required_missing, } if include_evidence_text: report["masked_text"] = result.text_masked report["forbidden_remaining"] = forbidden_remaining return report def evaluate_cases( cases: Iterable[Mapping[str, Any]], *, mask_func: MaskFunc = guardrail.mask_pii, include_evidence_text: bool = False, ) -> dict[str, Any]: case_list = list(cases) results = [ evaluate_case(case, mask_func=mask_func, include_evidence_text=include_evidence_text) for case in case_list ] total_expected_entities = 0 matched_expected_entities = 0 total_forbidden = 0 removed_forbidden = 0 unexpected_violations = 0 for case, result in zip(case_list, results): expected_entities = {str(item) for item in case.get("expected_entities") or []} forbidden = [str(item) for item in case.get("forbidden_substrings") or []] total_expected_entities += len(expected_entities) matched_expected_entities += len(expected_entities) - len(result["missing_entities"]) total_forbidden += len(forbidden) remaining_forbidden = int(result.get("forbidden_remaining_count", len(result.get("forbidden_remaining", [])))) removed_forbidden += len(forbidden) - remaining_forbidden unexpected_violations += len(result["unexpected_entities"]) passed_cases = sum(1 for result in results if result["passed"]) return { "schema_version": REPORT_SCHEMA_VERSION, "input_schema_version": INPUT_SCHEMA_VERSION, "run_mode": "technical_dry_run", "data_source": "local_fixture", "evidence_text_included": include_evidence_text, "passed": passed_cases == len(results), "cases_total": len(results), "cases_passed": passed_cases, "cases_failed": len(results) - passed_cases, "expected_entity_recall": _ratio(matched_expected_entities, total_expected_entities), "forbidden_substring_removal": _ratio(removed_forbidden, total_forbidden), "unexpected_entity_violations": unexpected_violations, "by_source": _breakdown(case_list, results, "source"), "by_category": _breakdown(case_list, results, "category"), "by_severity": _breakdown(case_list, results, "severity"), "results": results, } def evaluate_fixture( path: Path, *, mask_func: MaskFunc = guardrail.mask_pii, include_evidence_text: bool = False, ) -> dict[str, Any]: return evaluate_cases(load_cases(path), mask_func=mask_func, include_evidence_text=include_evidence_text) def _ratio(numerator: int, denominator: int) -> float: if denominator <= 0: return 1.0 return round(numerator / denominator, 4) def _case_meta(case: Mapping[str, Any], key: str) -> str: fallback = DEFAULT_CASE_META.get(key, "unspecified") return str(case.get(key) or fallback) def _breakdown( cases: list[Mapping[str, Any]], results: list[Mapping[str, Any]], key: str, ) -> dict[str, dict[str, int]]: grouped: dict[str, dict[str, int]] = {} for case, result in zip(cases, results): value = _case_meta(case, key) bucket = grouped.setdefault(value, {"cases_total": 0, "cases_passed": 0, "cases_failed": 0}) bucket["cases_total"] += 1 if result.get("passed"): bucket["cases_passed"] += 1 else: bucket["cases_failed"] += 1 return dict(sorted(grouped.items()))