90 lines
3.6 KiB
Python
90 lines
3.6 KiB
Python
"""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]
|
|
|
|
|
|
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) -> dict[str, Any]:
|
|
case_id = str(case.get("id") or "")
|
|
text = str(case.get("text") or "")
|
|
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)
|
|
|
|
return {
|
|
"id": case_id,
|
|
"passed": passed,
|
|
"entities": sorted(entities),
|
|
"masked_text": result.text_masked,
|
|
"missing_entities": missing_entities,
|
|
"unexpected_entities": unexpected_detected,
|
|
"forbidden_remaining": forbidden_remaining,
|
|
"required_missing": required_missing,
|
|
}
|
|
|
|
|
|
def evaluate_cases(
|
|
cases: Iterable[Mapping[str, Any]],
|
|
*,
|
|
mask_func: MaskFunc = guardrail.mask_pii,
|
|
) -> dict[str, Any]:
|
|
case_list = list(cases)
|
|
results = [evaluate_case(case, mask_func=mask_func) 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)
|
|
removed_forbidden += len(forbidden) - len(result["forbidden_remaining"])
|
|
unexpected_violations += len(result["unexpected_entities"])
|
|
|
|
passed_cases = sum(1 for result in results if result["passed"])
|
|
return {
|
|
"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,
|
|
"results": results,
|
|
}
|
|
|
|
|
|
def evaluate_fixture(path: Path, *, mask_func: MaskFunc = guardrail.mask_pii) -> dict[str, Any]:
|
|
return evaluate_cases(load_cases(path), mask_func=mask_func)
|
|
|
|
|
|
def _ratio(numerator: int, denominator: int) -> float:
|
|
if denominator <= 0:
|
|
return 1.0
|
|
return round(numerator / denominator, 4)
|