1098 lines
55 KiB
Python
1098 lines
55 KiB
Python
"""Synthetic tests only: python -I -S -B docker/test_windows_snapshot.py -v."""
|
|
|
|
import contextlib
|
|
import ctypes
|
|
from dataclasses import replace
|
|
import hashlib
|
|
import importlib.util
|
|
import io
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
import stat
|
|
import subprocess
|
|
import sys
|
|
import tarfile
|
|
import tempfile
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
from unittest.mock import Mock, patch
|
|
|
|
|
|
spec = importlib.util.spec_from_file_location('windows_snapshot', Path(__file__).with_name('windows_snapshot.py'))
|
|
snapshot = importlib.util.module_from_spec(spec)
|
|
sys.modules[spec.name] = snapshot
|
|
spec.loader.exec_module(snapshot)
|
|
REAL_SECURE_PATH = snapshot._secure_path
|
|
REAL_LOAD_SOURCE = snapshot._load_source
|
|
|
|
SECRET = 'synthetic-password-do-not-print'
|
|
|
|
|
|
class Fixture(unittest.TestCase):
|
|
def setUp(self):
|
|
temporary_root = Path(tempfile.gettempdir()) / 'opencode'
|
|
self.temp = tempfile.TemporaryDirectory(prefix='truf-snapshot-test-',
|
|
dir=temporary_root if temporary_root.is_dir() else None)
|
|
self.addCleanup(self.temp.cleanup)
|
|
self.base = Path(self.temp.name)
|
|
self.root, self.bundles = self.base / 'source', self.base / 'bundles'
|
|
self.imports = self.base / 'clone/docker/imports'
|
|
for path in (self.root / 'app', self.root / 'runtime', self.bundles, self.imports):
|
|
path.mkdir(parents=True)
|
|
self.output = self.imports / 'unique'
|
|
for name, value in (('SOURCE_ROOT', self.root), ('BUNDLE_ROOT', self.bundles),
|
|
('POSTGRES_DATA', self.base / 'pgdata'), ('IMPORTS_ROOT', self.imports)):
|
|
patcher = patch.object(snapshot, name, value)
|
|
patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
# An accidental integration call fails before reaching any original code
|
|
# or real child process. Individual tests substitute synthetic processes.
|
|
for name in ('_load_source',):
|
|
patcher = patch.object(snapshot, name, side_effect=AssertionError('source access forbidden'))
|
|
patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
for name in ('Popen', 'run'):
|
|
patcher = patch.object(snapshot.subprocess, name, side_effect=AssertionError('process forbidden'))
|
|
patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
patcher = patch.object(snapshot, '_secure_path')
|
|
self.secure = patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
self.report = Mock()
|
|
self.handlers = {number: snapshot.signal.default_int_handler
|
|
for number in (snapshot.signal.SIGINT, getattr(snapshot.signal, 'SIGBREAK', None))
|
|
if number is not None}
|
|
self.previous_handlers = self.handlers.copy()
|
|
|
|
def register(number, handler):
|
|
previous = self.handlers[number]
|
|
self.handlers[number] = handler
|
|
return previous
|
|
|
|
patcher = patch.object(snapshot.signal, 'signal', side_effect=register)
|
|
self.signal = patcher.start()
|
|
self.addCleanup(patcher.stop)
|
|
self.addCleanup(lambda: self.assertEqual(self.handlers, self.previous_handlers))
|
|
|
|
def send_signal(self, number):
|
|
# Invoke only this fixture's registered Python handler, not an OS signal.
|
|
self.handlers[number](number, None)
|
|
|
|
def put(self, relative, content=b'fixture', bundles=False):
|
|
path = (self.bundles if bundles else self.root) / relative
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
path.write_bytes(content)
|
|
return path
|
|
|
|
def runtime(self, state='STOPPED'):
|
|
calls = []
|
|
authority = SimpleNamespace(acquired=False)
|
|
|
|
def acquire():
|
|
authority.acquired = True
|
|
calls.append('acquire')
|
|
|
|
def release():
|
|
calls.append('release')
|
|
authority.acquired = False
|
|
|
|
authority.acquire, authority.release = Mock(side_effect=acquire), Mock(side_effect=release)
|
|
backend = SimpleNamespace(state=state, close=Mock())
|
|
backend.probe = Mock(side_effect=lambda: SimpleNamespace(kind=backend.state, detail=SECRET))
|
|
identity = {'pg_major': 16, 'system_identifier': '1234567890123456789',
|
|
'data_directory': str(snapshot.POSTGRES_DATA), 'database': 'fixture', 'user': 'fixture', 'port': 5432}
|
|
|
|
def start(config, backend):
|
|
self.assertTrue(authority.acquired)
|
|
calls.append('start')
|
|
backend.state = 'READY'
|
|
print(SECRET)
|
|
return SimpleNamespace(kind='READY')
|
|
|
|
def stop(config, backend):
|
|
self.assertTrue(authority.acquired)
|
|
calls.append('stop')
|
|
backend.state = 'STOPPED'
|
|
print(SECRET, file=sys.stderr)
|
|
return SimpleNamespace(completed=True, stopped=True)
|
|
|
|
pg = SimpleNamespace(
|
|
ProbeKind=SimpleNamespace(STOPPED='STOPPED', READY='READY'),
|
|
configured_cluster_values=Mock(return_value={'database': 'fixture', 'user': 'fixture', 'port': 5432}),
|
|
verify_cluster_identity=Mock(return_value=identity), PostgresBackend=Mock(return_value=backend),
|
|
maintenance_start=Mock(side_effect=start), maintenance_stop=Mock(side_effect=stop))
|
|
source = SimpleNamespace(pg=pg, security=SimpleNamespace(ClusterAuthorityLock=Mock(return_value=authority)),
|
|
config={}, dsn='postgresql://fixture:' + SECRET + '@127.0.0.1:5432/fixture', inputs={})
|
|
self.calls, self.authority, self.backend, self.identity, self.source = calls, authority, backend, identity, source
|
|
return source
|
|
|
|
@contextlib.contextmanager
|
|
def capture_context(self, database=None):
|
|
source = self.runtime()
|
|
|
|
def dump(*args):
|
|
self.assertTrue(self.authority.acquired)
|
|
self.assertEqual(self.backend.state, 'READY')
|
|
self.calls.append('database')
|
|
self.output.joinpath('database.dump').write_bytes(b'PGDMPfixture')
|
|
return {'version_num': 160004, 'system_identifier': self.identity['system_identifier'],
|
|
'database_name': 'fixture', 'table_counts': {'fixture': 2}, 'bytes': 12, 'sha256': '0' * 64,
|
|
'sequence_states': {}, 'sequence_count': 0}
|
|
|
|
with patch.object(snapshot, '_load_source', return_value=source), \
|
|
patch.object(snapshot, '_capture_database', side_effect=database or dump), \
|
|
patch.object(snapshot, '_verify_private_acl'), \
|
|
contextlib.redirect_stdout(io.StringIO()), contextlib.redirect_stderr(io.StringIO()):
|
|
yield source
|
|
|
|
|
|
class SelectionTests(Fixture):
|
|
def test_active_mappings_and_projection_sidecars(self):
|
|
expected = {}
|
|
for folder in snapshot.ACTIVE_DIRS:
|
|
name = 'runtime/' + folder + '/data.jsonl'
|
|
self.put(name)
|
|
expected['runtime-linux/' + folder + '/data.jsonl'] = self.root / name
|
|
for name in ('secrets.yaml', 'trufflehog-custom-detectors.yaml'):
|
|
self.put('app/' + name)
|
|
expected['config/' + name] = self.root / 'app' / name
|
|
self.put('runtime/proxy.txt', b'not-a-real-proxy')
|
|
expected['runtime-linux/proxy.txt'] = self.root / 'runtime/proxy.txt'
|
|
for folder in ('results', 'keychecks'):
|
|
for name in ('segment-00001.jsonl', 'scan_errors.log', 'scan_errors.log.22', 'stream.log.3',
|
|
'data.publication-ledger.sqlite3', 'data.publication-ledger.sqlite3-wal',
|
|
'data.publication-ledger.sqlite3-shm', 'data.sqlite-journal'):
|
|
relative = 'runtime/' + folder + '/' + name
|
|
self.put(relative)
|
|
expected['runtime-linux/' + folder + '/' + name] = self.root / relative
|
|
self.put('ready/id/segment.jsonl', bundles=True)
|
|
expected['scanner-result-bundles/ready/id/segment.jsonl'] = self.bundles / 'ready/id/segment.jsonl'
|
|
self.assertEqual({k: v.source for k, v in snapshot._inventory().files.items()}, expected)
|
|
|
|
def test_all_reviewed_archive_mappings(self):
|
|
names = (
|
|
'app/config.yaml', 'app/config.yaml.old', 'app/secrets.yaml.backup', 'app/.streamlit/config.toml',
|
|
'app/secrets.yaml.backup.log', 'scan_errors.log.1', 'scan_results.jsonl.segments/000001',
|
|
'found_secrets.jsonl.publication-ledger.sqlite3-wal',
|
|
'.env.postgres', 'docker-compose.postgres.yml', 'found_secrets.jsonl', 'scan_results.jsonl',
|
|
'scan_errors.log', 'checked_provider.txt', 'todo_provider.txt', 'state/legacy.json',
|
|
'runner_state.json', 'scanner.db', 'scanner.db-wal', 'app/scanner.db-shm',
|
|
'runtime/backups/old/scanner.db', 'runtime/' + snapshot.KEYCHECK_COPY + '/stream.log.1',
|
|
'runtime/keychecks.7z', 'runtime/orkey.txt', 'runtime/imports/old.dump', 'runtime/notes.md',
|
|
'runtime/check-openrouter-keys.ps1', 'runtime/control/migration-report.json',
|
|
)
|
|
for name in names:
|
|
self.put(name)
|
|
self.assertEqual(set(snapshot._inventory().files), {'windows-archive/' + name for name in names})
|
|
|
|
def test_state_scratch_archival_not_active(self):
|
|
names = ('scan_limiter.db', 'scan_limiter-copy.db-wal', 'scan_limiter.db-shm',
|
|
'nested/file.tmp.123', 'janitor.cursor.json', 'tmpdir.tmp/retained.json')
|
|
for name in names:
|
|
self.put('runtime/state/' + name)
|
|
self.put('runtime/state/useful.json')
|
|
self.assertEqual(set(snapshot._inventory().files),
|
|
{'windows-archive/runtime/state/' + name for name in names}
|
|
| {'runtime-linux/state/useful.json'})
|
|
|
|
def test_excludes_unreviewed_data_logs_locks_and_code_caches(self):
|
|
for name in (
|
|
'runtime/state/gharchive_cache/file.json.gz', 'runtime/state/debug.log.4',
|
|
'runtime/state/writer.lock.old', 'runtime/state/process.pid', 'runtime/queues/__pycache__/x.pyc',
|
|
'runtime/backups/ordinary.log.9', 'runtime/results/file.lock', 'runtime/results/__pycache__/x.pyc',
|
|
'runtime/downloads/archive.gz', 'runtime/git/repo/file', 'runtime/traces/t.json',
|
|
'runtime/freeze-diagnostics/report.json', 'runtime/postgres/data/physical',
|
|
'runtime/postgres/pgsql/bin/pg_dump.exe', 'app/scanner.py', 'app/config.yaml.lock',
|
|
'.opencode/private', 'tests/test.py',
|
|
):
|
|
self.put(name)
|
|
self.assertEqual(snapshot._inventory().files, {})
|
|
|
|
def test_control_never_includes_authority_or_non_reports(self):
|
|
for name in ('supervisor.instance.json', 'supervisor-report.json', 'client-token.json',
|
|
'authority.json', 'manifest.json', 'lock.json', 'worker.pid', 'data.txt', 'report.json'):
|
|
self.put('runtime/control/' + name)
|
|
self.assertEqual(set(snapshot._inventory().files), {'windows-archive/runtime/control/report.json'})
|
|
|
|
def test_special_scan_errors_and_ledger_names_outside_results_survive(self):
|
|
for name in ('scan_errors.log.8', 'name.log.publication-ledger.sqlite3-wal'):
|
|
self.put('runtime/backups/' + name)
|
|
self.assertEqual(len(snapshot._inventory().files), 2)
|
|
|
|
def test_selected_symlink_rejected(self):
|
|
link = self.put('runtime/proxy.txt')
|
|
original = Path.lstat
|
|
|
|
def lstat(path):
|
|
if path == link:
|
|
return SimpleNamespace(st_mode=stat.S_IFLNK, st_file_attributes=0, st_nlink=1)
|
|
return original(path)
|
|
|
|
with patch.object(Path, 'lstat', lstat), self.assertRaises(snapshot.Failure):
|
|
snapshot._inventory()
|
|
|
|
def test_selected_hardlink_rejected(self):
|
|
target = self.put('ordinary.txt')
|
|
os.link(target, self.root / 'runtime/proxy.txt')
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._inventory()
|
|
|
|
def test_reparse_directory_is_not_traversed(self):
|
|
child = self.put('runtime/results/junction/never-opened')
|
|
junction = child.parent
|
|
original = Path.lstat
|
|
|
|
def lstat(path):
|
|
if path == junction:
|
|
return SimpleNamespace(st_mode=stat.S_IFDIR, st_file_attributes=0x400, st_nlink=1)
|
|
self.assertNotEqual(path, child)
|
|
return original(path)
|
|
|
|
with patch.object(Path, 'lstat', lstat), patch.object(snapshot.os, 'scandir', wraps=os.scandir) as scan:
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._inventory()
|
|
self.assertNotIn(junction, [call.args[0] for call in scan.call_args_list])
|
|
|
|
def test_special_file_rejected(self):
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._check_type(SimpleNamespace(st_mode=stat.S_IFIFO, st_file_attributes=0, st_nlink=1))
|
|
|
|
def test_unsafe_tar_destinations(self):
|
|
for name in ('/absolute', '../up', 'a/../b', 'a//b', 'a/./b', 'a\\b', 'D:/file', 'a\x00b'):
|
|
with self.subTest(name=repr(name)), self.assertRaises(snapshot.Failure):
|
|
snapshot._destination(name, set(), set())
|
|
|
|
def test_casefold_duplicate_and_file_directory_conflicts(self):
|
|
for first, second in (('runtime-linux/X', 'runtime-linux/x'),
|
|
('runtime-linux/x', 'runtime-linux/X/child'),
|
|
('runtime-linux/X/child', 'runtime-linux/x'),
|
|
('runtime-linux/strasse', 'runtime-linux/stra\u00dfe')):
|
|
used, parents = set(), set()
|
|
snapshot._destination(first, used, parents)
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._destination(second, used, parents)
|
|
|
|
|
|
class ArchiveTests(Fixture):
|
|
def test_stable_large_file_path_and_handle_fingerprints_match(self):
|
|
for index in range(20):
|
|
path = self.put('runtime/results/large-' + str(index), b'x' * (snapshot.BLOCK + 7))
|
|
expected = snapshot._fingerprint(snapshot._file_info(path))
|
|
with path.open('rb', buffering=0) as handle:
|
|
actual = snapshot._fingerprint(os.fstat(handle.fileno()))
|
|
self.assertEqual(expected[:6] + expected[7:], actual[:6] + actual[7:])
|
|
handle.read(snapshot.BLOCK)
|
|
self.assertEqual(snapshot._fingerprint(os.fstat(handle.fileno())), actual)
|
|
|
|
def test_streamed_tar_exact_hashes_sizes_and_regular_members(self):
|
|
self.put('runtime/results/segment.jsonl', b'a' * (snapshot.BLOCK + 7))
|
|
self.put('runtime/state/file.json', b'{}')
|
|
self.output.mkdir()
|
|
inventory = snapshot._inventory()
|
|
files, archive = snapshot._write_tar(self.output, None, inventory, self.report)
|
|
data = self.output.joinpath('files.tar').read_bytes()
|
|
self.assertEqual(archive, {'bytes': len(data), 'sha256': hashlib.sha256(data).hexdigest()})
|
|
with tarfile.open(fileobj=io.BytesIO(data)) as tar:
|
|
self.assertEqual(len(tar.getmembers()), len(files))
|
|
for member, record in zip(tar.getmembers(), files):
|
|
content = tar.extractfile(member).read()
|
|
self.assertTrue(member.isfile())
|
|
self.assertFalse(member.issym() or member.islnk())
|
|
self.assertEqual(member.mode, 0o600)
|
|
self.assertEqual(record, {'path': member.name, 'size': len(content),
|
|
'sha256': hashlib.sha256(content).hexdigest()})
|
|
|
|
def test_changed_file_before_copy_rejected(self):
|
|
path = self.put('runtime/results/a', b'before')
|
|
inventory = snapshot._inventory()
|
|
path.write_bytes(b'after-and-longer')
|
|
self.output.mkdir()
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._write_tar(self.output, None, inventory, self.report)
|
|
|
|
def test_changed_file_during_copy_rejected_by_fstat(self):
|
|
path = self.put('runtime/results/a', b'before')
|
|
inventory = snapshot._inventory()
|
|
self.output.mkdir()
|
|
original = snapshot.HashReader.read
|
|
|
|
def read(reader, size):
|
|
block = original(reader, size)
|
|
with path.open('ab') as handle:
|
|
handle.write(b'changed')
|
|
return block
|
|
|
|
with patch.object(snapshot.HashReader, 'read', read), self.assertRaises(snapshot.Failure):
|
|
snapshot._write_tar(self.output, None, inventory, self.report)
|
|
|
|
def test_changed_open_handle_fails_before_read(self):
|
|
self.put('runtime/results/a')
|
|
inventory = snapshot._inventory()
|
|
name, entry = next(iter(inventory.files.items()))
|
|
inventory.files[name] = replace(entry, fingerprint=entry.fingerprint[:-1] + (999,))
|
|
self.output.mkdir()
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._write_tar(self.output, None, inventory, self.report)
|
|
|
|
def test_handle_ctime_change_during_copy_is_still_rejected(self):
|
|
self.put('runtime/results/a')
|
|
inventory = snapshot._inventory()
|
|
self.output.mkdir()
|
|
original = os.fstat
|
|
calls = []
|
|
|
|
def fstat(fd):
|
|
info = original(fd)
|
|
calls.append(fd)
|
|
if len(calls) == 3:
|
|
fields = ('st_dev', 'st_ino', 'st_mode', 'st_nlink', 'st_size', 'st_mtime_ns', 'st_file_attributes')
|
|
return SimpleNamespace(**{name: getattr(info, name, 0) for name in fields},
|
|
st_ctime_ns=info.st_ctime_ns + 1)
|
|
return info
|
|
|
|
with patch.object(snapshot.os, 'fstat', side_effect=fstat), self.assertRaises(snapshot.Failure):
|
|
snapshot._write_tar(self.output, None, inventory, self.report)
|
|
|
|
def test_inventory_detects_addition_deletion_and_rename(self):
|
|
path = self.put('runtime/results/a')
|
|
before = snapshot._inventory()
|
|
new = self.put('runtime/results/b')
|
|
self.assertNotEqual(snapshot._inventory(), before)
|
|
new.unlink()
|
|
path.rename(path.with_name('renamed'))
|
|
self.assertNotEqual(snapshot._inventory(), before)
|
|
|
|
|
|
class OutputTests(Fixture):
|
|
def test_output_requires_exact_windows_parent_and_safe_new_name(self):
|
|
with patch.object(snapshot, 'IMPORTS_ROOT', Path(r'D:\truf-docker\docker\imports')):
|
|
self.assertEqual(str(snapshot._output_path(r'D:\truf-docker\docker\imports\new-20260915')),
|
|
str(Path(r'D:\truf-docker\docker\imports\new-20260915')))
|
|
for path in (r'D:\other\new', r'D:truf-docker\docker\imports\new', 'relative',
|
|
r'D:\truf-docker\docker\imports\old\new',
|
|
r'D:\truf-docker\docker\imports\..\imports\new',
|
|
r'D:\truf-docker\docker\imports\new:stream',
|
|
r'D:\truf-docker\docker\imports\CON', r'D:\truf-docker\docker\imports\new.'):
|
|
with self.subTest(path=path), self.assertRaises(snapshot.Failure):
|
|
snapshot._output_path(path)
|
|
|
|
def test_existing_output_is_never_reused(self):
|
|
self.output.mkdir()
|
|
marker = self.output / 'marker'
|
|
marker.write_bytes(b'unchanged')
|
|
with self.assertRaises(FileExistsError):
|
|
snapshot._prepare_output(self.output, None)
|
|
self.assertEqual(marker.read_bytes(), b'unchanged')
|
|
self.secure.assert_not_called()
|
|
|
|
def test_output_private_before_sensitive_writes(self):
|
|
seen = []
|
|
|
|
def secure(path, security, directory=False):
|
|
seen.append((path.name, directory))
|
|
if not directory:
|
|
self.assertEqual(path.stat().st_size, 0)
|
|
|
|
self.secure.side_effect = secure
|
|
snapshot._prepare_output(self.output, None)
|
|
with snapshot._output_file(self.output / 'database.dump', None) as handle:
|
|
self.assertEqual(seen, [('unique', True), ('database.dump', False)])
|
|
handle.write(b'synthetic-sensitive-data')
|
|
|
|
def test_acl_failure_prevents_writes_and_manifest(self):
|
|
self.secure.side_effect = RuntimeError(SECRET)
|
|
with self.capture_context(), self.assertRaises(RuntimeError):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertEqual(list(self.output.iterdir()), [])
|
|
self.source.pg.maintenance_start.assert_not_called()
|
|
|
|
def test_private_acl_accepts_only_exact_user_and_system(self):
|
|
sid = 'S-1-5-21-111-222-333-1001'
|
|
security = SimpleNamespace(reject_reparse_components=Mock(), _windows_current_user_sid=lambda: sid)
|
|
good = 'O:' + sid + 'D:P(A;OICI;FA;;;' + sid + ')(A;OICI;FA;;;SY)'
|
|
for sddl in (good, good + '(A;OICI;FA;;;BA)', good.replace(';SY)', ';WD)'), good.replace('D:P', 'D:')):
|
|
security._windows_private_sddl = lambda path: sddl
|
|
if sddl == good:
|
|
snapshot._verify_private_acl(self.output, security, True)
|
|
else:
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot._verify_private_acl(self.output, security, True)
|
|
|
|
def test_acl_helper_is_narrowed_using_no_reparse_native_handle(self):
|
|
sid = 'S-1-5-21-111-222-333-1001'
|
|
descriptors = []
|
|
|
|
class Information(ctypes.Structure):
|
|
_fields_ = [('dwFileAttributes', ctypes.c_uint)]
|
|
|
|
def convert(sddl, *args):
|
|
descriptors.append(sddl)
|
|
return 1
|
|
|
|
security = SimpleNamespace(
|
|
harden_private_directory=Mock(), harden_private_file=Mock(),
|
|
reject_reparse_components=Mock(), _windows_current_user_sid=lambda: sid,
|
|
_windows_private_sddl=lambda path: 'O:' + sid + descriptors[-1],
|
|
_CONVERT_SDDL=Mock(side_effect=convert), _CREATE_FILE=Mock(return_value=123),
|
|
_BY_HANDLE_FILE_INFORMATION=Information, _GET_FILE_INFORMATION=Mock(return_value=1),
|
|
_SET_KERNEL_OBJECT_SECURITY=Mock(return_value=1), _CLOSE_HANDLE=Mock(), _LOCAL_FREE=Mock())
|
|
REAL_SECURE_PATH(self.output, security, directory=True)
|
|
security.harden_private_directory.assert_called_once_with(str(self.output))
|
|
self.assertEqual(descriptors, ['D:P(A;OICI;FA;;;' + sid + ')(A;OICI;FA;;;SY)'])
|
|
self.assertEqual(security._CREATE_FILE.call_args.args[5], 0x2200000)
|
|
self.assertEqual(security._SET_KERNEL_OBJECT_SECURITY.call_args.args[1], 0x80000004)
|
|
security._CLOSE_HANDLE.assert_called_once_with(123)
|
|
security._LOCAL_FREE.assert_called_once()
|
|
|
|
def test_partial_manifest_is_not_published_on_write_failure(self):
|
|
self.output.mkdir()
|
|
original = snapshot.HashWriter.write
|
|
|
|
def write(writer, block):
|
|
original(writer, block)
|
|
raise OSError(SECRET)
|
|
|
|
with patch.object(snapshot.HashWriter, 'write', write), self.assertRaises(OSError):
|
|
snapshot._publish_manifest(self.output, None, {'format': 'truf-windows-snapshot-v1'})
|
|
self.assertTrue(self.output.joinpath('manifest.json.partial').exists())
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
|
|
class LoaderTests(Fixture):
|
|
def test_loader_uses_original_bootstrap_and_original_environment_only(self):
|
|
for name in ('child_bootstrap.py', 'postgres_runtime.py', 'runtime_security.py', 'config.yaml'):
|
|
self.put('app/' + name)
|
|
self.put('.env.postgres')
|
|
self.put('runtime/postgres/cluster_identity.json')
|
|
bootstrap = SimpleNamespace(_enable_dependency_paths=Mock())
|
|
spec = SimpleNamespace(loader=SimpleNamespace(exec_module=Mock()))
|
|
config = {'global': {'root_dir': str(self.root), 'project_dir': str(self.root / 'app'),
|
|
'runtime_dir': str(self.root / 'runtime'), 'postgres_data_dir': str(snapshot.POSTGRES_DATA),
|
|
'result_bundle_dir': str(self.bundles)}}
|
|
source = self.runtime()
|
|
|
|
def load_environment(*args):
|
|
for key in ('TRUF_POSTGRES_PASSWORD', 'PGPASSWORD', 'DATABASE_URL', 'SCANNER_DB_URL', 'SCANNER_RUNTIME_DIR'):
|
|
self.assertNotIn(key, os.environ)
|
|
return str(self.root / '.env.postgres')
|
|
|
|
pg = SimpleNamespace(
|
|
__name__='postgres_runtime', __file__=str(self.root / 'app/postgres_runtime.py'),
|
|
_load_config=Mock(return_value=config), load_postgres_environment=Mock(side_effect=load_environment),
|
|
postgres_runtime_paths=Mock(return_value={'data_dir': str(snapshot.POSTGRES_DATA),
|
|
'postgres_dir': str(self.root / 'runtime/postgres')}),
|
|
canonical_database_url=Mock(return_value=source.dsn))
|
|
security = SimpleNamespace(__name__='runtime_security', __file__=str(self.root / 'app/runtime_security.py'),
|
|
preflight_lifecycle_paths=Mock())
|
|
inherited = {key: SECRET for key in ('TRUF_POSTGRES_PASSWORD', 'PGPASSWORD', 'DATABASE_URL',
|
|
'SCANNER_DB_URL', 'SCANNER_RUNTIME_DIR')}
|
|
with patch.object(snapshot.importlib.util, 'spec_from_file_location', return_value=spec), \
|
|
patch.object(snapshot.importlib.util, 'module_from_spec', return_value=bootstrap), \
|
|
patch.object(snapshot.importlib, 'import_module', side_effect=[pg, security]), \
|
|
patch.object(sys, 'path', sys.path[:]), patch.dict(os.environ, inherited):
|
|
loaded = REAL_LOAD_SOURCE()
|
|
bootstrap._enable_dependency_paths.assert_called_once_with('postgres-runtime')
|
|
spec.loader.exec_module.assert_called_once_with(bootstrap)
|
|
self.assertIs(loaded.pg, pg)
|
|
self.assertEqual(len(loaded.inputs), 3)
|
|
snapshot.subprocess.Popen.assert_not_called()
|
|
snapshot.subprocess.run.assert_not_called()
|
|
|
|
|
|
class FakeProcess:
|
|
def __init__(self, output=b'', returncode=0):
|
|
self.stdin = io.BytesIO()
|
|
self.stdout = io.BytesIO(output)
|
|
self.returncode = returncode
|
|
self.killed = False
|
|
|
|
def wait(self, timeout):
|
|
return self.returncode
|
|
|
|
def poll(self):
|
|
return self.returncode
|
|
|
|
def kill(self):
|
|
self.killed = True
|
|
self.returncode = -1
|
|
|
|
|
|
class DatabaseTests(Fixture):
|
|
def metadata(self):
|
|
return {'version_num': 160004, 'system_identifier': self.identity['system_identifier'],
|
|
'database_name': 'fixture', 'user_name': 'fixture', 'port': 5432,
|
|
'data_directory': str(snapshot.POSTGRES_DATA), 'in_recovery': False, 'read_only': 'on',
|
|
'snapshot': '00000003-0000001A-1', 'database_bytes': 100,
|
|
'tables': ['normal', 'a"; odd', 'runtime_operations',
|
|
'runtime_operations_control', 'runtime_audit_events'],
|
|
'other_clients': 0, 'all_sessions_visible': True,
|
|
'sequences': [['public', 'counter'], ['Other"Schema', 'Seq";name'],
|
|
['public', 'runtime_audit_events_id_seq']]}
|
|
|
|
def test_database_identity_requires_every_bound_value(self):
|
|
self.runtime()
|
|
snapshot._validate_database(self.metadata(), self.identity)
|
|
for key, value in (('version_num', 150004), ('system_identifier', 'wrong'), ('database_name', 'wrong'),
|
|
('user_name', 'wrong'), ('port', 6543), ('data_directory', 'wrong'),
|
|
('in_recovery', True), ('read_only', 'off'), ('other_clients', 1),
|
|
('all_sessions_visible', False), ('other_clients', False),
|
|
('snapshot', 'bad\ntext'), ('database_bytes', -1), ('tables', ['duplicate', 'duplicate']),
|
|
('sequences', [['public', 'same'], ['public', 'same']]), ('sequences', ['not-a-pair']),
|
|
('sequences', [['public', 'bad\0name']])):
|
|
with self.subTest(key=key), self.assertRaises(snapshot.Failure):
|
|
snapshot._validate_database(dict(self.metadata(), **{key: value}), self.identity)
|
|
|
|
def test_identity_compares_original_configuration_and_dsn(self):
|
|
source = self.runtime()
|
|
snapshot._validate_identity(source, self.identity)
|
|
for key, value in (('pg_major', 15), ('database', 'other'), ('user', 'other'), ('port', 6543),
|
|
('data_directory', 'wrong'), ('system_identifier', 'not-numeric')):
|
|
with self.subTest(key=key), self.assertRaises(snapshot.Failure):
|
|
snapshot._validate_identity(source, dict(self.identity, **{key: value}))
|
|
|
|
def test_client_environment_is_read_only_and_contains_no_inherited_overrides(self):
|
|
source = self.runtime()
|
|
with patch.dict(os.environ, {'PGSERVICE': SECRET, 'PGHOST': 'foreign', 'PGOPTIONS': 'unsafe',
|
|
'DATABASE_URL': SECRET, 'SCANNER_DB_URL': SECRET, 'PATH': SECRET}):
|
|
env = snapshot._client_environment(source.dsn)
|
|
self.assertEqual(env['PGPASSWORD'], SECRET)
|
|
self.assertEqual(env['PGHOST'], '127.0.0.1')
|
|
self.assertIn('default_transaction_read_only=on', env['PGOPTIONS'])
|
|
for key in ('PGSERVICE', 'DATABASE_URL', 'SCANNER_DB_URL', 'PATH'):
|
|
self.assertNotIn(key, env)
|
|
|
|
def test_query_uses_stdin_not_arguments_and_bounded_json(self):
|
|
process = FakeProcess(b'17\n')
|
|
sql = 'SELECT count(*) FROM ONLY "public"."a""; odd";'
|
|
with patch.object(snapshot, '_deadline', wraps=snapshot._deadline) as deadline:
|
|
self.assertEqual(snapshot._query(process, sql, 99), 17)
|
|
self.assertEqual(process.stdin.getvalue(), (sql + '\n').encode())
|
|
deadline.assert_called_once_with(process, 99)
|
|
for output in (SECRET.encode() + b'\n', b'1', b'x' * 33 + b'\n'):
|
|
with patch.object(snapshot, 'MAX_METADATA', 32), self.assertRaises(snapshot.Failure):
|
|
snapshot._query(FakeProcess(output), 'SELECT 1;')
|
|
|
|
def test_counts_and_dump_use_shared_snapshot_and_no_secret_arguments(self):
|
|
source = self.runtime()
|
|
for name in ('psql.exe', 'pg_dump.exe'):
|
|
self.put('runtime/postgres/pgsql/bin/' + name)
|
|
self.output.mkdir()
|
|
statements, commands = [], []
|
|
metadata = self.metadata()
|
|
psql, dump = FakeProcess(), FakeProcess(b'PGDMPsynthetic-dump')
|
|
|
|
@contextlib.contextmanager
|
|
def client(command, env, timeout, interactive=False):
|
|
commands.append((command, env, timeout))
|
|
yield psql if interactive else dump
|
|
|
|
def query(process, sql, timeout=snapshot.QUERY_TIMEOUT):
|
|
statements.append((sql, timeout))
|
|
if sql == snapshot.DATABASE_METADATA:
|
|
return metadata
|
|
if 'count(*) FROM ONLY' in sql:
|
|
return 7
|
|
if "'last_value', last_value" in sql:
|
|
return {'last_value': 42, 'is_called': True}
|
|
return 0
|
|
|
|
def version(command, **kwargs):
|
|
commands.append((command, kwargs['env'], kwargs['timeout']))
|
|
return SimpleNamespace(returncode=0, stdout=(Path(command[0]).stem + ' (PostgreSQL) 16.4\n').encode())
|
|
|
|
with patch.object(snapshot, '_client', client), patch.object(snapshot, '_query', query), \
|
|
patch.object(snapshot.subprocess, 'run', side_effect=version):
|
|
result = snapshot._capture_database(source, self.identity, self.output, snapshot._inventory(), self.report)
|
|
self.assertEqual(result['table_counts'], {
|
|
'normal': 7,
|
|
'a"; odd': 7,
|
|
'runtime_operations': 7,
|
|
'runtime_operations_control': 7,
|
|
'runtime_audit_events': 7,
|
|
})
|
|
self.assertEqual(result['sequence_states'], {
|
|
'Other"Schema': {'Seq";name': {'last_value': 42, 'is_called': True}},
|
|
'public': {
|
|
'counter': {'last_value': 42, 'is_called': True},
|
|
'runtime_audit_events_id_seq': {'last_value': 42, 'is_called': True},
|
|
},
|
|
})
|
|
self.assertEqual(result['sequence_count'], 3)
|
|
self.assertIn('non-MVCC', result['sequence_state_mode'])
|
|
sequence_sql = ("SELECT pg_catalog.json_build_object('last_value', last_value, 'is_called', is_called) "
|
|
'FROM "Other""Schema"."Seq"";name";')
|
|
self.assertEqual(statements.count((sequence_sql, snapshot.QUERY_TIMEOUT)), 2)
|
|
self.assertIn(('SELECT count(*) FROM ONLY "public"."a""; odd";', snapshot.COUNT_TIMEOUT), statements)
|
|
dump_command = commands[-1][0]
|
|
for flag in ('--format=custom', '--no-owner', '--no-acl', '--no-tablespaces', '--compress=1',
|
|
'--no-password', '--snapshot=' + metadata['snapshot']):
|
|
self.assertIn(flag, dump_command)
|
|
for command, env, timeout in commands:
|
|
self.assertNotIn(SECRET, repr(command))
|
|
self.assertEqual(env['PGPASSWORD'], SECRET)
|
|
self.assertGreater(timeout, 0)
|
|
data = self.output.joinpath('database.dump').read_bytes()
|
|
self.assertEqual(result['sha256'], hashlib.sha256(data).hexdigest())
|
|
self.assertEqual(result['bytes'], len(data))
|
|
|
|
def test_client_nonzero_and_stderr_suppression(self):
|
|
process = FakeProcess(returncode=9)
|
|
with patch.object(snapshot.subprocess, 'Popen', return_value=process) as popen:
|
|
with self.assertRaises(snapshot.Failure):
|
|
with snapshot._client(['fixture.exe'], {}, 5):
|
|
pass
|
|
self.assertEqual(popen.call_args.kwargs['stderr'], subprocess.DEVNULL)
|
|
self.assertTrue(process.stdout.closed)
|
|
|
|
def test_client_timeout_kills_process(self):
|
|
process = FakeProcess()
|
|
|
|
class Timer:
|
|
def __init__(self, seconds, callback):
|
|
self.callback = callback
|
|
|
|
def start(self):
|
|
self.callback()
|
|
|
|
def cancel(self):
|
|
pass
|
|
|
|
def join(self):
|
|
pass
|
|
|
|
with patch.object(snapshot.threading, 'Timer', Timer), self.assertRaises(snapshot.Failure) as caught:
|
|
with snapshot._deadline(process, 1):
|
|
pass
|
|
self.assertEqual(caught.exception.code, 124)
|
|
self.assertTrue(process.killed)
|
|
|
|
def test_timeout_is_preserved_when_killed_pipe_produces_an_error(self):
|
|
process = FakeProcess()
|
|
|
|
class Timer:
|
|
def __init__(self, seconds, callback):
|
|
self.callback = callback
|
|
|
|
def start(self):
|
|
self.callback()
|
|
|
|
def cancel(self):
|
|
pass
|
|
|
|
def join(self):
|
|
pass
|
|
|
|
with patch.object(snapshot.threading, 'Timer', Timer), self.assertRaises(snapshot.Failure) as caught:
|
|
with snapshot._deadline(process, 1):
|
|
raise OSError(SECRET)
|
|
self.assertEqual(caught.exception.code, 124)
|
|
|
|
def test_client_session_guard_refreshes_statistics(self):
|
|
with patch.object(snapshot, '_query', side_effect=[None, 1]) as query, self.assertRaises(snapshot.Failure):
|
|
snapshot._no_other_clients(FakeProcess())
|
|
self.assertIn('pg_stat_clear_snapshot()', query.call_args_list[0].args[1])
|
|
|
|
|
|
def test_sequence_state_accepts_uncalled_bigint_and_rejects_payloads(self):
|
|
value = {'last_value': -9223372036854775808, 'is_called': False}
|
|
with patch.object(snapshot, '_query', return_value=value):
|
|
self.assertEqual(snapshot._sequence_states(FakeProcess(), [['public', 'counter']]),
|
|
{'public': {'counter': value}})
|
|
for value in (None, {'last_value': True, 'is_called': False},
|
|
{'last_value': 3, 'is_called': 1}, {'last_value': 3},
|
|
{'last_value': 3, 'is_called': False, 'payload': SECRET}):
|
|
with self.subTest(value=value), patch.object(snapshot, '_query', return_value=value), \
|
|
self.assertRaises(snapshot.Failure):
|
|
snapshot._sequence_states(FakeProcess(), [['public', 'counter']])
|
|
|
|
def test_changed_sequence_state_invalidates_completed_dump(self):
|
|
source = self.runtime()
|
|
for name in ('psql.exe', 'pg_dump.exe'):
|
|
self.put('runtime/postgres/pgsql/bin/' + name)
|
|
self.output.mkdir()
|
|
metadata = self.metadata()
|
|
|
|
@contextlib.contextmanager
|
|
def client(command, env, timeout, interactive=False):
|
|
yield FakeProcess() if interactive else FakeProcess(b'PGDMPsynthetic-dump')
|
|
|
|
def query(process, sql, timeout=snapshot.QUERY_TIMEOUT):
|
|
return metadata if sql == snapshot.DATABASE_METADATA else 0
|
|
|
|
def version(command, **kwargs):
|
|
return SimpleNamespace(returncode=0, stdout=(Path(command[0]).stem + ' (PostgreSQL) 16.4\n').encode())
|
|
|
|
before = {'public': {'counter': {'last_value': 1, 'is_called': False}}}
|
|
for after in ({'public': {'counter': {'last_value': 2, 'is_called': False}}},
|
|
{'public': {'counter': {'last_value': 1, 'is_called': True}}}):
|
|
with self.subTest(after=after), patch.object(snapshot, '_client', client), \
|
|
patch.object(snapshot, '_query', query), \
|
|
patch.object(snapshot.subprocess, 'run', side_effect=version), \
|
|
patch.object(snapshot, '_sequence_states', side_effect=[before, after]), \
|
|
self.assertRaises(snapshot.Failure):
|
|
snapshot._capture_database(source, self.identity, self.output, snapshot._inventory(), self.report)
|
|
self.assertTrue(self.output.joinpath('database.dump').exists())
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.output.joinpath('database.dump').unlink()
|
|
|
|
|
|
class SignalTests(Fixture):
|
|
def test_handlers_only_set_flag_and_restore_previous_handlers(self):
|
|
for number in self.handlers:
|
|
with self.subTest(number=number):
|
|
with snapshot._defer_signals() as checkpoint:
|
|
checkpoint()
|
|
self.send_signal(number)
|
|
self.send_signal(number)
|
|
with self.assertRaises(snapshot.Failure):
|
|
checkpoint()
|
|
self.assertEqual(self.handlers, self.previous_handlers)
|
|
|
|
def test_partial_handler_installation_failure_never_loads_source(self):
|
|
if len(self.handlers) < 2:
|
|
self.skipTest('SIGBREAK is Windows-only')
|
|
register = self.signal.side_effect
|
|
last = list(self.handlers)[-1]
|
|
|
|
def fail(number, handler):
|
|
if number == last and handler is not self.previous_handlers[last]:
|
|
raise OSError(SECRET)
|
|
return register(number, handler)
|
|
|
|
self.signal.side_effect = fail
|
|
with self.assertRaises(OSError):
|
|
snapshot.capture(self.output, self.report)
|
|
snapshot._load_source.assert_not_called()
|
|
self.assertEqual(self.handlers, self.previous_handlers)
|
|
|
|
def test_cancel_before_start_neither_starts_nor_stops_source(self):
|
|
def report(number, count, size):
|
|
if number == 5:
|
|
self.send_signal(snapshot.signal.SIGINT)
|
|
|
|
with self.capture_context() as source, self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, report)
|
|
source.pg.maintenance_start.assert_not_called()
|
|
source.pg.maintenance_stop.assert_not_called()
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_spawn_window_signal_cannot_unwind_unpublished_child(self):
|
|
for number in self.handlers:
|
|
for outcome in ('READY', 'OWNED_START_UNCERTAIN'):
|
|
with self.subTest(number=number, outcome=outcome), self.capture_context() as source:
|
|
self.output = self.imports / (str(int(number)) + '-' + outcome)
|
|
backend = self.backend
|
|
backend._start_requested_wall_time = None
|
|
backend._accepted_start_at_monotonic = None
|
|
backend._started_postmaster_observed = False
|
|
backend._expected_process = None
|
|
child = SimpleNamespace(alive=False, visible=False)
|
|
events, window_probes, compensation_held, released_live, closed_states = [], [], [], [], []
|
|
|
|
def probe():
|
|
if child.alive and child.visible:
|
|
backend._accepted_start_at_monotonic = None
|
|
backend._started_postmaster_observed = True
|
|
backend._expected_process = child
|
|
return SimpleNamespace(kind='READY')
|
|
if backend._accepted_start_at_monotonic is not None:
|
|
return SimpleNamespace(kind='OWNED_START_UNCERTAIN')
|
|
# The real helper has no PID/listener before publication;
|
|
# without its accepted-start latch, it reports STOPPED.
|
|
return SimpleNamespace(kind='STOPPED')
|
|
|
|
def close():
|
|
closed_states.append((backend._accepted_start_at_monotonic,
|
|
backend._started_postmaster_observed))
|
|
backend._expected_process = None
|
|
|
|
def start(config, backend):
|
|
try:
|
|
backend._start_requested_wall_time = 1.0
|
|
child.alive = True
|
|
events.append('spawned')
|
|
window_probes.extend([backend.probe().kind, backend.probe().kind])
|
|
self.send_signal(number) # Popen -> 100ms sleep/poll gap.
|
|
events.append('signal-returned')
|
|
events.append('polled-still-running')
|
|
backend._accepted_start_at_monotonic = 10.0
|
|
events.append('accepted')
|
|
child.visible = outcome == 'READY'
|
|
result = backend.probe()
|
|
if result.kind != 'READY':
|
|
raise RuntimeError(SECRET)
|
|
return result
|
|
finally:
|
|
# maintenance_start's finally closes the process
|
|
# handle, not the accepted-start/observed latches.
|
|
backend.close()
|
|
|
|
def stop(config, backend):
|
|
try:
|
|
if backend.probe().kind == 'STOPPED':
|
|
return SimpleNamespace(completed=True, stopped=True)
|
|
if backend._accepted_start_at_monotonic is not None and backend._expected_process is None:
|
|
compensation_held.append(self.authority.acquired and child.alive and not child.visible)
|
|
child.visible = True
|
|
backend.probe()
|
|
events.append('identity-stop' if backend._expected_process is child else 'unsafe-stop')
|
|
child.alive = False
|
|
backend._accepted_start_at_monotonic = None
|
|
backend._start_requested_wall_time = None
|
|
backend._started_postmaster_observed = False
|
|
return SimpleNamespace(completed=True, stopped=True)
|
|
finally:
|
|
backend.close()
|
|
|
|
release = self.authority.release.side_effect
|
|
|
|
def release_checked():
|
|
released_live.append(child.alive)
|
|
events.append('release')
|
|
release()
|
|
|
|
backend.probe.side_effect = probe
|
|
backend.close.side_effect = close
|
|
source.pg.maintenance_start.side_effect = start
|
|
source.pg.maintenance_stop.side_effect = stop
|
|
self.authority.release.side_effect = release_checked
|
|
caught = None
|
|
try:
|
|
snapshot.capture(self.output, self.report)
|
|
except BaseException as exc:
|
|
caught = exc
|
|
self.assertIsInstance(caught, snapshot.Failure)
|
|
self.assertEqual(window_probes, ['STOPPED', 'STOPPED'])
|
|
self.assertEqual(events, ['spawned', 'signal-returned', 'polled-still-running',
|
|
'accepted', 'identity-stop', 'release'])
|
|
self.assertEqual(released_live, [False])
|
|
self.assertFalse(child.alive)
|
|
self.assertNotIn('database', self.calls)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertEqual(self.handlers, self.previous_handlers)
|
|
if outcome == 'OWNED_START_UNCERTAIN':
|
|
self.assertEqual(closed_states[0], (10.0, False))
|
|
self.assertEqual(compensation_held, [True])
|
|
else:
|
|
self.assertEqual(closed_states[0], (None, True))
|
|
|
|
def test_signal_during_stop_cannot_interrupt_confirmation_retries(self):
|
|
with self.capture_context() as source:
|
|
normal = source.pg.maintenance_stop.side_effect
|
|
|
|
def stop(config, backend):
|
|
self.assertTrue(self.authority.acquired)
|
|
if source.pg.maintenance_stop.call_count == 1:
|
|
self.send_signal(snapshot.signal.SIGINT)
|
|
return SimpleNamespace(completed=False, stopped=False)
|
|
return normal(config, backend)
|
|
|
|
source.pg.maintenance_stop.side_effect = stop
|
|
with patch.object(snapshot.time, 'sleep') as sleep, \
|
|
patch.object(snapshot, '_write_tar') as tar, self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, self.report)
|
|
sleep.assert_called_once_with(2)
|
|
tar.assert_not_called()
|
|
self.assertEqual(source.pg.maintenance_stop.call_count, 3)
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
self.assertFalse(self.authority.acquired)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
def test_signal_during_release_is_deferred_until_release_finishes(self):
|
|
with self.capture_context():
|
|
release = self.authority.release.side_effect
|
|
|
|
def interrupted_release():
|
|
self.send_signal(snapshot.signal.SIGINT)
|
|
release()
|
|
|
|
self.authority.release.side_effect = interrupted_release
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertFalse(self.authority.acquired)
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
self.assertTrue(self.output.joinpath('files.tar').exists())
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
|
|
class CaptureTests(Fixture):
|
|
def test_success_stops_before_tar_holds_authority_until_manifest(self):
|
|
self.put('runtime/results/segment.jsonl', b'fixture-row\n')
|
|
original_tar, original_publish = snapshot._write_tar, snapshot._publish_manifest
|
|
|
|
def tar(*args):
|
|
self.assertTrue(self.authority.acquired)
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
self.calls.append('tar')
|
|
return original_tar(*args)
|
|
|
|
def publish(*args):
|
|
self.assertTrue(self.authority.acquired)
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
self.calls.append('manifest')
|
|
return original_publish(*args)
|
|
|
|
with self.capture_context(), patch.object(snapshot, '_write_tar', side_effect=tar), \
|
|
patch.object(snapshot, '_publish_manifest', side_effect=publish):
|
|
snapshot.capture(self.output, self.report)
|
|
manifest = json.loads(self.output.joinpath('manifest.json').read_text())
|
|
self.assertEqual(manifest['format'], 'truf-windows-snapshot-v1')
|
|
self.assertTrue(manifest['source']['supervisor_stopped'])
|
|
self.assertTrue(manifest['source']['postgres_stopped'])
|
|
self.assertEqual(manifest['counts']['files'], 1)
|
|
self.assertEqual(self.calls, ['acquire', 'start', 'database', 'stop', 'tar', 'stop', 'manifest', 'release'])
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_supervisor_metadata_refuses_start_without_stopping_existing_source(self):
|
|
for relative in ('runtime/control/supervisor.instance.json', 'runtime/logs/supervisor.instance.json',
|
|
'runtime/logs/supervisor.pid'):
|
|
path = self.put(relative)
|
|
self.output = self.imports / ('case-' + str(len(list(self.imports.iterdir()))))
|
|
with self.capture_context() as source, self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, self.report)
|
|
source.pg.maintenance_start.assert_not_called()
|
|
source.pg.maintenance_stop.assert_not_called()
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertFalse(self.authority.acquired)
|
|
path.unlink()
|
|
|
|
def test_authenticated_stopped_guard_rejects_all_other_states(self):
|
|
for state in ('READY', 'RECOVERING', 'FOREIGN_OR_CONFIG_ERROR', 'OWNED_START_UNCERTAIN'):
|
|
self.output = self.imports / state
|
|
with self.capture_context() as source, self.assertRaises(snapshot.Failure):
|
|
self.backend.state = state
|
|
snapshot.capture(self.output, self.report)
|
|
source.pg.maintenance_start.assert_not_called()
|
|
source.pg.maintenance_stop.assert_not_called()
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
def test_stopped_guard_is_rechecked_after_inventory_before_start(self):
|
|
with self.capture_context() as source, self.assertRaises(snapshot.Failure):
|
|
self.backend.probe.side_effect = [SimpleNamespace(kind='STOPPED'), SimpleNamespace(kind='READY')]
|
|
snapshot.capture(self.output, self.report)
|
|
source.pg.maintenance_start.assert_not_called()
|
|
source.pg.maintenance_stop.assert_not_called()
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_partial_maintenance_start_still_gets_confirmed_stop(self):
|
|
with self.capture_context() as source:
|
|
source.pg.maintenance_start.side_effect = RuntimeError(SECRET)
|
|
with self.assertRaises(RuntimeError):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertGreaterEqual(source.pg.maintenance_stop.call_count, 1)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_failed_database_keeps_partial_output_without_manifest(self):
|
|
def failure(*args):
|
|
self.output.joinpath('database.dump').write_bytes(b'partial')
|
|
raise RuntimeError(SECRET)
|
|
|
|
with self.capture_context(database=failure), self.assertRaises(RuntimeError):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertTrue(self.output.joinpath('database.dump').exists())
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
self.assertEqual(self.calls[-1], 'release')
|
|
|
|
def test_failed_tar_keeps_dump_without_manifest(self):
|
|
with self.capture_context(), patch.object(snapshot, '_write_tar', side_effect=RuntimeError(SECRET)), \
|
|
self.assertRaises(RuntimeError):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertTrue(self.output.joinpath('database.dump').exists())
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertEqual(self.backend.state, 'STOPPED')
|
|
|
|
def test_changed_inventory_prevents_publication(self):
|
|
original = snapshot._write_tar
|
|
|
|
def tar(*args):
|
|
result = original(*args)
|
|
self.put('runtime/state/new.json')
|
|
return result
|
|
|
|
with self.capture_context(), patch.object(snapshot, '_write_tar', side_effect=tar), \
|
|
self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
def test_source_change_during_final_stop_is_not_published(self):
|
|
with self.capture_context() as source:
|
|
normal = source.pg.maintenance_stop.side_effect
|
|
|
|
def stop(*args, **kwargs):
|
|
result = normal(*args, **kwargs)
|
|
if source.pg.maintenance_stop.call_count == 2:
|
|
self.put('runtime/state/late-change')
|
|
return result
|
|
|
|
source.pg.maintenance_stop.side_effect = stop
|
|
with self.assertRaises(snapshot.Failure):
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
def test_post_publication_cleanup_failure_revokes_manifest(self):
|
|
with self.capture_context(), self.assertRaises(OSError):
|
|
self.backend.close.side_effect = OSError(SECRET)
|
|
snapshot.capture(self.output, self.report)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
self.assertTrue(self.output.joinpath('files.tar').exists())
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_stop_retries_retain_authority_and_no_manifest(self):
|
|
with self.capture_context() as source:
|
|
normal_stop = source.pg.maintenance_stop.side_effect
|
|
tries = []
|
|
|
|
def stop(config, backend):
|
|
self.assertTrue(self.authority.acquired)
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
tries.append(1)
|
|
if len(tries) == 1:
|
|
raise KeyboardInterrupt(SECRET)
|
|
return normal_stop(config, backend)
|
|
|
|
source.pg.maintenance_stop.side_effect = stop
|
|
with patch.object(snapshot.time, 'sleep') as sleep:
|
|
snapshot.capture(self.output, self.report)
|
|
sleep.assert_called_once_with(2)
|
|
self.assertTrue(self.output.joinpath('manifest.json').exists())
|
|
self.assertFalse(self.authority.acquired)
|
|
|
|
def test_changed_identity_input_prevents_start(self):
|
|
path = self.put('app/config.yaml')
|
|
with self.capture_context() as source, self.assertRaises(snapshot.Failure):
|
|
source.inputs[path] = snapshot._fingerprint(path.stat())
|
|
path.write_bytes(b'changed-fixture')
|
|
snapshot.capture(self.output, self.report)
|
|
source.pg.maintenance_start.assert_not_called()
|
|
self.assertFalse(self.output.joinpath('manifest.json').exists())
|
|
|
|
def test_main_redacts_raw_source_failure_and_restores_environment(self):
|
|
output, error = io.StringIO(), io.StringIO()
|
|
|
|
def failure(*args):
|
|
print(SECRET)
|
|
os.environ['SYNTHETIC_SNAPSHOT_TEST'] = SECRET
|
|
raise RuntimeError(SECRET)
|
|
|
|
with patch.object(snapshot.os, 'name', 'nt'), patch.object(snapshot, '_output_path', return_value=self.output), \
|
|
patch.object(snapshot, '_load_source', side_effect=failure), \
|
|
contextlib.redirect_stdout(output), contextlib.redirect_stderr(error):
|
|
result = snapshot.main(['capture', '--output', 'fixture'])
|
|
self.assertEqual(result, 1)
|
|
self.assertNotIn(SECRET, output.getvalue() + error.getvalue())
|
|
self.assertRegex(error.getvalue(), r'^2 1\n$')
|
|
self.assertRegex(output.getvalue(), r'^2 0 0\n$')
|
|
self.assertNotIn('SYNTHETIC_SNAPSHOT_TEST', os.environ)
|
|
|
|
def test_non_windows_refused_before_source_loading(self):
|
|
with patch.object(snapshot.os, 'name', 'posix'), contextlib.redirect_stderr(io.StringIO()):
|
|
self.assertEqual(snapshot.main(['capture', '--output', 'fixture']), 1)
|
|
snapshot._load_source.assert_not_called()
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|