Files
werkjournal/tests/operations/test_operator_restore.py
2026-09-09 14:26:09 +02:00

185 lines
8.7 KiB
Python

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))