#!/usr/bin/env python3 """Audit raw smart-turn data and emit a leakage-grouped JSONL manifest.""" from __future__ import annotations import argparse import os import sys from collections.abc import Iterable, Iterator, Mapping from pathlib import Path from typing import Any PROJECT_ROOT = Path(__file__).resolve().parents[1] SRC_ROOT = PROJECT_ROOT / "src" if str(SRC_ROOT) not in sys.path: sys.path.insert(0, str(SRC_ROOT)) from turn_detection.data import ( # noqa: E402 DatasetReadError, GroupingConfig, OptionalDependencyError, audit_records, iter_records, write_json, write_manifest, ) DEFAULT_DATASET = PROJECT_ROOT / "data" / "raw" / "smart-turn-data-v3.2-train" DEFAULT_MANIFEST = PROJECT_ROOT / "data" / "processed" / "manifest.jsonl" DEFAULT_REPORT = PROJECT_ROOT / "artifacts" / "data_audit.json" def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description=( "Stream local Parquet/JSON/CSV or a Hugging Face dataset, validate every row, " "hash exact audio, and construct transitive leakage groups." ) ) parser.add_argument( "inputs", nargs="*", help=( "Local file/directory paths, or one Hugging Face dataset id. " f"Defaults to {DEFAULT_DATASET.relative_to(PROJECT_ROOT)}." ), ) parser.add_argument( "--output", type=Path, default=DEFAULT_MANIFEST, help="Output JSONL manifest" ) parser.add_argument( "--report", type=Path, default=DEFAULT_REPORT, help="Output audit report JSON" ) parser.add_argument("--hf-split", default="train", help="Hugging Face split name") parser.add_argument("--revision", help="Pinned Hugging Face dataset revision") parser.add_argument( "--token-env", default="HF_TOKEN", help="Environment variable containing the HF token" ) parser.add_argument( "--env-file", type=Path, default=PROJECT_ROOT / ".env", help="Optional KEY=VALUE file used only when --token-env is unset", ) parser.add_argument( "--batch-size", type=int, default=256, help="Parquet rows per streamed Arrow batch" ) parser.add_argument("--limit", type=int, help="Audit at most N records (useful for smoke runs)") parser.add_argument( "--progress-every", type=int, default=1000, help="Print progress every N input rows; set 0 to disable", ) parser.add_argument( "--no-text-grouping", action="store_true", help="Do not link repeated normalized transcripts/prompts", ) parser.add_argument( "--fail-on-error", action="store_true", help="Return a failure status when any invalid record is found (outputs are still written)", ) return parser def _token_from_env(name: str, env_file: Path | None) -> str | None: token = os.environ.get(name) if token or env_file is None or not env_file.is_file(): return token # Minimal dotenv parsing avoids a required dependency and never logs values. for raw_line in env_file.read_text(encoding="utf-8").splitlines(): line = raw_line.strip() if not line or line.startswith("#") or "=" not in line: continue key, value = line.split("=", 1) if key.strip() in (name, name.casefold()): value = value.strip() if len(value) >= 2 and value[0] == value[-1] and value[0] in ("'", '"'): value = value[1:-1] return value or None return None def _progress( records: Iterable[Mapping[str, Any]], *, every: int, ) -> Iterator[Mapping[str, Any]]: for count, record in enumerate(records, start=1): if every > 0 and count % every == 0: print(f"audited input rows: {count:,}", file=sys.stderr, flush=True) yield record def main(argv: list[str] | None = None) -> int: args = _parser().parse_args(argv) if args.batch_size <= 0: raise SystemExit("--batch-size must be positive") if args.limit is not None and args.limit < 0: raise SystemExit("--limit cannot be negative") if args.progress_every < 0: raise SystemExit("--progress-every cannot be negative") inputs = args.inputs or [str(DEFAULT_DATASET)] source: str | Path | list[str | Path] source = inputs[0] if len(inputs) == 1 else inputs token = _token_from_env(args.token_env, args.env_file) try: records = iter_records( source, split=args.hf_split, revision=args.revision, token=token, batch_size=args.batch_size, limit=args.limit, ) manifest, report = audit_records( _progress(records, every=args.progress_every), grouping_config=GroupingConfig(include_text=not args.no_text_grouping), ) write_manifest(args.output, manifest) write_json(args.report, report) except (FileNotFoundError, DatasetReadError, OptionalDependencyError, ValueError) as exc: print(f"audit failed: {exc}", file=sys.stderr) return 2 invalid = int(report["records"]["invalid"]) print( f"wrote {len(manifest):,} rows to {args.output} " f"({invalid:,} invalid); report: {args.report}", file=sys.stderr, ) return 1 if args.fail_on_error and invalid else 0 if __name__ == "__main__": raise SystemExit(main())