import copy import json from pathlib import Path import tempfile import unittest from data_transition import (SPLITS, PROTECTED_SPLITS, data_signature, sha256, verify_train_data_transition) class DataTransitionTests(unittest.TestCase): def setUp(self): self.directory = tempfile.TemporaryDirectory() self.addCleanup(self.directory.cleanup) root = Path(self.directory.name) self.old, self.new = root / "old", root / "new" self.provenance = {"revision": "pinned-model"} for path in (self.old, self.new): path.mkdir() for split in SPLITS: (path / f"{split}.jsonl").write_text('{"id":"original"}\n') (self.new / "train.jsonl").write_text('{"id":"original"}\n{"id":"new"}\n') self.parent = {"split_sha256": self.hashes(self.old)} (self.old / "manifest.json").write_text(json.dumps(self.parent)) self.current = {"split_sha256": self.hashes(self.new), "base_dataset": { "path": str(self.old), "manifest_sha256": sha256(self.old / "manifest.json"), "split_sha256": self.parent["split_sha256"]}} self.write_current() self.saved = {"config": {"dataset": str(self.old), "max_tokens": 512}, "model_provenance": self.provenance, "data_signature": data_signature(self.parent, self.provenance, 512)} self.config = {"dataset": str(self.new), "max_tokens": 512} def hashes(self, path): return {split: sha256(path / f"{split}.jsonl") for split in SPLITS} def write_current(self): (self.new / "manifest.json").write_text(json.dumps(self.current)) def verify(self): return verify_train_data_transition(self.saved, self.config, data_signature(self.current, self.provenance, 512)) def test_expansion_records_both_signatures_and_protected_hashes(self): untouched = copy.deepcopy(self.saved) result = self.verify() self.assertEqual(self.saved, untouched) self.assertEqual(result["parent_data_signature"], self.saved["data_signature"]) self.assertNotEqual(result["parent_data_signature"], result["data_signature"]) self.assertEqual(set(result["protected_split_sha256"]), set(PROTECTED_SPLITS)) def test_reserved_changes_rejected_even_with_valid_updated_hash(self): for split in PROTECTED_SPLITS: with self.subTest(split=split): original = (self.new / f"{split}.jsonl").read_bytes() (self.new / f"{split}.jsonl").write_text('changed') self.current["split_sha256"] = self.hashes(self.new) self.write_current() with self.assertRaisesRegex(ValueError, "reserved"): self.verify() (self.new / f"{split}.jsonl").write_bytes(original) def test_actual_bytes_missing_splits_and_false_lineage_rejected(self): (self.new / "train.jsonl").write_text('corrupted') with self.assertRaisesRegex(ValueError, "checksum"): self.verify() self.current["split_sha256"] = self.hashes(self.new) self.current["base_dataset"]["manifest_sha256"] = "wrong" self.write_current() with self.assertRaisesRegex(ValueError, "lineage"): self.verify() del self.current["split_sha256"]["holdout"] self.write_current() with self.assertRaisesRegex(ValueError, "five registered"): self.verify() def test_wrong_parent_signature_or_token_policy_rejected(self): self.saved["data_signature"] = "incorrect" with self.assertRaisesRegex(ValueError, "Parent data signature"): self.verify() self.config["max_tokens"] = 1024 with self.assertRaisesRegex(ValueError, "max_tokens"): self.verify() if __name__ == "__main__": unittest.main()