"""Shared data contracts for the purpose-classifier pipeline.""" from __future__ import annotations import hashlib import json import math import re import unicodedata from collections import Counter, defaultdict from dataclasses import dataclass from pathlib import Path from typing import Any, Callable, Iterable, Sequence LABELS = ( "planning", "backendImpl", "frontendImpl", "quickFix", "refactor", "debugging", "review", "writing", ) LABEL_SET = frozenset(LABELS) SOURCE_FIELDS = frozenset( {"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"} ) SLICES = frozenset( {"core", "boundary", "mixed", "pasted-context", "vague-eval"} ) HARD_SLICES = frozenset({"boundary", "mixed", "pasted-context", "vague-eval"}) WORD_RE = re.compile(r"\w+", re.UNICODE) class DataError(ValueError): """A deterministic data-contract failure.""" @dataclass(frozen=True) class SourceRecord: value: dict[str, Any] source: Path line: int @dataclass(frozen=True) class Duplicate: dropped: SourceRecord matched_prompt_hash: str kind: str similarity: float @dataclass(frozen=True) class CurationResult: records: list[SourceRecord] duplicates: list[Duplicate] @dataclass(frozen=True) class SplitResult: train: list[SourceRecord] validation: list[SourceRecord] test: list[SourceRecord] fixture_count: int @property def logical_test_count(self) -> int: return len(self.test) + self.fixture_count def normalize_prompt(prompt: str) -> str: """Match the runtime's whitespace collapse and add stable Unicode normalization.""" normalized = unicodedata.normalize("NFKC", prompt) return " ".join(normalized.split()) def normalized_key(prompt: str) -> str: return normalize_prompt(prompt).casefold() def prompt_hash(prompt: str) -> str: return hashlib.sha256(normalized_key(prompt).encode("utf-8")).hexdigest() def file_sha256(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 canonical_json(record: dict[str, Any]) -> str: return json.dumps(record, ensure_ascii=False, separators=(",", ":")) def jsonl_bytes(records: Iterable[dict[str, Any]]) -> bytes: return ("".join(f"{canonical_json(record)}\n" for record in records)).encode("utf-8") def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) path.write_bytes(jsonl_bytes(records)) def write_json(path: Path, value: Any) -> None: path.parent.mkdir(parents=True, exist_ok=True) path.write_text( json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8", ) def load_jsonl(path: Path) -> list[dict[str, Any]]: records: list[dict[str, Any]] = [] try: lines = path.read_text(encoding="utf-8").splitlines() except (OSError, UnicodeError) as exc: raise DataError(f"{path}: cannot read UTF-8 JSONL: {exc}") from exc for line_number, line in enumerate(lines, 1): if not line.strip(): raise DataError(f"{path}:{line_number}: blank JSONL line") try: value = json.loads(line) except json.JSONDecodeError as exc: raise DataError(f"{path}:{line_number}: invalid JSON: {exc}") from exc if not isinstance(value, dict): raise DataError(f"{path}:{line_number}: expected a JSON object") records.append(value) return records def validate_source_record(record: dict[str, Any], location: str) -> None: keys = set(record) if keys != SOURCE_FIELDS: missing = sorted(SOURCE_FIELDS - keys) extra = sorted(keys - SOURCE_FIELDS) details = [] if missing: details.append(f"missing {', '.join(missing)}") if extra: details.append(f"unexpected {', '.join(extra)}") raise DataError(f"{location}: invalid fields ({'; '.join(details)})") prompt = record["prompt"] if not isinstance(prompt, str) or not prompt.strip(): raise DataError(f"{location}: prompt must be a non-empty string") if record["purpose"] not in LABEL_SET: raise DataError(f"{location}: invalid purpose {record['purpose']!r}") secondary = record["secondary"] mixed = record["mixed"] if type(mixed) is not bool: raise DataError(f"{location}: mixed must be a boolean") if secondary is not None and secondary not in LABEL_SET: raise DataError(f"{location}: invalid secondary purpose {secondary!r}") if mixed != (secondary is not None): raise DataError(f"{location}: mixed and secondary disagree") if secondary == record["purpose"]: raise DataError(f"{location}: secondary must differ from purpose") difficulty = record["difficulty"] if ( isinstance(difficulty, bool) or not isinstance(difficulty, (int, float)) or not math.isfinite(difficulty) or not 0.0 <= difficulty <= 1.0 ): raise DataError(f"{location}: difficulty must be a finite value from 0 to 1") if record["slice"] not in SLICES: raise DataError(f"{location}: invalid slice {record['slice']!r}") if (record["slice"] == "mixed") != mixed: raise DataError(f"{location}: the mixed slice and mixed field disagree") if not isinstance(record["lang"], str) or not record["lang"]: raise DataError(f"{location}: lang must be a non-empty string") def load_sources(paths: Sequence[Path]) -> list[SourceRecord]: records: list[SourceRecord] = [] for path in paths: for line, value in enumerate(load_jsonl(path), 1): validate_source_record(value, f"{path}:{line}") records.append(SourceRecord(value=value, source=path, line=line)) return records def load_classifiable_fixtures(path: Path) -> list[dict[str, str]]: try: value = json.loads(path.read_text(encoding="utf-8")) except (OSError, UnicodeError, json.JSONDecodeError) as exc: raise DataError(f"{path}: cannot load fixtures: {exc}") from exc if not isinstance(value, list): raise DataError(f"{path}: fixture root must be an array") fixtures: list[dict[str, str]] = [] for index, item in enumerate(value): if not isinstance(item, dict): raise DataError(f"{path}: fixture {index} must be an object") prompt = item.get("prompt") purpose = item.get("purpose") if not isinstance(prompt, str) or not prompt.strip(): raise DataError(f"{path}: fixture {index} has an invalid prompt") if purpose == "general": continue if purpose not in LABEL_SET: raise DataError(f"{path}: fixture {index} has invalid purpose {purpose!r}") fixtures.append({"prompt": prompt, "purpose": purpose}) return fixtures def _word_shingles(prompt: str) -> frozenset[str]: words = WORD_RE.findall(normalized_key(prompt)) if len(words) < 8: return frozenset() return frozenset( "\x1f".join(words[index : index + 3]) for index in range(len(words) - 2) ) def _simhash(features: frozenset[str]) -> int: weights = [0] * 64 for feature in features: value = int.from_bytes( hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest(), "big" ) for bit in range(64): weights[bit] += 1 if value & (1 << bit) else -1 signature = 0 for bit, weight in enumerate(weights): if weight >= 0: signature |= 1 << bit return signature class _NearDuplicateIndex: """Small dependency-free LSH pass for generated template duplicates. This is intentionally conservative. Short prompts are handled only by exact-match checks because a one-word change can completely change their purpose. Longer prompts become word-trigram sets. SimHash bands produce candidates; exact Jaccard similarity decides whether a row is dropped. """ def __init__(self, threshold: float) -> None: self.threshold = threshold self.items: list[tuple[frozenset[str], str, str]] = [] self.buckets: dict[tuple[int, int], list[int]] = defaultdict(list) @staticmethod def _bands(signature: int) -> Iterable[tuple[int, int]]: # Requiring two matching 8-bit bands keeps random candidate sets small while # retaining every pair whose SimHashes differ in at most six bands. for band in range(8): yield band, (signature >> (band * 8)) & 0xFF def find(self, prompt: str) -> tuple[str, str, float] | None: features = _word_shingles(prompt) if not features: return None signature = _simhash(features) hits: Counter[int] = Counter() for band in self._bands(signature): hits.update(self.buckets.get(band, ())) best: tuple[str, str, float] | None = None for index, matching_bands in hits.items(): if matching_bands < 2: continue other_features, other_label, other_hash = self.items[index] union = len(features | other_features) similarity = len(features & other_features) / union if union else 1.0 if similarity >= self.threshold and ( best is None or similarity > best[2] ): best = (other_label, other_hash, similarity) return best def add(self, prompt: str, label: str) -> None: features = _word_shingles(prompt) if not features: return signature = _simhash(features) index = len(self.items) self.items.append((features, label, prompt_hash(prompt))) for band in self._bands(signature): self.buckets[band].append(index) def curate_records( records: Sequence[SourceRecord], fixtures: Sequence[dict[str, str]], *, near_duplicate_threshold: float = 0.92, ) -> CurationResult: if not 0.0 < near_duplicate_threshold <= 1.0: raise DataError("near-duplicate threshold must be in (0, 1]") exact: dict[str, tuple[str, str]] = {} near = _NearDuplicateIndex(near_duplicate_threshold) for fixture in fixtures: key = normalized_key(fixture["prompt"]) previous = exact.get(key) if previous is not None and previous[0] != fixture["purpose"]: raise DataError("shipped fixtures contain an exact prompt with two labels") exact[key] = (fixture["purpose"], prompt_hash(fixture["prompt"])) near.add(fixture["prompt"], fixture["purpose"]) kept: list[SourceRecord] = [] duplicates: list[Duplicate] = [] conflicts: list[str] = [] for record in records: prompt = record.value["prompt"] purpose = record.value["purpose"] key = normalized_key(prompt) previous = exact.get(key) if previous is not None: previous_label, previous_hash = previous if previous_label != purpose: conflicts.append( f"{record.source}:{record.line}: exact duplicate has labels " f"{previous_label!r} and {purpose!r}" ) else: duplicates.append( Duplicate(record, previous_hash, "exact", 1.0) ) continue match = near.find(prompt) if match is not None: previous_label, previous_hash, similarity = match if previous_label != purpose: conflicts.append( f"{record.source}:{record.line}: {similarity:.1%}-similar prompt " f"has labels {previous_label!r} and {purpose!r}" ) # Keep the row for now so all conflicts are reported without causing a # cascade of duplicates against a record that may later be relabeled. exact[key] = (purpose, prompt_hash(prompt)) near.add(prompt, purpose) kept.append(record) else: duplicates.append( Duplicate(record, previous_hash, "near", similarity) ) continue exact[key] = (purpose, prompt_hash(prompt)) near.add(prompt, purpose) kept.append(record) if conflicts: preview = "\n".join(conflicts[:20]) remainder = len(conflicts) - min(20, len(conflicts)) suffix = f"\n... {remainder} more conflict(s)" if remainder else "" raise DataError(f"near-duplicate label conflicts require review:\n{preview}{suffix}") return CurationResult(records=kept, duplicates=duplicates) def _stable_digest(seed: int, prompt: str) -> str: material = f"{seed}\0{normalized_key(prompt)}".encode("utf-8") return hashlib.sha256(material).hexdigest() def _stratified_order( records: Sequence[SourceRecord], *, seed: int, strata: Callable[[SourceRecord], tuple[str, ...]], ) -> list[SourceRecord]: groups: dict[tuple[str, ...], list[SourceRecord]] = defaultdict(list) for record in records: groups[strata(record)].append(record) ranked: list[tuple[float, str, SourceRecord]] = [] for key in sorted(groups): group = sorted( groups[key], key=lambda record: _stable_digest(seed, record.value["prompt"]), ) size = len(group) for index, record in enumerate(group): quantile = (index + 0.5) / size ranked.append( (quantile, _stable_digest(seed + 1, record.value["prompt"]), record) ) return [item[2] for item in sorted(ranked, key=lambda item: (item[0], item[1]))] def split_records( records: Sequence[SourceRecord], *, fixture_count: int, seed: int = 0xC1A551F1, train_ratio: float = 0.8, validation_ratio: float = 0.1, ) -> SplitResult: if not records: raise DataError("cannot split an empty dataset") if fixture_count < 0: raise DataError("fixture_count cannot be negative") if not 0.0 < train_ratio < 1.0 or not 0.0 < validation_ratio < 1.0: raise DataError("split ratios must be in (0, 1)") if train_ratio + validation_ratio >= 1.0: raise DataError("train + validation ratios must leave room for test") logical_total = len(records) + fixture_count target_train = round(logical_total * train_ratio) target_validation = round(logical_total * validation_ratio) target_test = logical_total - target_train - target_validation if fixture_count > target_test: raise DataError("fixture count exceeds the target test split") vague = [record for record in records if record.value["slice"] == "vague-eval"] regular = [record for record in records if record.value["slice"] != "vague-eval"] eval_capacity = target_validation + target_test - fixture_count if len(vague) > eval_capacity: raise DataError( "vague-eval records exceed validation/test capacity; lower train_ratio" ) # Allocate vague records between validation and test in proportion to each split's # remaining capacity. None may enter training. test_source_capacity = target_test - fixture_count vague_validation_count = round( len(vague) * target_validation / (target_validation + test_source_capacity) ) vague_validation_count = min(vague_validation_count, target_validation) vague_test_count = len(vague) - vague_validation_count if vague_test_count > test_source_capacity: overflow = vague_test_count - test_source_capacity vague_validation_count += overflow vague_test_count -= overflow vague_order = _stratified_order( vague, seed=seed + 7, strata=lambda record: (record.value["purpose"],), ) vague_validation = vague_order[:vague_validation_count] vague_test = vague_order[vague_validation_count:] regular_train_count = target_train regular_validation_count = target_validation - len(vague_validation) regular_test_count = test_source_capacity - len(vague_test) if ( regular_train_count + regular_validation_count + regular_test_count != len(regular) ): raise DataError("internal split accounting mismatch") regular_order = _stratified_order( regular, seed=seed, strata=lambda record: ( record.value["purpose"], record.value["slice"], record.value["lang"].split("-", 1)[0].casefold() if record.value["lang"].split("-", 1)[0].casefold() != "en" else "en", ), ) train = regular_order[:regular_train_count] validation_end = regular_train_count + regular_validation_count validation = regular_order[regular_train_count:validation_end] + vague_validation test = regular_order[validation_end:] + vague_test # A second stable ordering makes file content independent of stratum dictionary order # and gives the trainer a deterministic shuffle before its epoch sampler takes over. def stable_sort(items: Sequence[SourceRecord], offset: int) -> list[SourceRecord]: return sorted( items, key=lambda record: _stable_digest(seed + offset, record.value["prompt"]), ) result = SplitResult( train=stable_sort(train, 11), validation=stable_sort(validation, 13), test=stable_sort(test, 17), fixture_count=fixture_count, ) if any(record.value["slice"] == "vague-eval" for record in result.train): raise DataError("vague-eval leakage into training") if len(result.train) != target_train: raise DataError("train split missed its target size") if len(result.validation) != target_validation: raise DataError("validation split missed its target size") if result.logical_test_count != target_test: raise DataError("test split missed its target size") return result def distribution(records: Sequence[SourceRecord]) -> dict[str, dict[str, int]]: return { "purpose": dict( sorted(Counter(record.value["purpose"] for record in records).items()) ), "slice": dict( sorted(Counter(record.value["slice"] for record in records).items()) ), "language": dict( sorted( Counter( record.value["lang"].split("-", 1)[0].casefold() for record in records ).items() ) ), }