import contextlib import hashlib import io import json from pathlib import Path import runpy import shutil import subprocess import tempfile import unittest from unittest.mock import patch from test_backup import Controlled, backup ROOT = Path(__file__).resolve().parents[2] scope = runpy.run_path(str(ROOT / 'scripts/deploy-update-remote.py')) scope.update(runpy.run_path(str(ROOT / 'scripts/backup.py'))) remote = runpy.run_path(str(ROOT / 'scripts/restore-remote.py'), init_globals=scope) client = runpy.run_path(str(ROOT / 'scripts/operator-restore.py')) class Service(Controlled): fail_commit = None def preflight(self, check_data=True, allow_restore=False): if not allow_restore and (self.app / 'pending-restore.json').exists(): raise RuntimeError('Pending operator restore') def stop(self): self.calls.append('stop') self.running = False def ready(self, expected): self.calls.append('ready') if expected['commit'] == self.fail_commit: raise RuntimeError('Injected restored application failure') class OperatorRestoreTest(unittest.TestCase): def setUp(self): directory = tempfile.TemporaryDirectory() self.addCleanup(directory.cleanup) self.root = Path(directory.name) self.d = Service(self.root) self.d.database.write_bytes(b'old snapshot data') stream = io.BytesIO() backup(self.d, stream) self.archive = self.root / 'snapshot.tgz' self.archive.write_bytes(stream.getvalue()) self.snapshot_release = dict(self.d.expected) self.d.expected = dict(self.d.expected, commit='b' * 40) newer = self.d.app / 'releases' / ('b' * 40) shutil.copytree(self.d.app / 'current', newer) (newer / 'manifest.json').write_text(json.dumps(self.d.expected)) (self.d.app / 'current').unlink() (self.d.app / 'current').symlink_to(newer) (self.d.app / 'deployed.json').write_text(json.dumps(self.d.expected)) self.d.database.write_bytes(b'current newer records') self.d.calls.clear() self.operation = remote['OperatorRestore'](self.d) def test_confirmed_restore_keeps_prior_pair_and_installs_matching_snapshot(self): result = self.operation.restore(self.archive) self.assertEqual('RESTORED', result['result']) self.assertEqual(b'old snapshot data', self.d.database.read_bytes()) self.assertEqual('a' * 40, (self.d.app / 'current').resolve().name) self.assertEqual(self.snapshot_release, json.loads((self.d.app / 'deployed.json').read_text())) prior = Path(result['priorSnapshot']) self.assertEqual(b'current newer records', (prior / 'data/werkjournal.mv.db').read_bytes()) self.assertEqual(self.d.expected, remote['verify_tree'](prior)['release']) self.assertTrue(self.d.running) self.assertFalse(self.d.flag.exists()) self.assertFalse(self.operation.pending.exists()) def test_failed_restored_application_rolls_back_and_still_reports_failure(self): self.d.fail_commit = 'a' * 40 with self.assertRaisesRegex(RuntimeError, 'Injected'): self.operation.restore(self.archive) self.assertEqual(b'current newer records', self.d.database.read_bytes()) self.assertEqual('b' * 40, (self.d.app / 'current').resolve().name) self.assertFalse(self.operation.pending.exists()) self.assertFalse(self.d.flag.exists()) self.assertTrue(self.d.running) def interrupt_at(self, phase): record = self.operation.record def interrupted(state, value): record(state, value) if value == phase: raise KeyboardInterrupt('simulated process death') with patch.object(self.operation, 'record', side_effect=interrupted): with self.assertRaises(KeyboardInterrupt): self.operation.restore(self.archive) def test_interrupted_install_recovers_previous_database_and_application(self): self.interrupt_at('installed') self.assertTrue(self.d.flag.exists()) self.assertEqual(b'old snapshot data', self.d.database.read_bytes()) self.operation.recover() self.assertEqual(b'current newer records', self.d.database.read_bytes()) self.assertEqual('b' * 40, (self.d.app / 'current').resolve().name) self.assertFalse(self.operation.pending.exists()) def test_interruption_with_missing_live_data_directory_recovers_prior_copy(self): def interrupted(source, bundle): (self.d.app / 'data').rename(bundle / 'interrupted-data') raise KeyboardInterrupt('between directory renames') with patch.object(self.operation, 'install_data', side_effect=interrupted): with self.assertRaises(KeyboardInterrupt): self.operation.restore(self.archive) self.assertFalse((self.d.app / 'data').exists()) self.operation.recover() self.assertEqual(b'current newer records', self.d.database.read_bytes()) def test_failed_rollback_remains_closed_and_can_be_recovered_later(self): with patch.object(self.d, 'ready', side_effect=RuntimeError('cannot start')), contextlib.redirect_stderr(io.StringIO()): with self.assertRaisesRegex(RuntimeError, 'cannot start'): self.operation.restore(self.archive) self.assertTrue(self.operation.pending.exists()) self.assertTrue(self.d.flag.exists()) self.assertFalse(self.d.running) self.operation.recover() self.assertEqual(b'current newer records', self.d.database.read_bytes()) self.assertFalse(self.operation.pending.exists()) def test_recovery_after_publication_preserves_newly_recorded_data(self): def interrupt(state): self.d.flag.unlink() raise KeyboardInterrupt('after publication') with patch.object(self.operation, 'finish', side_effect=interrupt): with self.assertRaises(KeyboardInterrupt): self.operation.restore(self.archive) self.d.database.write_bytes(b'records accepted after publication') self.operation.recover() self.assertEqual(b'records accepted after publication', self.d.database.read_bytes()) self.assertEqual('a' * 40, (self.d.app / 'current').resolve().name) self.assertFalse(self.operation.pending.exists()) def test_corrupt_prior_snapshot_keeps_maintenance_instead_of_restoring_corruption(self): self.interrupt_at('installed') state = json.loads(self.operation.pending.read_text()) prior = self.operation.bundle(state) / 'prior/data/werkjournal.mv.db' prior.write_bytes(b'corrupted') with self.assertRaisesRegex(ValueError, 'checksum'): self.operation.recover() self.assertTrue(self.d.flag.exists()) self.assertTrue(self.operation.pending.exists()) self.assertFalse(self.d.running) def test_stopped_service_and_operator_maintenance_remain_after_success(self): self.d.running = False self.d.flag.write_text('operator maintenance') self.operation.restore(self.archive) self.assertFalse(self.d.running) self.assertEqual('operator maintenance', self.d.flag.read_text()) def test_pending_restore_blocks_other_deployment_operations(self): self.operation.pending.write_text('{}') deployment = scope['Deployment'](self.root) with self.assertRaisesRegex(RuntimeError, 'Pending operator restore'): deployment.preflight() def test_cancel_and_invalid_archive_do_not_open_ssh(self): with patch.object(subprocess, 'call') as ssh, contextlib.redirect_stdout(io.StringIO()): with self.assertRaisesRegex(RuntimeError, 'cancelled'): client['restore']('fixture@host', self.archive, confirm=lambda _: 'no') ssh.assert_not_called() self.archive.write_bytes(b'not a backup') with patch.object(subprocess, 'call') as ssh: with self.assertRaises(Exception): client['restore']('fixture@host', self.archive, confirm=lambda _: 'ignored') ssh.assert_not_called() def test_confirmation_transfers_frozen_verified_bytes(self): original = self.archive.read_bytes() def confirm(prompt): self.archive.write_bytes(b'changed original during confirmation') return 'RESTORE ' + 'a' * 40 def ssh(command, stdin): self.assertEqual(original, stdin.read()) self.assertEqual(['ssh', '-o', 'BatchMode=yes', 'fixture@host'], command[:4]) return 0 with patch.object(subprocess, 'call', side_effect=ssh), contextlib.redirect_stdout(io.StringIO()): self.assertEqual(0, client['restore']('fixture@host', self.archive, confirm=confirm))