vignette/apps/api/app/services/dataset_export.py
2026-06-27 16:08:41 +09:00

400 lines
14 KiB
Python

"""Phase 3 recursive-learning dataset export helpers.
The exporter is intentionally conservative: it only emits masked text, keeps raw
database identifiers out of JSONL records, and never upgrades an artifact to an
approved seed dataset unless the explicit approval and agreement gates pass.
"""
from __future__ import annotations
import hashlib
import json
import math
import re
from collections import Counter, defaultdict
from dataclasses import dataclass, field
from datetime import UTC, date, datetime
from decimal import Decimal
from pathlib import Path
from typing import Any, Iterable, Mapping, Sequence
from uuid import UUID
DATASET_ITEM_SCHEMA = "phase3_dataset_item_v1"
APPROVED_EXPORT_STATUS = "approved_for_recursive_learning_seed"
DRY_RUN_EXPORT_STATUS = "technical_dry_run"
BLOCKED_EXPORT_STATUS = "blocked"
EXPORT_STATUSES = {APPROVED_EXPORT_STATUS, DRY_RUN_EXPORT_STATUS, BLOCKED_EXPORT_STATUS}
BLOCKED_FIELD_NAMES = {
"name",
"email",
"phone",
"student_id",
"national_id",
"address",
"date_of_birth",
"raw_audio_path",
"raw_voice",
"raw_source_case",
"identity_map",
"api_key",
"api_keys",
"access_token",
"refresh_token",
"token",
"cookie",
"credentials",
"credential",
"secret",
}
PII_PATTERNS: tuple[tuple[str, re.Pattern[str]], ...] = (
("email", re.compile(r"\b[A-Z0-9._%+-]+@[A-Z0-9.-]+\.[A-Z]{2,}\b", re.I)),
(
"phone",
re.compile(r"\b(?:\+?82[-. ]?)?(?:0?1[016789]|0[2-9]\d?)[-. ]?\d{3,4}[-. ]?\d{4}\b"),
),
("national_id", re.compile(r"\b\d{6}[- ]?[1-4]\d{6}\b")),
("student_id", re.compile(r"\b20\d{2}[- ]?\d{4,8}\b")),
(
"secret",
re.compile(
r"(?i)\b(?:api[_-]?key|access[_-]?token|refresh[_-]?token|session[_-]?secret|cookie)\b\s*[:=]\s*\S+"
),
),
("secret", re.compile(r"\b(?:sk-[A-Za-z0-9_-]{12,}|AIza[0-9A-Za-z_-]{20,}|xox[baprs]-[A-Za-z0-9-]+)\b")),
)
@dataclass
class ExportKeyMaps:
participant: dict[str, str] = field(default_factory=dict)
session: dict[str, str] = field(default_factory=dict)
def participant_key(self, raw_id: Any) -> str:
key = str(raw_id or "unknown-participant")
if key not in self.participant:
self.participant[key] = f"PX-{len(self.participant) + 1:04d}"
return self.participant[key]
def session_key(self, raw_id: Any) -> str:
key = str(raw_id or "unknown-session")
if key not in self.session:
self.session[key] = f"SX-{len(self.session) + 1:04d}"
return self.session[key]
def json_safe(value: Any) -> Any:
if isinstance(value, (datetime, date)):
if isinstance(value, datetime) and value.tzinfo is None:
value = value.replace(tzinfo=UTC)
return value.isoformat().replace("+00:00", "Z")
if isinstance(value, (UUID, Decimal)):
return str(value)
if isinstance(value, Mapping):
return {str(key): json_safe(item) for key, item in value.items()}
if isinstance(value, list):
return [json_safe(item) for item in value]
if isinstance(value, tuple):
return [json_safe(item) for item in value]
return value
def normalize_json_value(value: Any) -> Any:
if isinstance(value, str):
stripped = value.strip()
if stripped.startswith(("{", "[")):
try:
return json.loads(stripped)
except json.JSONDecodeError:
return value
return value
def _redacted_sample(kind: str, value: str) -> str:
if kind == "email" and "@" in value:
return f"<email:{value.rsplit('@', 1)[1].lower()}>"
return f"<{kind}>"
def _scan_text(value: str, path: str) -> list[dict[str, Any]]:
findings: list[dict[str, Any]] = []
for kind, pattern in PII_PATTERNS:
for match in pattern.finditer(value):
findings.append(
{
"kind": kind,
"path": path,
"sample": _redacted_sample(kind, match.group(0)),
}
)
return findings
def scan_for_pii(value: Any, path: str = "$") -> list[dict[str, Any]]:
findings: list[dict[str, Any]] = []
if isinstance(value, Mapping):
for raw_key, item in value.items():
key = str(raw_key)
next_path = f"{path}.{key}"
if key.lower() in BLOCKED_FIELD_NAMES:
findings.append({"kind": "blocked_field", "path": next_path, "sample": f"<{key.lower()}>"})
findings.extend(scan_for_pii(item, next_path))
return findings
if isinstance(value, list):
for index, item in enumerate(value):
findings.extend(scan_for_pii(item, f"{path}[{index}]"))
return findings
if isinstance(value, str):
findings.extend(_scan_text(value, path))
return findings
def build_dataset_record(
row: Mapping[str, Any],
*,
item_index: int,
export_manifest_id: str,
keys: ExportKeyMaps,
pii_scan_status: str = "pending",
consent_scope: str = "recursive_learning_seed",
) -> dict[str, Any]:
text_masked = str(row.get("text_masked") or "").strip()
if not text_masked:
raise ValueError("dataset export requires non-empty text_masked")
participant_key = keys.participant_key(row.get("learner_id"))
session_key = keys.session_key(row.get("session_id"))
persona_code = row.get("persona_code") or row.get("persona_id") or "unknown"
return {
"schema": DATASET_ITEM_SCHEMA,
"item_id": f"DI-{item_index:06d}",
"participant_key": participant_key,
"session_key": session_key,
"turn_key": f"TX-{item_index:06d}",
"persona_id": str(persona_code),
"stage": row.get("stage") or "",
"speaker": row.get("speaker") or "",
"text_masked": text_masked,
"techniques": json_safe(normalize_json_value(row.get("techniques") or [])),
"client_states": json_safe(normalize_json_value(row.get("client_states") or [])),
"feedback_scores": json_safe(normalize_json_value(row.get("feedback_scores") or [])),
"supervisor_comments": json_safe(normalize_json_value(row.get("supervisor_comments") or [])),
"source_refs": {
"session_started_at": json_safe(row.get("session_started_at") or row.get("started_at")),
"export_manifest_id": export_manifest_id,
},
"privacy": {
"direct_identifiers_removed": True,
"pii_scan_status": pii_scan_status,
"consent_scope": consent_scope,
},
}
def write_jsonl(records: Sequence[Mapping[str, Any]], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8", newline="\n") as handle:
for record in records:
handle.write(json.dumps(json_safe(record), ensure_ascii=False, sort_keys=True, separators=(",", ":")))
handle.write("\n")
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def cohen_kappa(annotations: Iterable[Mapping[str, Any]], label_key: str) -> float | None:
pairs: list[tuple[Any, Any]] = []
by_item: dict[Any, list[Any]] = defaultdict(list)
for annotation in annotations:
labels = normalize_json_value(annotation.get("labels") or {})
if not isinstance(labels, Mapping) or label_key not in labels:
continue
by_item[annotation.get("item_id")].append(labels[label_key])
for values in by_item.values():
if len(values) >= 2:
pairs.append((values[0], values[1]))
if not pairs:
return None
total = len(pairs)
observed = sum(1 for left, right in pairs if left == right) / total
left_counts = Counter(left for left, _ in pairs)
right_counts = Counter(right for _, right in pairs)
expected = sum((left_counts[label] / total) * (right_counts[label] / total) for label in set(left_counts) | set(right_counts))
if math.isclose(1.0, expected):
return 1.0 if math.isclose(1.0, observed) else None
return round((observed - expected) / (1.0 - expected), 4)
def intraclass_correlation(annotations: Iterable[Mapping[str, Any]], score_key: str) -> float | None:
by_item: dict[Any, list[float]] = defaultdict(list)
for annotation in annotations:
labels = normalize_json_value(annotation.get("labels") or {})
if not isinstance(labels, Mapping) or score_key not in labels:
continue
try:
by_item[annotation.get("item_id")].append(float(labels[score_key]))
except (TypeError, ValueError):
continue
matrix = [values[:2] for values in by_item.values() if len(values) >= 2]
if len(matrix) < 2:
return None
n = len(matrix)
k = 2
row_means = [sum(row) / k for row in matrix]
col_means = [sum(row[col] for row in matrix) / n for col in range(k)]
grand_mean = sum(row_means) / n
msr = k * sum((mean - grand_mean) ** 2 for mean in row_means) / (n - 1)
msc = n * sum((mean - grand_mean) ** 2 for mean in col_means) / (k - 1)
residual = 0.0
for row_index, row in enumerate(matrix):
for col_index, value in enumerate(row):
residual += (value - row_means[row_index] - col_means[col_index] + grand_mean) ** 2
mse = residual / ((n - 1) * (k - 1))
denominator = msr + (k - 1) * mse + (k * (msc - mse) / n)
if math.isclose(denominator, 0.0):
return None
return round((msr - mse) / denominator, 4)
def infer_source_window(records: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
starts = [record.get("source_refs", {}).get("session_started_at") for record in records]
starts = [value for value in starts if value]
return {
"started_at": min(starts) if starts else "",
"ended_at": max(starts) if starts else "",
}
def build_manifest(
*,
export_id: str,
dataset_name: str,
export_status: str,
purpose: str,
records: Sequence[Mapping[str, Any]],
jsonl_path: str,
jsonl_sha256: str,
pii_findings: Sequence[Mapping[str, Any]],
participants_included: int,
participants_excluded: int = 0,
cohort_id: str = "phase3",
consent_version: str = "",
agreement: Mapping[str, Any] | None = None,
approvals: Mapping[str, str] | None = None,
known_limitations: Sequence[str] | None = None,
created_at: datetime | None = None,
) -> dict[str, Any]:
if export_status not in EXPORT_STATUSES:
raise ValueError(f"unsupported export_status: {export_status}")
agreement_payload = {
"kappa": None,
"icc": None,
"gold_status": "not_gold",
}
if agreement:
agreement_payload.update(dict(agreement))
approvals_payload = {
"data_steward": "",
"legal_or_privacy_reviewer": "",
"technical_operator": "",
"approved_at": "",
}
if approvals:
approvals_payload.update({key: value for key, value in approvals.items() if value is not None})
pii_status = "pass" if not pii_findings else "fail"
limitations = list(known_limitations or [])
if export_status != APPROVED_EXPORT_STATUS:
limitations.append("technical dry-run only; data-steward/legal approval is not complete")
if pii_findings:
limitations.append("PII scan found records requiring reviewer disposition")
manifest = {
"export_id": export_id,
"dataset_name": dataset_name,
"export_status": export_status,
"created_at": json_safe(created_at or datetime.now(UTC)),
"purpose": purpose,
"source_window": infer_source_window(records),
"source_tables": [
"app.sessions",
"app.turns",
"app.feedback_scores",
"app.turn_technique",
"app.turn_client_state",
"app.supervisor_comment",
"ds.annotation",
"ds.export_manifest",
],
"selection_criteria": {
"cohort_id": cohort_id,
"min_completed_sessions": 0,
"include_withdrawn": False,
"excluded_safety_scope": ["self_harm_scenario_primary"],
},
"consent_scope": {
"consent_version": consent_version,
"allowed_uses": ["education_quality_review", "recursive_learning_seed"],
"withdrawal_cutoff_applied_at": "",
"participants_included": participants_included,
"participants_excluded": participants_excluded,
},
"anonymization": {
"participant_key": "pseudonymous export key; no identity map included",
"text_transform": "masked_text_only",
"direct_identifier_policy": "blocked",
"salt_or_identity_map_location": "not in export",
},
"pii_scan": {
"tool": "vignette.dataset_export.regex",
"version": "1",
"ran_at": json_safe(datetime.now(UTC)),
"status": pii_status,
"findings": [json_safe(finding) for finding in pii_findings],
},
"agreement": agreement_payload,
"files": [
{
"path": jsonl_path,
"rows": len(records),
"sha256": jsonl_sha256,
"schema": DATASET_ITEM_SCHEMA,
}
],
"approvals": approvals_payload,
"known_limitations": sorted(set(limitations)),
}
validate_manifest_gate(manifest)
return manifest
def validate_manifest_gate(manifest: Mapping[str, Any]) -> None:
if manifest.get("export_status") != APPROVED_EXPORT_STATUS:
return
errors: list[str] = []
pii_scan = manifest.get("pii_scan") or {}
agreement = manifest.get("agreement") or {}
approvals = manifest.get("approvals") or {}
if pii_scan.get("status") != "pass":
errors.append("PII scan must pass")
if (agreement.get("kappa") or 0) < 0.60:
errors.append("kappa must be >= 0.60")
if (agreement.get("icc") or 0) < 0.75:
errors.append("ICC must be >= 0.75")
for key in ("data_steward", "legal_or_privacy_reviewer", "technical_operator", "approved_at"):
if not str(approvals.get(key) or "").strip():
errors.append(f"approval missing: {key}")
if errors:
raise ValueError("; ".join(errors))