"""Dependency-free, content-addressed experiment manifests.""" from __future__ import annotations from collections.abc import Iterator, Mapping from dataclasses import dataclass from hashlib import sha256 import json import math import re from typing import Any EXACT_DEPENDENCY = re.compile( r"[A-Za-z0-9][A-Za-z0-9._-]*(?:\[[A-Za-z0-9_,.-]+\])?" r"==[A-Za-z0-9][A-Za-z0-9._+!-]*" ) @dataclass(frozen=True) class FrozenMapping(Mapping[str, Any]): """An immutable, key-sorted mapping whose values are recursively frozen.""" _items: tuple[tuple[str, Any], ...] def __post_init__(self) -> None: normalized: list[tuple[str, Any]] = [] for key, value in tuple(self._items): if not isinstance(key, str): raise TypeError("canonical JSON object keys must be strings") normalized.append((key, _freeze_json(value))) normalized.sort(key=lambda item: item[0]) keys = [key for key, _ in normalized] if len(keys) != len(set(keys)): raise ValueError("canonical JSON object keys must be unique") object.__setattr__(self, "_items", tuple(normalized)) def __getitem__(self, key: str) -> Any: for candidate, value in self._items: if candidate == key: return value raise KeyError(key) def __iter__(self) -> Iterator[str]: return (key for key, _ in self._items) def __len__(self) -> int: return len(self._items) def _freeze_json(value: Any) -> Any: if isinstance(value, FrozenMapping): return FrozenMapping(value._items) if isinstance(value, Mapping): if any(not isinstance(key, str) for key in value): raise TypeError("canonical JSON object keys must be strings") return FrozenMapping( tuple( (key, _freeze_json(child)) for key, child in sorted(value.items(), key=lambda item: item[0]) ) ) if isinstance(value, (list, tuple)): return tuple(_freeze_json(child) for child in value) if value is None or isinstance(value, (str, bool, int)): return value if isinstance(value, float): if not math.isfinite(value): raise ValueError("canonical JSON numbers must be finite") return value raise TypeError(f"value is not canonical JSON: {type(value).__name__}") def _json_value(value: Any) -> Any: if isinstance(value, FrozenMapping): return {key: _json_value(child) for key, child in value._items} if isinstance(value, tuple): return [_json_value(child) for child in value] return value def _freeze_mapping(value: Mapping[str, Any], name: str) -> FrozenMapping: frozen = _freeze_json(value) if not isinstance(frozen, FrozenMapping): raise TypeError(f"{name} must be a mapping") return frozen def digest_bytes(content: bytes) -> str: return sha256(content).hexdigest() def canonical_json(value: Any) -> str: return json.dumps( _json_value(_freeze_json(value)), allow_nan=False, ensure_ascii=False, separators=(",", ":"), sort_keys=True, ) @dataclass(frozen=True) class EvaluationContract: data_sha256: str code_sha256: str metrics: tuple[str, ...] slices: tuple[str, ...] def __post_init__(self) -> None: object.__setattr__(self, "metrics", tuple(sorted(self.metrics))) object.__setattr__(self, "slices", tuple(sorted(self.slices))) @dataclass(frozen=True) class RandomnessPolicy: generators: tuple[tuple[str, int], ...] worker_policy: str deterministic_algorithms: str def __post_init__(self) -> None: object.__setattr__( self, "generators", tuple(sorted(tuple(pair) for pair in self.generators)), ) @dataclass(frozen=True) class ExperimentManifest: experiment: str code_sha256: str data_sha256: str config: Mapping[str, Any] seed: int dependencies: tuple[str, ...] runtime: Mapping[str, str] hardware: Mapping[str, str] entry_command: tuple[str, ...] evaluation_contract: EvaluationContract randomness: RandomnessPolicy def __post_init__(self) -> None: object.__setattr__(self, "config", _freeze_mapping(self.config, "config")) object.__setattr__(self, "runtime", _freeze_mapping(self.runtime, "runtime")) object.__setattr__(self, "hardware", _freeze_mapping(self.hardware, "hardware")) object.__setattr__(self, "dependencies", tuple(sorted(self.dependencies))) object.__setattr__(self, "entry_command", tuple(self.entry_command)) validate_manifest(self) @property def run_id(self) -> str: evaluation = self.evaluation_contract randomness = self.randomness identity = { "experiment": self.experiment, "code_sha256": self.code_sha256, "data_sha256": self.data_sha256, "config": self.config, "seed": self.seed, "dependencies": self.dependencies, "runtime": self.runtime, "hardware": self.hardware, "entry_command": self.entry_command, "evaluation_contract": { "data_sha256": evaluation.data_sha256, "code_sha256": evaluation.code_sha256, "metrics": evaluation.metrics, "slices": evaluation.slices, }, "randomness": { "generators": randomness.generators, "worker_policy": randomness.worker_policy, "deterministic_algorithms": randomness.deterministic_algorithms, }, } return sha256(canonical_json(identity).encode("utf-8")).hexdigest()[:12] def _validate_digest(name: str, digest: str) -> None: if len(digest) != 64 or any( character not in "0123456789abcdef" for character in digest ): raise ValueError(f"{name} must be a lowercase SHA-256 digest") def validate_manifest(manifest: ExperimentManifest) -> None: if not manifest.experiment.strip(): raise ValueError("experiment must be explicit") _validate_digest("code_sha256", manifest.code_sha256) _validate_digest("data_sha256", manifest.data_sha256) if isinstance(manifest.seed, bool) or not isinstance(manifest.seed, int): raise ValueError("seed must be an integer") if not manifest.dependencies: raise ValueError("dependencies must be recorded") if len(manifest.dependencies) != len(set(manifest.dependencies)): raise ValueError("dependencies must be unique") for dependency in manifest.dependencies: if EXACT_DEPENDENCY.fullmatch(dependency) is None: raise ValueError(f"dependency is not exactly pinned: {dependency}") for group_name, values, required in ( ("runtime", manifest.runtime, {"python", "os", "architecture"}), ("hardware", manifest.hardware, {"accelerator", "precision"}), ): missing = required - values.keys() if missing: raise ValueError(f"{group_name} missing: {', '.join(sorted(missing))}") if any(not isinstance(value, str) or not value.strip() for value in values.values()): raise ValueError(f"{group_name} values must be explicit strings") if not manifest.entry_command or any( not isinstance(part, str) or not part.strip() for part in manifest.entry_command ): raise ValueError("entry_command must contain explicit arguments") evaluation = manifest.evaluation_contract if not isinstance(evaluation, EvaluationContract): raise ValueError("evaluation_contract must be explicit") _validate_digest("evaluation data_sha256", evaluation.data_sha256) _validate_digest("evaluation code_sha256", evaluation.code_sha256) for name, values in (("metrics", evaluation.metrics), ("slices", evaluation.slices)): if not values or len(values) != len(set(values)) or any( not isinstance(value, str) or not value.strip() for value in values ): raise ValueError(f"evaluation {name} must be unique and explicit") randomness = manifest.randomness if not isinstance(randomness, RandomnessPolicy): raise ValueError("randomness policy must be explicit") if not randomness.generators: raise ValueError("at least one RNG policy is required") generator_names: list[str] = [] for generator in randomness.generators: if len(generator) != 2: raise ValueError("RNG policy entries must be name/seed pairs") name, generator_seed = generator if not isinstance(name, str) or not name.strip(): raise ValueError("RNG names must be explicit") if isinstance(generator_seed, bool) or not isinstance(generator_seed, int): raise ValueError(f"RNG seed must be an integer: {name}") generator_names.append(name) if len(generator_names) != len(set(generator_names)): raise ValueError("RNG names must be unique") if not isinstance(randomness.worker_policy, str) or not randomness.worker_policy.strip(): raise ValueError("worker_policy must be explicit") if ( not isinstance(randomness.deterministic_algorithms, str) or not randomness.deterministic_algorithms.strip() ): raise ValueError("deterministic_algorithms policy must be explicit") def build_manifest( *, experiment: str, code: bytes, data: bytes, config: Mapping[str, Any], seed: int, dependencies: tuple[str, ...], runtime: Mapping[str, str], hardware: Mapping[str, str], entry_command: tuple[str, ...], evaluation_contract: EvaluationContract, randomness: RandomnessPolicy, ) -> ExperimentManifest: manifest = ExperimentManifest( experiment=experiment, code_sha256=digest_bytes(code), data_sha256=digest_bytes(data), config=config, seed=seed, dependencies=dependencies, runtime=runtime, hardware=hardware, entry_command=entry_command, evaluation_contract=evaluation_contract, randomness=randomness, ) validate_manifest(manifest) return manifest def verify_inputs(manifest: ExperimentManifest, *, code: bytes, data: bytes) -> bool: return ( manifest.code_sha256 == digest_bytes(code) and manifest.data_sha256 == digest_bytes(data) ) def require_inputs(manifest: ExperimentManifest, *, code: bytes, data: bytes) -> None: if not verify_inputs(manifest, code=code, data=data): raise ValueError("code or data does not match the experiment manifest") EXAMPLE_CODE = b"score = weight * feature\n" EXAMPLE_DATA = b"id,label\nr1,0\nr2,1\n" EXAMPLE_EVALUATION_CODE = b"metric = correct / total\n" EXAMPLE_EVALUATION_DATA = b"id,slice,label\ne1,all,1\n" def example_manifest() -> ExperimentManifest: return build_manifest( experiment="refund-intent-baseline", code=EXAMPLE_CODE, data=EXAMPLE_DATA, config={"epochs": 3, "learning_rate": 0.1}, seed=17, dependencies=("numpy==2.3.2",), runtime={"python": "3.12", "os": "linux", "architecture": "x86_64"}, hardware={"accelerator": "none", "precision": "float64"}, entry_command=("python3", "train.py", "--manifest", "manifest.json"), evaluation_contract=EvaluationContract( data_sha256=digest_bytes(EXAMPLE_EVALUATION_DATA), code_sha256=digest_bytes(EXAMPLE_EVALUATION_CODE), metrics=("accuracy",), slices=("all",), ), randomness=RandomnessPolicy( generators=(("numpy", 17), ("python", 17)), worker_policy="worker_seed=base_seed+worker_id", deterministic_algorithms="required-or-fail", ), ) if __name__ == "__main__": example = example_manifest() print(f"run_id={example.run_id}") print(f"code_sha256={example.code_sha256[:12]}") print(f"data_sha256={example.data_sha256[:12]}") print(f"inputs_verified={verify_inputs(example, code=EXAMPLE_CODE, data=EXAMPLE_DATA)}")