diff --git a/tests/test_validate_data.py b/tests/test_validate_data.py new file mode 100644 index 0000000..62a0520 --- /dev/null +++ b/tests/test_validate_data.py @@ -0,0 +1,173 @@ +import importlib.util +import json +import sys +import tempfile +import unittest +from pathlib import Path + + +MODULE_PATH = Path(__file__).resolve().parents[1] / "validate-data.py" +SPEC = importlib.util.spec_from_file_location("validate_data", MODULE_PATH) +assert SPEC is not None and SPEC.loader is not None +validate_data = importlib.util.module_from_spec(SPEC) +sys.modules[SPEC.name] = validate_data +SPEC.loader.exec_module(validate_data) + + +def example(**overrides): + value = { + "prompt": "Add a health-check endpoint", + "purpose": "backendImpl", + "secondary": None, + "mixed": False, + "difficulty": 0.4, + "slice": "core", + "lang": "en", + } + value.update(overrides) + return value + + +class RecordValidationTests(unittest.TestCase): + def test_accepts_valid_record(self): + validator = validate_data.Validator( + fixture_path=None, process_checks=False + ) + + self.assertTrue(validator._check_record(Path("batch.jsonl"), 1, example())) + self.assertEqual([], validator.issues) + + def test_rejects_mixed_field_inconsistencies(self): + validator = validate_data.Validator( + fixture_path=None, process_checks=False + ) + + valid = validator._check_record( + Path("batch.jsonl"), + 7, + example(mixed=True, secondary=None, slice="core"), + ) + + self.assertFalse(valid) + messages = [issue.message for issue in validator.issues] + self.assertIn("mixed=true requires a secondary purpose", messages) + self.assertIn("the mixed slice and mixed field must agree", messages) + + def test_rejects_unknown_fields_and_non_finite_difficulty(self): + validator = validate_data.Validator( + fixture_path=None, process_checks=False + ) + + valid = validator._check_record( + Path("batch.jsonl"), + 3, + example(difficulty=float("nan"), surprise="value"), + ) + + self.assertFalse(valid) + messages = [issue.message for issue in validator.issues] + self.assertIn("unexpected fields: surprise", messages) + self.assertIn( + "difficulty must be a finite number from 0.0 to 1.0", messages + ) + + def test_bcp47_structure(self): + for tag in ("en", "pt-BR", "zh-Hans-CN", "de-DE-1996", "x-project"): + with self.subTest(tag=tag): + self.assertTrue(validate_data.is_bcp47(tag)) + for tag in ("", "e", "en_US", "-en", "en-", "en-☃"): + with self.subTest(tag=tag): + self.assertFalse(validate_data.is_bcp47(tag)) + + +class DatasetValidationTests(unittest.TestCase): + def test_valid_complete_batch_has_no_errors(self): + slices = ( + ["core"] * 110 + + ["boundary"] * 40 + + ["mixed"] * 20 + + ["pasted-context"] * 20 + + ["vague-eval"] * 10 + ) + languages = ["en"] * 190 + [ + "es", + "de", + "fr", + "pt", + "zh", + "ja", + "es-MX", + "de-DE", + "fr-CA", + "pt-BR", + ] + rows = [] + for index, (slice_name, lang) in enumerate(zip(slices, languages)): + if index < 30: + prompt = f"sample{index} task" + elif index < 40 or slice_name == "pasted-context": + prompt = f"sample{index} " + "context " * 110 + else: + prompt = ( + f"sample{index} perform a realistic scoped change in the " + "project with the listed constraints" + ) + purpose = validate_data.PURPOSES[index % len(validate_data.PURPOSES)] + is_mixed = slice_name == "mixed" + secondary = ( + validate_data.PURPOSES[(index + 1) % len(validate_data.PURPOSES)] + if is_mixed + else None + ) + rows.append( + example( + prompt=prompt, + purpose=purpose, + secondary=secondary, + mixed=is_mixed, + slice=slice_name, + lang=lang, + ) + ) + + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + batch = root / "batch.jsonl" + batch.write_text( + "".join(json.dumps(row) + "\n" for row in rows), encoding="utf-8" + ) + validator = validate_data.Validator( + batch_size=200, + expected_total=200, + fixture_path=None, + process_checks=True, + ) + records = validator.validate([batch], [root]) + + self.assertEqual(200, len(records)) + self.assertEqual( + [], + [issue for issue in validator.issues if issue.severity == "error"], + ) + + def test_reports_bad_json_and_normalized_duplicates(self): + with tempfile.TemporaryDirectory() as directory: + batch = Path(directory) / "batch.jsonl" + batch.write_text( + json.dumps(example(prompt="Fix spacing")) + "\n" + + "{not json}\n" + + json.dumps(example(prompt=" fix spacing ")) + "\n", + encoding="utf-8", + ) + validator = validate_data.Validator( + fixture_path=None, process_checks=False + ) + validator.validate([batch], [batch]) + + messages = [issue.message for issue in validator.issues] + self.assertTrue(any(message.startswith("invalid JSON:") for message in messages)) + self.assertTrue(any(message.startswith("duplicate prompt;") for message in messages)) + + +if __name__ == "__main__": + unittest.main() diff --git a/validate-data.py b/validate-data.py new file mode 100755 index 0000000..76df5ce --- /dev/null +++ b/validate-data.py @@ -0,0 +1,762 @@ +#!/usr/bin/env python3 +"""Validate purpose-classifier JSONL data. + +The checks in this script are derived from datagen-prompt.md. Definite format and +process violations are errors. Approximate targets (the requirements written as +"~" or "≈") are warnings; pass --strict to make warnings fail the command too. + +Usage: + python3 ml/purpose-classifier/validate-data.py + python3 ml/purpose-classifier/validate-data.py path/to/batch.jsonl + python3 ml/purpose-classifier/validate-data.py --strict +""" + +from __future__ import annotations + +import argparse +import json +import math +import re +import sys +import unicodedata +from collections import Counter +from dataclasses import dataclass +from datetime import date +from pathlib import Path +from typing import Any, Iterable + + +SCRIPT_DIR = Path(__file__).resolve().parent +DEFAULT_DATA_DIR = SCRIPT_DIR / "data" +DEFAULT_FIXTURES = ( + SCRIPT_DIR.parent.parent + / "Tests" + / "NucleicCoreTests" + / "Fixtures" + / "purpose-prompts.json" +) + +FIELDS = frozenset( + {"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"} +) +PURPOSES = ( + "planning", + "backendImpl", + "frontendImpl", + "quickFix", + "refactor", + "debugging", + "review", + "writing", +) +PURPOSE_SET = frozenset(PURPOSES) +SLICES = ("core", "boundary", "mixed", "pasted-context", "vague-eval") +SLICE_SET = frozenset(SLICES) +SLICE_TARGETS = { + "core": 0.55, + "boundary": 0.20, + "mixed": 0.10, + "pasted-context": 0.10, + "vague-eval": 0.05, +} +EXPECTED_NON_ENGLISH = frozenset({"es", "de", "fr", "pt", "zh", "ja"}) + +# This is deliberately a conservative list. It is only used for the prompt's +# mechanical "same opening verb no more than three times per 50 examples" check; +# unknown first words are not guessed to be verbs. +OPENING_VERBS = frozenset( + { + "add", + "analyze", + "architect", + "audit", + "break", + "build", + "bump", + "change", + "check", + "clean", + "compare", + "consolidate", + "correct", + "create", + "debug", + "decouple", + "design", + "diagnose", + "document", + "draft", + "evaluate", + "explain", + "extract", + "figure", + "find", + "fix", + "flip", + "help", + "implement", + "investigate", + "look", + "make", + "map", + "migrate", + "modularize", + "move", + "outline", + "plan", + "polish", + "proofread", + "propose", + "refactor", + "remove", + "rename", + "replace", + "restructure", + "review", + "rewrite", + "set", + "simplify", + "sketch", + "split", + "summarize", + "track", + "translate", + "tweak", + "update", + "walk", + "wire", + "write", + } +) + +TOKEN_RE = re.compile(r"\w+|[^\w\s]", re.UNICODE) +WORD_RE = re.compile(r"[^\W_]+(?:['’-][^\W_]+)*", re.UNICODE) +ENUMERATED_TASK_RE = re.compile(r"^\s*(?:[-*]\s*)?task\s+\d+\s*:", re.IGNORECASE) +LABEL_LEAK_RE = re.compile( + r"\b(?:this|it)\s+is\s+(?:an?\s+)?" + r"(?:planning|backendimpl|frontendimpl|quickfix|refactor|debugging|review|writing)" + r"\s+(?:task|prompt)\b", + re.IGNORECASE, +) +META_RE = re.compile(r"\b(?:classify this prompt|as an ai)\b", re.IGNORECASE) +MID_CONVERSATION_RE = re.compile( + r"^\s*(?:" + r"yes[,\s]+do\s+option\s+\d+|" + r"that\s+didn['’]?t\s+work|" + r"try\s+again\b|" + r"same\s+error\s+as\s+before|" + r"looks\s+good[,\s]+ship\s+it|" + r"no[,\s]+the\s+other\b" + r")", + re.IGNORECASE, +) + + +@dataclass(frozen=True) +class Issue: + severity: str + path: Path + message: str + line: int | None = None + + def render(self) -> str: + location = str(self.path) + if self.line is not None: + location += f":{self.line}" + return f"{location}: {self.severity}: {self.message}" + + +@dataclass(frozen=True) +class Record: + path: Path + line: int + value: dict[str, Any] + + +class Validator: + def __init__( + self, + *, + batch_size: int = 200, + expected_total: int = 8_000, + fixture_path: Path | None = DEFAULT_FIXTURES, + process_checks: bool = True, + ) -> None: + self.batch_size = batch_size + self.expected_total = expected_total + self.fixture_path = fixture_path + self.process_checks = process_checks + self.issues: list[Issue] = [] + + def error(self, path: Path, message: str, line: int | None = None) -> None: + self.issues.append(Issue("error", path, message, line)) + + def warning(self, path: Path, message: str, line: int | None = None) -> None: + self.issues.append(Issue("warning", path, message, line)) + + def validate(self, paths: list[Path], roots: list[Path]) -> list[Record]: + records_by_file: dict[Path, list[Record]] = {} + for path in paths: + records_by_file[path] = self._read_jsonl(path) + + records = [record for path in paths for record in records_by_file[path]] + self._check_duplicates(records) + self._check_fixture_contamination(records) + + if self.process_checks: + for path, file_records in records_by_file.items(): + self._check_file_batches(path, file_records) + self._check_global_distribution(records, roots) + self._check_manifests(roots) + + return records + + def _read_jsonl(self, path: Path) -> list[Record]: + records: list[Record] = [] + try: + lines = path.read_text(encoding="utf-8").splitlines() + except (OSError, UnicodeError) as exc: + self.error(path, f"cannot read UTF-8 JSONL: {exc}") + return records + + if not lines: + self.error(path, "file is empty") + return records + + for line_number, line in enumerate(lines, 1): + if not line.strip(): + self.error(path, "blank lines are not allowed in strict JSONL", line_number) + continue + try: + value = json.loads(line, parse_constant=self._reject_json_constant) + except (json.JSONDecodeError, ValueError) as exc: + self.error(path, f"invalid JSON: {exc}", line_number) + continue + if not isinstance(value, dict): + self.error(path, "each JSONL line must be an object", line_number) + continue + if self._check_record(path, line_number, value): + records.append(Record(path, line_number, value)) + return records + + @staticmethod + def _reject_json_constant(value: str) -> None: + raise ValueError(f"{value} is not valid strict JSON") + + def _check_record(self, path: Path, line: int, value: dict[str, Any]) -> bool: + valid = True + keys = set(value) + missing = sorted(FIELDS - keys) + extra = sorted(keys - FIELDS) + if missing: + self.error(path, f"missing fields: {', '.join(missing)}", line) + valid = False + if extra: + self.error(path, f"unexpected fields: {', '.join(extra)}", line) + valid = False + if missing: + return False + + prompt = value["prompt"] + if not isinstance(prompt, str): + self.error(path, "prompt must be a string", line) + valid = False + elif not prompt.strip(): + self.error(path, "prompt must not be empty or whitespace-only", line) + valid = False + + purpose = value["purpose"] + if not isinstance(purpose, str) or purpose not in PURPOSE_SET: + self.error(path, f"purpose must be one of: {', '.join(PURPOSES)}", line) + valid = False + + secondary = value["secondary"] + if secondary is not None and ( + not isinstance(secondary, str) or secondary not in PURPOSE_SET + ): + self.error(path, "secondary must be null or a valid purpose label", line) + valid = False + + mixed = value["mixed"] + if type(mixed) is not bool: + self.error(path, "mixed must be a boolean", line) + valid = False + + difficulty = value["difficulty"] + if ( + isinstance(difficulty, bool) + or not isinstance(difficulty, (int, float)) + or not math.isfinite(difficulty) + or not 0.0 <= difficulty <= 1.0 + ): + self.error(path, "difficulty must be a finite number from 0.0 to 1.0", line) + valid = False + + slice_name = value["slice"] + if not isinstance(slice_name, str) or slice_name not in SLICE_SET: + self.error(path, f"slice must be one of: {', '.join(SLICES)}", line) + valid = False + + lang = value["lang"] + if not isinstance(lang, str) or not is_bcp47(lang): + self.error(path, "lang must be a syntactically valid BCP 47 tag", line) + valid = False + + if type(mixed) is bool: + if mixed and secondary is None: + self.error(path, "mixed=true requires a secondary purpose", line) + valid = False + if not mixed and secondary is not None: + self.error(path, "mixed=false requires secondary=null", line) + valid = False + if (slice_name == "mixed") != mixed: + self.error(path, "the mixed slice and mixed field must agree", line) + valid = False + + if purpose in PURPOSE_SET and secondary == purpose: + self.error(path, "secondary must differ from the primary purpose", line) + valid = False + + if isinstance(prompt, str) and prompt.strip(): + self._check_prompt_antipatterns(path, line, prompt) + tokens = token_count(prompt) + if slice_name == "pasted-context" and not 100 <= tokens <= 400: + self.warning( + path, + "pasted-context prompt should be approximately 100-400 tokens " + f"(estimated {tokens})", + line, + ) + elif tokens > 400: + self.warning( + path, + f"prompt exceeds the approximately 400-token maximum (estimated {tokens})", + line, + ) + + return valid + + def _check_prompt_antipatterns(self, path: Path, line: int, prompt: str) -> None: + if ENUMERATED_TASK_RE.search(prompt): + self.error(path, "enumerated 'Task N:' prefix is forbidden", line) + if LABEL_LEAK_RE.search(prompt): + self.error(path, "prompt leaks its dataset label as a task/prompt hint", line) + if META_RE.search(prompt): + self.error(path, "assistant-directed classification/meta prompt is forbidden", line) + if MID_CONVERSATION_RE.search(prompt): + self.error(path, "prompt reads as a mid-conversation reply", line) + + def _check_duplicates(self, records: list[Record]) -> None: + seen: dict[str, Record] = {} + for record in records: + key = normalize_prompt(record.value["prompt"]) + previous = seen.get(key) + if previous is not None: + self.error( + record.path, + f"duplicate prompt; first seen at {previous.path}:{previous.line}", + record.line, + ) + else: + seen[key] = record + + def _check_fixture_contamination(self, records: list[Record]) -> None: + if self.fixture_path is None: + return + try: + fixture_data = json.loads(self.fixture_path.read_text(encoding="utf-8")) + fixture_prompts = { + normalize_prompt(item["prompt"]) + for item in fixture_data + if isinstance(item, dict) and isinstance(item.get("prompt"), str) + } + except (OSError, UnicodeError, json.JSONDecodeError, TypeError, KeyError) as exc: + self.warning(self.fixture_path, f"could not load eval-only fixtures: {exc}") + return + + for record in records: + if normalize_prompt(record.value["prompt"]) in fixture_prompts: + self.error( + record.path, + "prompt duplicates an eval-only purpose-prompts.json fixture", + record.line, + ) + + def _check_file_batches(self, path: Path, records: list[Record]) -> None: + if self.batch_size <= 0: + return + if len(records) % self.batch_size: + self.error( + path, + f"contains {len(records)} valid records; generation batches must contain " + f"{self.batch_size} records", + ) + + for start in range(0, len(records), self.batch_size): + batch = records[start : start + self.batch_size] + if batch: + self._check_batch(path, start // self.batch_size + 1, batch) + + def _check_batch(self, path: Path, number: int, records: list[Record]) -> None: + name = f"batch {number}" + total = len(records) + minimum_length_count = math.ceil(total * 0.15) + short_count = sum(word_count(r.value["prompt"]) < 8 for r in records) + long_count = sum(token_count(r.value["prompt"]) > 60 for r in records) + if short_count < minimum_length_count: + self.error( + path, + f"{name} has {short_count}/{total} prompts under 8 words; " + f"at least {minimum_length_count} required", + ) + if long_count < minimum_length_count: + self.error( + path, + f"{name} has {long_count}/{total} prompts over 60 estimated tokens; " + f"at least {minimum_length_count} required", + ) + + slices = Counter(r.value["slice"] for r in records) + for slice_name, target in SLICE_TARGETS.items(): + actual = slices[slice_name] / total + if abs(actual - target) > 0.05: + self.warning( + path, + f"{name} {slice_name} share is {actual:.1%}; target is approximately " + f"{target:.0%} (±5 percentage points)", + ) + + mixed_share = sum(r.value["mixed"] for r in records) / total + if mixed_share > 0.12: + self.error( + path, + f"{name} mixed-intent share is {mixed_share:.1%}; maximum is approximately 12%", + ) + + english = sum(primary_language(r.value["lang"]) == "en" for r in records) + english_share = english / total + if not 0.90 <= english_share <= 0.98: + self.warning( + path, + f"{name} English share is {english_share:.1%}; target is approximately 95%", + ) + + foreign_languages = { + primary_language(r.value["lang"]) + for r in records + if primary_language(r.value["lang"]) != "en" + } + unexpected = sorted(foreign_languages - EXPECTED_NON_ENGLISH) + if unexpected: + self.warning( + path, + f"{name} uses non-English languages outside the requested set: " + f"{', '.join(unexpected)}", + ) + + for window_start in range(0, total, 50): + window = records[window_start : window_start + 50] + verbs: dict[str, list[Record]] = {} + for record in window: + verb = opening_verb(record.value["prompt"]) + if verb is not None: + verbs.setdefault(verb, []).append(record) + for verb, matches in sorted(verbs.items()): + if len(matches) > 3: + locations = ", ".join(str(r.line) for r in matches) + self.error( + path, + f"{name}, records {window_start + 1}-{window_start + len(window)} " + f"open with '{verb}' {len(matches)} times (lines {locations}); maximum is 3", + ) + + def _check_global_distribution(self, records: list[Record], roots: list[Path]) -> None: + label = roots[0] if roots else DEFAULT_DATA_DIR + total = len(records) + if self.expected_total > 0 and total != self.expected_total: + self.warning( + label, + f"dataset contains {total} valid records; generation target is " + f"{self.expected_total}", + ) + if not records: + return + + non_vague = [r for r in records if r.value["slice"] != "vague-eval"] + if len(non_vague) >= len(PURPOSES): + counts = Counter(r.value["purpose"] for r in non_vague) + expected = len(non_vague) / len(PURPOSES) + low = expected * 0.85 + high = expected * 1.15 + outside = [ + f"{purpose}={counts[purpose]}" + for purpose in PURPOSES + if not low <= counts[purpose] <= high + ] + if outside: + self.error( + label, + "non-vague primary labels are not within ±15% of uniform " + f"(expected about {expected:.1f} each): {', '.join(outside)}", + ) + + mixed_share = sum(r.value["mixed"] for r in records) / total + if mixed_share > 0.12: + self.error( + label, + f"dataset mixed-intent share is {mixed_share:.1%}; maximum is approximately 12%", + ) + + slices = Counter(r.value["slice"] for r in records) + for slice_name, target in SLICE_TARGETS.items(): + actual = slices[slice_name] / total + if abs(actual - target) > 0.05: + self.warning( + label, + f"dataset {slice_name} share is {actual:.1%}; target is approximately " + f"{target:.0%} (±5 percentage points)", + ) + + def _check_manifests(self, roots: list[Path]) -> None: + directories = sorted({root if root.is_dir() else root.parent for root in roots}) + for directory in directories: + candidates = sorted(directory.glob("*manifest*.json")) + if not candidates: + self.warning( + directory, + "generation manifest not found; record the model, date, and batch topics", + ) + continue + for candidate in candidates: + self._check_manifest(candidate) + + def _check_manifest(self, path: Path) -> None: + try: + value = json.loads(path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + self.error(path, f"invalid generation manifest: {exc}") + return + if not isinstance(value, dict): + self.error(path, "generation manifest must be a JSON object") + return + + values_by_key: dict[str, list[Any]] = {} + for key, item in walk_mapping_items(value): + normalized_key = re.sub(r"[-_]", "", key).casefold() + values_by_key.setdefault(normalized_key, []).append(item) + + models = values_by_key.get("model", []) + values_by_key.get("generatingmodel", []) + if not any(isinstance(item, str) and item.strip() for item in models): + self.error(path, "generation manifest must record a non-empty model") + + dates = values_by_key.get("date", []) + values_by_key.get("generationdate", []) + valid_dates = [item for item in dates if isinstance(item, str) and is_iso_date(item)] + if not valid_dates: + self.error(path, "generation manifest must record an ISO date (YYYY-MM-DD)") + + topics = values_by_key.get("topics", []) + values_by_key.get("batchtopics", []) + if not any( + (isinstance(item, str) and item.strip()) + or (isinstance(item, list) and len(item) > 0) + for item in topics + ): + self.error(path, "generation manifest must record non-empty batch topics") + + +def is_bcp47(value: str) -> bool: + """A dependency-free structural BCP 47 check. + + Full registry validation would make this script network- or package-dependent. + This rejects the common malformed forms while accepting normal language, + script, region, variant, extension, and private-use tags. + """ + + if not value or "_" in value or value.startswith("-") or value.endswith("-"): + return False + parts = value.split("-") + if any(not part.isascii() or not part.isalnum() or not 1 <= len(part) <= 8 for part in parts): + return False + if parts[0].casefold() == "x": + return len(parts) > 1 + return parts[0].isalpha() and 2 <= len(parts[0]) <= 8 + + +def normalize_prompt(prompt: str) -> str: + normalized = unicodedata.normalize("NFKC", prompt).casefold() + return " ".join(normalized.split()) + + +def word_count(prompt: str) -> int: + return len(WORD_RE.findall(prompt)) + + +def token_count(prompt: str) -> int: + return len(TOKEN_RE.findall(prompt)) + + +def primary_language(tag: str) -> str: + return tag.split("-", 1)[0].casefold() + + +def opening_verb(prompt: str) -> str | None: + words = [word.casefold() for word in WORD_RE.findall(prompt[:160])] + if not words: + return None + if words[0] in OPENING_VERBS: + return words[0] + + # Recognize common request wrappers without mistaking a later noun ("a small + # change") for the sentence's opening verb. + start = 0 + if words[0] in {"please", "kindly"}: + start = 1 + elif len(words) >= 2 and words[0] in {"can", "could", "would", "will"}: + start = 2 if words[1] == "you" else 1 + elif len(words) >= 2 and words[0] in {"i", "we"} and words[1] in { + "need", + "want", + "would", + }: + start = 2 + elif len(words) >= 2 and words[0] == "help" and words[1] in {"me", "us"}: + start = 2 + else: + return None + + for word in words[start : start + 3]: + if word in OPENING_VERBS: + return word + return None + + +def walk_mapping_items(value: Any) -> Iterable[tuple[str, Any]]: + if isinstance(value, dict): + for key, item in value.items(): + if isinstance(key, str): + yield key, item + yield from walk_mapping_items(item) + elif isinstance(value, list): + for item in value: + yield from walk_mapping_items(item) + + +def is_iso_date(value: str) -> bool: + try: + date.fromisoformat(value) + except ValueError: + return False + return bool(re.fullmatch(r"\d{4}-\d{2}-\d{2}", value)) + + +def discover_paths(targets: list[Path]) -> tuple[list[Path], list[Path], list[str]]: + files: set[Path] = set() + roots: list[Path] = [] + errors: list[str] = [] + for target in targets: + target = target.resolve() + if not target.exists(): + errors.append(f"{target}: path does not exist") + continue + roots.append(target) + if target.is_dir(): + files.update(path.resolve() for path in target.rglob("*.jsonl") if path.is_file()) + elif target.is_file() and target.suffix.casefold() == ".jsonl": + files.add(target) + else: + errors.append(f"{target}: expected a .jsonl file or directory") + if not files and not errors: + errors.append("no .jsonl files found") + return sorted(files), roots, errors + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Validate purpose-classifier JSONL data against datagen-prompt.md." + ) + parser.add_argument( + "paths", + nargs="*", + type=Path, + help=f"JSONL files or directories (default: {DEFAULT_DATA_DIR})", + ) + parser.add_argument( + "--strict", + action="store_true", + help="return failure for approximate-target warnings as well as errors", + ) + parser.add_argument( + "--batch-size", + type=int, + default=200, + help="required generation batch size; use 0 to disable (default: 200)", + ) + parser.add_argument( + "--expected-total", + type=int, + default=8_000, + help="expected total record count; use 0 to disable (default: 8000)", + ) + parser.add_argument( + "--fixtures", + type=Path, + default=DEFAULT_FIXTURES, + help="eval-only fixture JSON checked for contamination", + ) + parser.add_argument( + "--no-fixture-check", + action="store_true", + help="do not check prompts against the eval-only fixture set", + ) + parser.add_argument( + "--no-process-checks", + action="store_true", + help="only check JSONL records, schema, duplicates, and fixture contamination", + ) + parser.add_argument( + "--max-issues", + type=int, + default=200, + help="maximum diagnostics to print; 0 prints all (default: 200)", + ) + return parser + + +def main(argv: list[str] | None = None) -> int: + args = build_parser().parse_args(argv) + if args.batch_size < 0 or args.expected_total < 0 or args.max_issues < 0: + print("error: numeric options must be non-negative", file=sys.stderr) + return 2 + + targets = args.paths or [DEFAULT_DATA_DIR] + paths, roots, discovery_errors = discover_paths(targets) + if discovery_errors: + for message in discovery_errors: + print(f"error: {message}", file=sys.stderr) + return 2 + + validator = Validator( + batch_size=args.batch_size, + expected_total=args.expected_total, + fixture_path=None if args.no_fixture_check else args.fixtures.resolve(), + process_checks=not args.no_process_checks, + ) + records = validator.validate(paths, roots) + errors = sum(issue.severity == "error" for issue in validator.issues) + warnings = sum(issue.severity == "warning" for issue in validator.issues) + + visible = validator.issues if args.max_issues == 0 else validator.issues[: args.max_issues] + for issue in visible: + print(issue.render(), file=sys.stderr) + hidden = len(validator.issues) - len(visible) + if hidden: + print(f"... {hidden} additional issues omitted", file=sys.stderr) + + print( + f"Validated {len(records)} records in {len(paths)} JSONL files: " + f"{errors} error(s), {warnings} warning(s)." + ) + return 1 if errors or (args.strict and warnings) else 0 + + +if __name__ == "__main__": + raise SystemExit(main())