import contextlib import io import json import os from pathlib import Path import tempfile import time import unittest from unittest.mock import patch from checkpoint_archive import Archive, atomic_json, sha256_file, stable_stat import cleanup_run1_run40 as cleanup class CleanupGuardTest(unittest.TestCase): def setUp(self): self.tmp = tempfile.TemporaryDirectory() self.base = Path(self.tmp.name).resolve() self.root = self.base / 'checkpoints' self.migration = self.root / 'migration' self.migration.mkdir(parents=True) self.patch = patch.multiple(cleanup, BASE=self.base, ROOT=self.root, MIGRATION=self.migration) self.patch.start() self.source = self.base / 'source' self.source.mkdir() self.weight = self.source / 'svae.pt' self.weight.write_bytes(b'complete final model') old = time.time() - 7200 os.utime(self.weight, (old, old)) self.cache = self.source / 'features_rank0_part00000.pt' self.cache.write_bytes(b'regenerable feature cache') os.utime(self.cache, (old, old)) self.keep = self.source / 'config.yaml' self.keep.write_text('preserve: true') self.archive = Archive(self.root) # The one-time migration guard requires exactly 53 entries. for i in range(53): self.archive.add_model(key=f'svae/run{i:02d}', name=f'run{i:02d}_test', source=self.source, files=['svae.pt']) with contextlib.redirect_stderr(io.StringIO()): self.archive.verify(reconstruct=True, report=self.migration / 'verification.json', workers=4) atomic_json(self.migration / 'plan.json', {'fixture': True}) rows = [{'path': str(p), 'category': cat, 'archive_id': 'svae/run00', 'stat': stable_stat(p)} for p, cat in [(self.weight, 'archived_svae_original'), (self.cache, 'regenerable_feature_shard')]] atomic_json(self.migration / 'cleanup-plan.json', {'approved_plan_sha256': sha256_file(self.migration / 'plan.json'), 'files': rows, 'remove_if_empty': [], 'preserved': ['configuration']}) atomic_json(self.migration / 'load-checks.json', {'status': 'passed', 'eval_dry_run_exit_code': 0, 'cpu_load_successes': ['synthetic-fixture'], 'catalog_sha256': sha256_file(self.root / 'catalog.json')}) def tearDown(self): self.patch.stop() self.tmp.cleanup() def test_only_verified_manifest_files_are_deleted(self): with contextlib.redirect_stdout(io.StringIO()): cleanup.apply_plan() self.assertFalse(self.weight.exists()) self.assertFalse(self.cache.exists()) self.assertEqual(self.keep.read_text(), 'preserve: true') self.assertEqual((self.root / 'svae/run00_test/private/svae.pt').read_bytes(), b'complete final model') report = json.loads((self.migration / 'cleanup-report.json').read_text()) self.assertEqual(report['removed_files'], 2) def test_failed_validation_prevents_all_deletion(self): atomic_json(self.migration / 'load-checks.json', {'status': 'failed'}) with self.assertRaises(ValueError): cleanup.apply_plan() self.assertTrue(self.weight.exists()) self.assertTrue(self.cache.exists()) def test_changed_source_prevents_all_deletion(self): self.cache.write_bytes(b'new training in progress') with self.assertRaises(ValueError): cleanup.apply_plan() self.assertTrue(self.weight.exists()) self.assertTrue(self.cache.exists()) if __name__ == '__main__': unittest.main()