OpenWAM-3sys / tools /test_cleanup_guards.py
qiuly's picture
Add files using upload-large-folder tool
e7aa2d6 verified
Raw History Blame Contribute Delete
3.65 kB
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()