"""Validate a train/validation/test manifest against causal availability rules. The validator treats a split as deployment simulation, not a random partition. It binds every row to a content-addressed contract and manifest, checks when features and full-horizon labels could exist, and enforces entity and group isolation before reporting a manifest as usable. """ from __future__ import annotations from dataclasses import asdict, dataclass import hashlib import json import re from typing import Any _IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/@-]{0,127}$") _SPLITS = ("train", "validation", "test") _MAX_ALLOWED_RECORDS = 100_000 _FEATURE_RULE = "feature-as-of-not-after-decision-v1" _LABEL_RULE = "full-horizon-label-maturity-v1" def _require_identifier(name: str, value: str) -> None: if not isinstance(value, str) or not _IDENTIFIER.fullmatch(value): raise ValueError(f"{name} must be a non-empty stable identifier") def _require_owner(name: str, value: str) -> None: _require_identifier(name, value) if not value.startswith("team:"): raise ValueError(f"{name} must name an accountable team: owner") def _require_int(name: str, value: int, *, minimum: int = 0) -> None: if isinstance(value, bool) or not isinstance(value, int) or value < minimum: raise ValueError(f"{name} must be an integer >= {minimum}") def _canonical(value: Any) -> Any: if isinstance(value, dict): return {key: _canonical(item) for key, item in value.items()} if isinstance(value, tuple): return [_canonical(item) for item in value] return value def _content_id(prefix: str, payload: dict[str, Any]) -> str: encoded = json.dumps( _canonical(payload), sort_keys=True, separators=(",", ":"), allow_nan=False ).encode("utf-8") return f"{prefix}@sha256:{hashlib.sha256(encoded).hexdigest()}" @dataclass(frozen=True) class SplitContract: contract_version: str task_id: str target_id: str population_id: str dataset_revision: str provenance_snapshot_id: str feature_cutoff_rule: str label_availability_rule: str train_decision_start_day: int train_decision_end_day: int validation_decision_start_day: int validation_decision_end_day: int test_decision_start_day: int test_decision_end_day: int test_evaluation_cutoff_day: int prediction_horizon_days: int embargo_days: int isolate_entities: bool isolate_groups: bool minimum_records_per_split: int max_records: int data_owner: str evaluation_owner: str def __post_init__(self) -> None: for name in ( "contract_version", "task_id", "target_id", "population_id", "dataset_revision", "provenance_snapshot_id", ): _require_identifier(name, getattr(self, name)) if self.feature_cutoff_rule != _FEATURE_RULE: raise ValueError(f"feature_cutoff_rule must be {_FEATURE_RULE}") if self.label_availability_rule != _LABEL_RULE: raise ValueError(f"label_availability_rule must be {_LABEL_RULE}") for name in ( "train_decision_start_day", "train_decision_end_day", "validation_decision_start_day", "validation_decision_end_day", "test_decision_start_day", "test_decision_end_day", "test_evaluation_cutoff_day", ): _require_int(name, getattr(self, name)) _require_int("prediction_horizon_days", self.prediction_horizon_days, minimum=1) _require_int("embargo_days", self.embargo_days) _require_int("minimum_records_per_split", self.minimum_records_per_split, minimum=1) _require_int("max_records", self.max_records, minimum=3) if self.max_records > _MAX_ALLOWED_RECORDS: raise ValueError(f"max_records cannot exceed {_MAX_ALLOWED_RECORDS}") if self.minimum_records_per_split * len(_SPLITS) > self.max_records: raise ValueError("minimum split coverage cannot exceed max_records") if type(self.isolate_entities) is not bool or not self.isolate_entities: raise ValueError("isolate_entities must be the literal True") if type(self.isolate_groups) is not bool or not self.isolate_groups: raise ValueError("isolate_groups must be the literal True") if not self.train_decision_start_day <= self.train_decision_end_day: raise ValueError("train decision window is invalid") if not self.train_decision_end_day < self.validation_decision_start_day: raise ValueError("validation must begin after the train decision window") if not self.validation_decision_start_day <= self.validation_decision_end_day: raise ValueError("validation decision window is invalid") if not self.validation_decision_end_day < self.test_decision_start_day: raise ValueError("test must begin after the validation decision window") if not self.test_decision_start_day <= self.test_decision_end_day: raise ValueError("test decision window is invalid") if ( self.train_decision_start_day + self.prediction_horizon_days + self.embargo_days > self.validation_decision_start_day ): raise ValueError("train window cannot produce a mature label before validation") if ( self.validation_decision_start_day + self.prediction_horizon_days + self.embargo_days > self.test_decision_start_day ): raise ValueError("validation window cannot produce a mature label before test") if self.test_evaluation_cutoff_day < self.test_decision_end_day + self.prediction_horizon_days: raise ValueError("test_evaluation_cutoff_day must mature the final test horizon") _require_owner("data_owner", self.data_owner) _require_owner("evaluation_owner", self.evaluation_owner) @property def contract_id(self) -> str: return _content_id(self.contract_version, asdict(self)) @dataclass(frozen=True) class SplitRecord: record_id: str entity_id: str group_id: str split: str task_id: str target_id: str population_id: str source_revision: str feature_as_of_day: int decision_day: int outcome: int event_day: int | None label_as_of_day: int @dataclass(frozen=True) class SplitManifest: contract_id: str manifest_id: str task_id: str target_id: str population_id: str dataset_revision: str provenance_snapshot_id: str records: tuple[SplitRecord, ...] @property def evidence_id(self) -> str: return _content_id("causal-split-evidence-v1", asdict(self)) @dataclass(frozen=True) class SplitAudit: contract_id: str manifest_id: str evidence_id: str counts: tuple[tuple[str, int], ...] minimum_boundary_gap_days: int group_overlap_count: int entity_overlap_count: int decision: str def _validate_label(contract: SplitContract, record: SplitRecord, index: int) -> None: _require_int(f"records[{index}].label_as_of_day", record.label_as_of_day) maturity_day = record.decision_day + contract.prediction_horizon_days if record.label_as_of_day < maturity_day: raise ValueError(f"records[{index}] label is immature for the full prediction horizon") if isinstance(record.outcome, bool) or type(record.outcome) is not int: raise TypeError(f"records[{index}].outcome must be the integer 0 or 1") if record.outcome not in (0, 1): raise ValueError(f"records[{index}].outcome must be 0 or 1") if record.outcome == 0: if record.event_day is not None: raise ValueError(f"records[{index}] negative label cannot carry an event_day") else: if record.event_day is None: raise ValueError(f"records[{index}] positive label requires event_day") _require_int(f"records[{index}].event_day", record.event_day) if not record.decision_day < record.event_day <= maturity_day: raise ValueError(f"records[{index}] event_day is outside the prediction horizon") if record.event_day > record.label_as_of_day: raise ValueError(f"records[{index}] event_day is after label_as_of_day") def _validate_time(contract: SplitContract, record: SplitRecord, index: int) -> None: _require_int(f"records[{index}].feature_as_of_day", record.feature_as_of_day) _require_int(f"records[{index}].decision_day", record.decision_day) if record.feature_as_of_day > record.decision_day: raise ValueError(f"records[{index}] uses a feature unavailable at decision time") _validate_label(contract, record, index) if record.split == "train": if not contract.train_decision_start_day <= record.decision_day <= contract.train_decision_end_day: raise ValueError(f"records[{index}] is outside the train decision window") if record.label_as_of_day + contract.embargo_days > contract.validation_decision_start_day: raise ValueError(f"records[{index}] train label is unavailable before validation") elif record.split == "validation": if not contract.validation_decision_start_day <= record.decision_day <= contract.validation_decision_end_day: raise ValueError(f"records[{index}] is outside the validation decision window") if record.label_as_of_day + contract.embargo_days > contract.test_decision_start_day: raise ValueError(f"records[{index}] validation label is unavailable before test") elif record.split == "test": if not contract.test_decision_start_day <= record.decision_day <= contract.test_decision_end_day: raise ValueError(f"records[{index}] is outside the test decision window") if record.label_as_of_day > contract.test_evaluation_cutoff_day: raise ValueError(f"records[{index}] test label is after the evaluation cutoff") else: raise ValueError(f"records[{index}].split must be train, validation, or test") def validate_causal_split(contract: SplitContract, manifest: SplitManifest) -> SplitAudit: """Fail closed on scope, provenance, availability, and partition overlap.""" if type(contract) is not SplitContract: raise TypeError("contract must be a concrete SplitContract") SplitContract.__post_init__(contract) if type(manifest) is not SplitManifest: raise TypeError("manifest must be a concrete SplitManifest") if manifest.contract_id != contract.contract_id: raise ValueError("manifest contract_id does not match the full contract identity") _require_identifier("manifest_id", manifest.manifest_id) for field in ("task_id", "target_id", "population_id"): if getattr(manifest, field) != getattr(contract, field): raise ValueError(f"manifest {field} does not match contract") if manifest.dataset_revision != contract.dataset_revision: raise ValueError("manifest dataset_revision does not match contract") if manifest.provenance_snapshot_id != contract.provenance_snapshot_id: raise ValueError("manifest provenance_snapshot_id does not match contract") if not isinstance(manifest.records, tuple): raise TypeError("records must be an immutable tuple") if len(manifest.records) > contract.max_records: raise ValueError("record count exceeds the contract's bounded workload") counts = {split: 0 for split in _SPLITS} record_ids: set[str] = set() entity_splits: dict[str, str] = {} group_splits: dict[str, str] = {} for index, record in enumerate(manifest.records): if type(record) is not SplitRecord: raise TypeError(f"records[{index}] must be a concrete SplitRecord") _require_identifier(f"records[{index}].record_id", record.record_id) _require_identifier(f"records[{index}].entity_id", record.entity_id) _require_identifier(f"records[{index}].group_id", record.group_id) if record.record_id in record_ids: raise ValueError("record_id values must be unique") record_ids.add(record.record_id) for field in ("task_id", "target_id", "population_id"): if getattr(record, field) != getattr(contract, field): raise ValueError(f"records[{index}].{field} does not match contract") if record.source_revision != contract.dataset_revision: raise ValueError(f"records[{index}] source_revision does not match contract") _validate_time(contract, record, index) counts[record.split] += 1 prior_entity_split = entity_splits.setdefault(record.entity_id, record.split) if prior_entity_split != record.split: raise ValueError("entity_id appears in more than one split") prior_group_split = group_splits.setdefault(record.group_id, record.split) if prior_group_split != record.split: raise ValueError("group_id appears in more than one split") for split, count in counts.items(): if count < contract.minimum_records_per_split: raise ValueError(f"{split} is below minimum_records_per_split") minimum_gap = min( contract.validation_decision_start_day - contract.train_decision_end_day, contract.test_decision_start_day - contract.validation_decision_end_day, ) return SplitAudit( contract_id=contract.contract_id, manifest_id=manifest.manifest_id, evidence_id=manifest.evidence_id, counts=tuple((split, counts[split]) for split in _SPLITS), minimum_boundary_gap_days=minimum_gap, group_overlap_count=0, entity_overlap_count=0, decision="PASS", ) ILLUSTRATIVE_CONTRACT = SplitContract( contract_version="causal-split-v1", task_id="illustrative-renewal-forecast", target_id="cancellation-within-30-days", population_id="active-consumer-accounts", dataset_revision="renewal-events-2026-08-20", provenance_snapshot_id="renewal-lineage-snapshot-0042", feature_cutoff_rule=_FEATURE_RULE, label_availability_rule=_LABEL_RULE, train_decision_start_day=10, train_decision_end_day=60, validation_decision_start_day=100, validation_decision_end_day=110, test_decision_start_day=150, test_decision_end_day=160, test_evaluation_cutoff_day=190, prediction_horizon_days=30, embargo_days=2, isolate_entities=True, isolate_groups=True, minimum_records_per_split=2, max_records=1000, data_owner="team:account-data", evaluation_owner="team:retention-modeling", ) def _row( record_id: str, entity_id: str, group_id: str, split: str, feature_day: int, decision_day: int, outcome: int, event_day: int | None, label_day: int, ) -> SplitRecord: return SplitRecord( record_id, entity_id, group_id, split, ILLUSTRATIVE_CONTRACT.task_id, ILLUSTRATIVE_CONTRACT.target_id, ILLUSTRATIVE_CONTRACT.population_id, ILLUSTRATIVE_CONTRACT.dataset_revision, feature_day, decision_day, outcome, event_day, label_day, ) ILLUSTRATIVE_MANIFEST = SplitManifest( contract_id=ILLUSTRATIVE_CONTRACT.contract_id, manifest_id="renewal-causal-manifest-0042", task_id=ILLUSTRATIVE_CONTRACT.task_id, target_id=ILLUSTRATIVE_CONTRACT.target_id, population_id=ILLUSTRATIVE_CONTRACT.population_id, dataset_revision=ILLUSTRATIVE_CONTRACT.dataset_revision, provenance_snapshot_id=ILLUSTRATIVE_CONTRACT.provenance_snapshot_id, records=( _row("row-1", "account-1", "region-a", "train", 19, 20, 0, None, 60), _row("row-2", "account-2", "region-a", "train", 29, 30, 1, 45, 60), _row("row-3", "account-3", "region-b", "validation", 99, 100, 0, None, 140), _row("row-4", "account-4", "region-b", "validation", 104, 105, 1, 120, 140), _row("row-5", "account-5", "region-c", "test", 149, 150, 0, None, 190), _row("row-6", "account-6", "region-c", "test", 154, 155, 1, 170, 190), ), ) def format_example() -> str: audit = validate_causal_split(ILLUSTRATIVE_CONTRACT, ILLUSTRATIVE_MANIFEST) counts = ",".join(f"{name}:{count}" for name, count in audit.counts) return "\n".join( ( "example=illustrative_only", f"contract_version={ILLUSTRATIVE_CONTRACT.contract_version}", f"manifest={audit.manifest_id}", f"evidence_id={audit.evidence_id}", f"counts={counts}", f"minimum_boundary_gap_days={audit.minimum_boundary_gap_days}", f"group_overlap={audit.group_overlap_count}", f"entity_overlap={audit.entity_overlap_count}", f"decision={audit.decision}", ) ) if __name__ == "__main__": print(format_example())