vignette/scripts/export-recursive-dataset.py
2026-06-27 16:08:41 +09:00

324 lines
11 KiB
Python

from __future__ import annotations
import argparse
import asyncio
import json
import os
import sys
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
API_ROOT = REPO_ROOT / "apps" / "api"
sys.path.insert(0, str(API_ROOT))
import asyncpg # noqa: E402
from app.services.dataset_export import ( # noqa: E402
APPROVED_EXPORT_STATUS,
DRY_RUN_EXPORT_STATUS,
ExportKeyMaps,
build_dataset_record,
build_manifest,
cohen_kappa,
intraclass_correlation,
scan_for_pii,
sha256_file,
validate_manifest_gate,
write_jsonl,
)
TURN_QUERY = """
SELECT
s.id::text AS session_id,
s.learner_id::text AS learner_id,
COALESCE(pc.code, s.persona_id::text, 'unknown') AS persona_code,
s.started_at AS session_started_at,
s.ended_at AS session_ended_at,
t.id::text AS turn_id,
t.seq,
t.speaker,
t.stage,
t.text_masked,
t.created_at,
COALESCE(
jsonb_agg(DISTINCT jsonb_build_object(
'dimension', fs.dimension,
'score', fs.score,
'rationale', fs.rationale,
'loop', fs.loop
)) FILTER (WHERE fs.dimension IS NOT NULL),
'[]'::jsonb
) AS feedback_scores,
COALESCE(
jsonb_agg(DISTINCT jsonb_build_object(
'code', tech.code,
'display_name', tech.display_name,
'category', tech.category
)) FILTER (WHERE tech.code IS NOT NULL),
'[]'::jsonb
) AS techniques,
COALESCE(
jsonb_agg(DISTINCT jsonb_build_object(
'code', cs.code,
'display_name', cs.display_name
)) FILTER (WHERE cs.code IS NOT NULL),
'[]'::jsonb
) AS client_states,
COALESCE(
jsonb_agg(DISTINCT jsonb_build_object(
'kind', sc.kind,
'text', sc.text,
'intent_deviation', sc.intent_deviation
)) FILTER (WHERE sc.id IS NOT NULL),
'[]'::jsonb
) AS supervisor_comments
FROM app.turns t
JOIN app.sessions s ON s.id = t.session_id
JOIN app.app_user u ON u.user_id = s.learner_id
LEFT JOIN app.persona_card pc
ON pc.persona_id = s.persona_id AND pc.version = s.persona_version
LEFT JOIN app.feedback_scores fs ON fs.turn_id = t.id
LEFT JOIN app.turn_technique tt ON tt.turn_id = t.id
LEFT JOIN app.technique_label_def tech ON tech.label_id = tt.label_id
LEFT JOIN app.turn_client_state tcs ON tcs.turn_id = t.id
LEFT JOIN app.client_state_def cs ON cs.label_id = tcs.label_id
LEFT JOIN app.supervisor_comment sc ON sc.turn_id = t.id
WHERE u.is_active
AND t.text_masked IS NOT NULL
AND btrim(t.text_masked) <> ''
AND ($1::timestamptz IS NULL OR t.created_at >= $1::timestamptz)
AND ($2::timestamptz IS NULL OR t.created_at < $2::timestamptz)
AND ($4::boolean = false OR u.cohort = $3::text)
GROUP BY s.id, s.learner_id, pc.code, t.id, t.seq, t.speaker, t.stage, t.text_masked, t.created_at
ORDER BY t.created_at, s.id, t.seq
LIMIT $5
"""
ANNOTATION_QUERY = """
SELECT
di.item_id::text AS item_id,
ar.round_id::text AS round_id,
a.annotator,
a.labels
FROM ds.annotation a
JOIN ds.annotation_round ar ON ar.round_id = a.round_id
JOIN ds.dataset_item di ON di.item_id = a.item_id
WHERE ar.dataset_id = $1
ORDER BY di.item_id, ar.round_id, a.annotation_id
"""
def parse_time(value: str | None) -> datetime | None:
if not value:
return None
normalized = value.replace("Z", "+00:00")
parsed = datetime.fromisoformat(normalized)
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=UTC)
return parsed
def make_export_id() -> str:
return f"phase3-rl-seed-{datetime.now(UTC).strftime('%Y%m%d')}-001"
def row_to_dict(row: asyncpg.Record) -> dict[str, Any]:
return {key: row[key] for key in row.keys()}
async def fetch_rows(conn: asyncpg.Connection, args: argparse.Namespace) -> list[dict[str, Any]]:
rows = await conn.fetch(
TURN_QUERY,
parse_time(args.started_at),
parse_time(args.ended_at),
args.cohort_id,
args.apply_cohort_filter,
args.limit,
)
return [row_to_dict(row) for row in rows]
async def write_dataset_rows(
conn: asyncpg.Connection,
*,
dataset_name: str,
purpose: str,
raw_rows: list[dict[str, Any]],
records: list[dict[str, Any]],
) -> int:
dataset_id = await conn.fetchval(
"INSERT INTO ds.dataset (name, purpose) VALUES ($1, $2) RETURNING dataset_id",
dataset_name,
purpose,
)
for raw_row, record in zip(raw_rows, records):
await conn.execute(
"INSERT INTO ds.dataset_item (dataset_id, turn_id, payload) VALUES ($1, $2::uuid, $3::jsonb)",
dataset_id,
raw_row["turn_id"],
json.dumps(record, ensure_ascii=False),
)
return int(dataset_id)
async def fetch_agreement(conn: asyncpg.Connection, dataset_id: int | None, args: argparse.Namespace) -> dict[str, Any]:
if dataset_id is None:
return {"kappa": None, "icc": None, "gold_status": "not_gold"}
rows = [row_to_dict(row) for row in await conn.fetch(ANNOTATION_QUERY, dataset_id)]
kappa = cohen_kappa(rows, args.kappa_label)
icc = intraclass_correlation(rows, args.icc_label)
gold_status = "gold_candidate" if kappa is not None and kappa >= 0.60 and icc is not None and icc >= 0.75 else "not_gold"
return {"kappa": kappa, "icc": icc, "gold_status": gold_status}
async def insert_manifest_row(
conn: asyncpg.Connection,
*,
dataset_id: int,
agreement: dict[str, Any],
jsonl_path: str,
) -> None:
await conn.execute(
"""
INSERT INTO ds.export_manifest (dataset_id, iaa_kappa, iaa_icc, jsonl_path)
VALUES ($1, $2, $3, $4)
""",
dataset_id,
agreement.get("kappa"),
agreement.get("icc"),
jsonl_path,
)
async def run(args: argparse.Namespace) -> int:
if args.export_status == APPROVED_EXPORT_STATUS and not args.allow_approved:
raise SystemExit("--export-status approved_for_recursive_learning_seed requires --allow-approved")
if args.write_manifest_row and not args.write_dataset:
raise SystemExit("--write-manifest-row requires --write-dataset")
output_root = Path(args.output_root).resolve()
export_dir = output_root / "03-export"
jsonl_path = export_dir / "anonymized_dataset.jsonl"
manifest_path = export_dir / "export_manifest.json"
conn = await asyncpg.connect(args.database_url)
try:
raw_rows = await fetch_rows(conn, args)
keys = ExportKeyMaps()
records = [
build_dataset_record(
row,
item_index=index,
export_manifest_id=args.export_id,
keys=keys,
pii_scan_status="pending",
)
for index, row in enumerate(raw_rows, start=1)
]
pii_findings: list[dict[str, Any]] = []
for record in records:
for finding in scan_for_pii(record):
pii_findings.append({"item_id": record["item_id"], **finding})
if pii_findings and args.fail_on_pii:
raise SystemExit(f"PII scan found {len(pii_findings)} finding(s)")
pii_status = "pass" if not pii_findings else "fail"
for record in records:
record["privacy"]["pii_scan_status"] = pii_status
write_jsonl(records, jsonl_path)
digest = sha256_file(jsonl_path)
dataset_id: int | None = None
if args.write_dataset:
dataset_id = await write_dataset_rows(
conn,
dataset_name=args.dataset_name,
purpose=args.purpose,
raw_rows=raw_rows,
records=records,
)
agreement = await fetch_agreement(conn, dataset_id, args)
approvals = {
"data_steward": args.data_steward,
"legal_or_privacy_reviewer": args.legal_or_privacy_reviewer,
"technical_operator": args.technical_operator or os.environ.get("USERNAME", ""),
"approved_at": args.approved_at,
}
manifest = build_manifest(
export_id=args.export_id,
dataset_name=args.dataset_name,
export_status=args.export_status,
purpose=args.purpose,
records=records,
jsonl_path=str(jsonl_path.relative_to(output_root)).replace("\\", "/"),
jsonl_sha256=digest,
pii_findings=pii_findings,
participants_included=len(keys.participant),
cohort_id=args.cohort_id,
consent_version=args.consent_version,
agreement=agreement,
approvals=approvals,
known_limitations=args.known_limitation,
)
validate_manifest_gate(manifest)
manifest_path.write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
if args.write_manifest_row and dataset_id is not None:
await insert_manifest_row(
conn,
dataset_id=dataset_id,
agreement=agreement,
jsonl_path=str(jsonl_path),
)
finally:
await conn.close()
print(json.dumps({"manifest": str(manifest_path), "jsonl": str(jsonl_path), "rows": len(records)}, ensure_ascii=False))
return 0
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Export a Phase 3 recursive-learning JSONL dry-run dataset.")
parser.add_argument("--database-url", default=os.environ.get("DATABASE_URL"), required=os.environ.get("DATABASE_URL") is None)
parser.add_argument("--output-root", default=str(REPO_ROOT / "data" / "phase3-dry-run"))
parser.add_argument("--export-id", default=make_export_id())
parser.add_argument("--dataset-name", default="vignette_phase3_recursive_learning_seed")
parser.add_argument("--purpose", default="recursive-learning seed dataset for education simulator improvement")
parser.add_argument("--export-status", choices=[DRY_RUN_EXPORT_STATUS, "blocked", APPROVED_EXPORT_STATUS], default=DRY_RUN_EXPORT_STATUS)
parser.add_argument("--allow-approved", action="store_true")
parser.add_argument("--started-at")
parser.add_argument("--ended-at")
parser.add_argument("--limit", type=int, default=1000)
parser.add_argument("--cohort-id", default="phase3")
parser.add_argument("--apply-cohort-filter", action="store_true")
parser.add_argument("--consent-version", default="")
parser.add_argument("--data-steward", default="")
parser.add_argument("--legal-or-privacy-reviewer", default="")
parser.add_argument("--technical-operator", default="")
parser.add_argument("--approved-at", default="")
parser.add_argument("--known-limitation", action="append", default=[])
parser.add_argument("--fail-on-pii", action="store_true")
parser.add_argument("--kappa-label", default="appropriateness")
parser.add_argument("--icc-label", default="rapport_signal")
parser.add_argument("--write-dataset", action="store_true")
parser.add_argument("--write-manifest-row", action="store_true")
return parser
def main() -> int:
parser = build_parser()
args = parser.parse_args()
return asyncio.run(run(args))
if __name__ == "__main__":
raise SystemExit(main())