"""Run directly: python -I -S -B tests/test_container_import.py ContainerImportTests. Stdlib only. All data is synthetic and temporary. Subprocesses, mount metadata, runtime/DB modules, and POSIX-only metadata are mocked; no operational app is loaded. The pure target_identity helper is loaded without importing the app. """ import contextlib import copy import hashlib import io import json import os from pathlib import Path, PurePosixPath import stat import subprocess import sys import tarfile import tempfile from types import ModuleType, SimpleNamespace import unittest from unittest import mock APP = Path(__file__).resolve().parents[1] / 'app' PASSWORD = 'generated-linux-fixture-password-' + 'x' * 32 PRIVATE_VALUE = 'synthetic-private-provider-value' def load(name, path): module = ModuleType(name) module.__file__ = str(path) exec(compile(path.read_text(encoding='utf-8'), str(path), 'exec'), module.__dict__) return module def sha(payload): return hashlib.sha256(payload).hexdigest() class ContainerImportTests(unittest.TestCase): def setUp(self): self.module = m = load('container_import_under_test', APP / 'container_import.py') self.temporary = tempfile.TemporaryDirectory(prefix='truf-import-test-') self.addCleanup(self.temporary.cleanup) self.root = Path(self.temporary.name) data, incoming, app = (self.root / name for name in ('data', 'import', 'app')) for path in (data, incoming, app): path.mkdir(mode=0o700) self.trace, self.locked, self.stopped = [], False, False self.runtime = r = SimpleNamespace( DATA=data, DEFAULT_CONFIG=app / 'config.linux.yaml', PROVISIONED=data / '.provisioned.json', INITIALIZED=data / 'initialized.json', INITIALIZE_LOCK=data / 'initialize.lock', PASSWORD=data / 'postgres-password', FORMAT='truf-container-data-v1', _shutdown_requested=False, DIRECTORIES=('home', 'config', 'runtime-linux', 'runtime-linux/state', 'runtime-linux/state/gharchive_cache', 'runtime-linux/results', 'runtime-linux/queues', 'runtime-linux/keychecks', 'runtime-linux/postman_cache', 'runtime-linux/result_spool', 'runtime-linux/logs', 'runtime-linux/postgres', 'runtime-linux/postgres/logs', 'postgres-linux', 'scanner-work', 'scanner-result-bundles', 'scanner-result-bundles/tmp', 'scanner-result-bundles/ready', 'scanner-result-bundles/quarantine'), ) for name in r.DIRECTORIES: (data / name).mkdir(mode=0o700) r.PROVISIONED.write_bytes(json.dumps({'format': r.FORMAT, 'uid': 10001, 'gid': 10001}).encode('ascii')) for name in ('.provision.lock', 'initialize.lock', 'runtime-linux/proxy.txt'): (data / name).write_bytes(b'') (data / 'config/secrets.yaml').write_bytes(b'{}\n') r.PASSWORD.write_bytes((PASSWORD + '\n').encode('ascii')) r.DEFAULT_CONFIG.write_text('{}\n', encoding='ascii') def private(path, *, directory=False): path = Path(path) self.assertTrue(path.is_relative_to(self.root), 'fixture escaped its temp root') self.assertTrue(path.is_dir() if directory else path.is_file()) self.assertFalse(path.is_symlink()) return path r.private_path = mock.Mock(side_effect=private) r._read_json = lambda path: json.loads(Path(path).read_bytes()) r._bootstrap_command = lambda target, *args: ['fixture-bootstrap', target, *args] r.prepare_environment = mock.Mock(side_effect=AssertionError('unexpected environment preparation')) m.IMPORT = incoming m.os = SimpleNamespace(**vars(os)) for name in ('O_NOFOLLOW', 'O_DIRECTORY', 'O_NONBLOCK'): setattr(m.os, name, getattr(os, name, 0)) if os.name == 'nt': fingerprint = m._fingerprint # Windows 3.12 stat/fstat use different ctime bases. Model Linux's # shared basis without weakening the production fingerprint check. m._fingerprint = lambda info: fingerprint(info)[:-1] + (0,) dsn = 'postgresql://truf:' + PASSWORD + '@127.0.0.1:5432/truf' m.os.environ = {key: dsn for key in ('SCANNER_DB_URL', 'DATABASE_URL', 'TRUF_MANAGED_POSTGRES_DSN')} m.os.environ.update(TRUF_POSTGRES_PASSWORD=PASSWORD) m.subprocess = SimpleNamespace(**vars(subprocess)) m.subprocess.Popen = mock.Mock(side_effect=AssertionError('a real subprocess is forbidden')) self.native_fsync_dir = m._fsync_dir m._fsync_dir = mock.Mock() m.shutil = SimpleNamespace(disk_usage=mock.Mock(return_value=SimpleNamespace(free=200 * m.GIB))) m.time = SimpleNamespace(monotonic=mock.Mock(return_value=0), sleep=mock.Mock()) m.signal = SimpleNamespace(SIGINT=2, signal=mock.Mock(return_value='fixture-prior-handler')) self.progress = mock.Mock() self.identity = {'pg_major': 16, 'system_identifier': '2222222222', 'database': 'truf', 'user': 'truf', 'port': 5432, 'data_directory': str(data / 'postgres-linux')} @contextlib.contextmanager def initialize_lock(path): self.assertEqual(Path(path), r.INITIALIZE_LOCK) self.assertFalse(self.locked) self.locked = True self.trace.append('initialize-lock') try: yield finally: self.trace.append('initialize-unlock') self.locked = False def publish(path, value): self.assertTrue(self.locked) self.assertTrue(self.stopped, 'marker before positive stop') self.assertTrue((data / 'config/windows-import-report.json').is_file()) self.trace.append('marker') with Path(path).open('xb') as handle: handle.write(json.dumps(value).encode('ascii')) self.security = SimpleNamespace(PrivateFileLock=initialize_lock, write_private_json_exclusive=mock.Mock(side_effect=publish), require_trusted_native_executable=mock.Mock(side_effect=lambda path: path)) self.module_patch = mock.patch.dict(sys.modules, { 'runtime_security': self.security, 'target_identity': load('target_identity_fixture', APP / 'target_identity.py'), 'postgres_runtime': SimpleNamespace(), 'scanner_db': SimpleNamespace(), }) self.module_patch.start() self.addCleanup(self.module_patch.stop) self.payloads = { 'windows-archive/app/config.yaml': b'{"global": {}}\n', 'config/secrets.yaml': ('fixture: ' + PRIVATE_VALUE + '\r\n').encode('ascii'), 'config/trufflehog-custom-detectors.yaml': b'detectors: []\r\n', 'runtime-linux/proxy.txt': b'127.0.0.1:9:fixture:private\r\n', } self.snapshot() def snapshot(self, entries=None, manifest_files=None): entries = list(self.payloads.items()) if entries is None else entries buffer = io.BytesIO() with tarfile.open(fileobj=buffer, mode='w', format=tarfile.PAX_FORMAT) as archive: for name, value in entries: if isinstance(value, tarfile.TarInfo): archive.addfile(value) else: info = tarfile.TarInfo(name) info.size, info.mode, info.mtime = len(value), 0o600, 0 archive.addfile(info, io.BytesIO(value)) self.tar_bytes = buffer.getvalue() dump = b'PGDMP-synthetic-logical-dump-not-a-real-database' self.manifest = { 'format': self.module.SNAPSHOT_FORMAT, 'source': {'root': r'D:\truf', 'postgres_data_dir': r'S:\postgres-data', 'supervisor_stopped': True, 'postgres_stopped': True}, 'database': {'version_num': 160014, 'system_identifier': '1111111111', 'database_name': 'windows_db', 'user_name': 'windows_user', 'port': 15432, 'data_directory': r'S:\postgres-data', 'table_counts': {'target_queue': 2}, 'bytes': len(dump), 'sha256': sha(dump), 'sequence_states': {'public': {'target_queue_id_seq': {'last_value': 7, 'is_called': True}}}, 'sequence_count': 1}, 'archive': {'bytes': len(self.tar_bytes), 'sha256': sha(self.tar_bytes)}, 'files': manifest_files if manifest_files is not None else [ {'path': name, 'size': len(value), 'sha256': sha(value)} for name, value in self.payloads.items() ], } (self.module.IMPORT / 'files.tar').write_bytes(self.tar_bytes) (self.module.IMPORT / 'database.dump').write_bytes(dump) self.save_manifest() def save_manifest(self): self.manifest_bytes = json.dumps(self.manifest, ensure_ascii=True).encode('ascii') self.manifest_sha = sha(self.manifest_bytes) (self.module.IMPORT / 'manifest.json').write_bytes(self.manifest_bytes) def parsed(self): self.save_manifest() return self.module._manifest(self.manifest_bytes, self.manifest_sha)[1] def archive(self, *, extract=False): files = self.parsed() placeholders = self.module._fresh(self.runtime, files) if extract else None self.module._archive(self.runtime, io.BytesIO(self.tar_bytes), self.manifest['archive'], files, self.progress, placeholders) def public(self): stdout, stderr = io.StringIO(), io.StringIO() with contextlib.redirect_stdout(stdout), contextlib.redirect_stderr(stderr): code = self.module.import_snapshot(self.runtime, self.manifest_sha) return code, stdout.getvalue(), stderr.getvalue() def fake_public_database(self): m, r = self.module, self.runtime m._mounts = mock.Mock() def configuration(runtime): self.assertTrue(self.locked) path = r.DATA / 'config/windows-import.yaml' digest = m._write(r, path, b'global: {}\n') return path, {}, ['global.root_dir'], digest def restore(runtime, path, config, manifest, files, dump, fingerprint, progress, report): self.assertTrue(self.locked) self.assertFalse(r.INITIALIZED.exists()) self.trace.append('restore') report.update(raw={'tables': 1, 'rows': 2, 'sequences_verified': 1}, cutover_sha256='c' * 64, migration_rows=29, postman_reviewed=2, postman_adjusted=2) self.stopped = True self.trace.append('maintenance-stop') return self.identity @contextlib.contextmanager def authority(runtime, config, source, progress, *, stopped=False): self.assertTrue(self.locked) self.assertTrue(stopped) self.assertTrue(self.stopped) self.trace.append('stop-confirmed-under-authority') yield self.identity self.trace.append('cluster-unlock') m._configuration = mock.Mock(side_effect=configuration) m._restore = mock.Mock(side_effect=restore) m._authority = mock.Mock(side_effect=authority) def fake_restore_dependencies(self, failure=None): m, r = self.module, self.runtime self.commands = [] self.sql_trace = [] def command(runtime, argv, timeout, progress, **kwargs): action = argv[2] if argv[1:2] == ['postgres-runtime'] else ( 'migrate' if argv[1:2] == ['migrate-runtime-safety'] else 'pg_restore') self.commands.append((action, argv, timeout, kwargs)) self.trace.append(action) if action == failure: raise m.Failure() @contextlib.contextmanager def authority(*args, **kwargs): yield self.identity @contextlib.contextmanager def connection(*args, **kwargs): conn = mock.Mock() def execute(sql, *params): self.sql_trace.append((sql, tuple(self.trace))) return mock.Mock() conn.execute.side_effect = execute yield conn m._command = mock.Mock(side_effect=command) m._authority = mock.Mock(side_effect=authority) m._connect = mock.Mock(side_effect=connection) m._database_equality = mock.Mock(return_value={'tables': 1, 'rows': 2, 'sequences_verified': 1}) m._hash_input = mock.Mock() m._preserved = mock.Mock(return_value={'pipeline_quarantine': {'count': 1, 'sha256': 'd' * 64}}) m._projection_files = mock.Mock() m._pipeline_counts = mock.Mock(return_value={'worker_leases': 0}) m._rebase_postman = mock.Mock(return_value={'postman_reviewed': 1, 'postman_adjusted': 1}) m._online = mock.Mock() self.db = mock.Mock(enabled=True) self.db.conn.is_postgres = True self.db.conn.execute.return_value.fetchone.return_value = {'count': 29} self.db.final_cutover_status.return_value = {'evidence_sha256': 'c' * 64} sys.modules['scanner_db'] = SimpleNamespace(ScannerDB=mock.Mock(return_value=self.db)) def restore(self): files = self.parsed() self.restore_report = {'manifest_sha256': self.manifest_sha} return self.module._restore(self.runtime, self.runtime.DATA / 'config/windows-import.yaml', {}, self.manifest, files, io.BytesIO(), (), self.progress, self.restore_report) def fake_authority_dependencies(self): lock, backend = mock.Mock(), mock.Mock() lock.acquire.side_effect = lambda: self.trace.append('cluster-lock') lock.release.side_effect = lambda: self.trace.append('cluster-release') self.security.ClusterAuthorityLock = mock.Mock(return_value=lock) backend.probe.return_value.kind = 'ready' backend.close.side_effect = lambda: self.trace.append('backend-close') def stop(config, backend): self.stopped = True self.trace.append('confirmed-stop') backend.probe.return_value.kind = 'stopped' return SimpleNamespace(completed=True, stopped=True, detail=PRIVATE_VALUE) pg = SimpleNamespace(PostgresBackend=mock.Mock(return_value=backend), ProbeKind=SimpleNamespace(READY='ready', STOPPED='stopped'), verify_cluster_identity=mock.Mock(return_value=self.identity), maintenance_stop=mock.Mock(side_effect=stop)) sys.modules['postgres_runtime'] = pg return pg, backend, lock def test_manifest_hash_bound_and_duplicate_json_keys(self): m = self.module self.assertEqual(set(self.parsed()), set(self.payloads)) for digest in ('', 'x' * 64, '0' * 64): with self.subTest(digest=digest), self.assertRaises(m.Failure): m._manifest(self.manifest_bytes, digest) self.assertTrue(m._manifest(self.manifest_bytes, self.manifest_sha.upper())) bad = b'{"format":"x","format":"y"}' with self.assertRaises(m.Failure): m._manifest(bad, sha(bad)) with mock.patch.object(m, 'MAX_MANIFEST', len(self.manifest_bytes) - 1), self.assertRaises(m.Failure): m._manifest(self.manifest_bytes, self.manifest_sha) def test_manifest_requires_stopped_source_exact_version_and_paths(self): baseline = copy.deepcopy(self.manifest) for section, key, value in ( ('source', 'supervisor_stopped', 1), ('source', 'postgres_stopped', False), ('source', 'root', r'D:\other'), ('database', 'data_directory', '/data/postgres-linux'), ('database', 'version_num', 150010), ('database', 'port', True), ): self.manifest = copy.deepcopy(baseline) self.manifest[section][key] = value with self.subTest(key=key), self.assertRaises(self.module.Failure): self.parsed() def test_manifest_rejects_traversal_absolute_and_nonnormal_paths(self): for suffix in ('../escape', '/absolute', 'a//b', 'a/./b', 'a/../b', 'a\\b', 'C:payload', 'space /x', 'dot./x', 'nul\x00x', 'line\nx', 'e\u0301/x'): with self.subTest(suffix=suffix), self.assertRaises(self.module.Failure): self.module._destination('runtime-linux/results/' + suffix) for name in ('/data/runtime-linux/results/x', 'postgres-linux/PG_VERSION', 'postgres-password', 'runtime-linux/logs/x', 'scanner-result-bundles/other/x', 'proxy.txt'): with self.subTest(name=name), self.assertRaises(self.module.Failure): self.module._destination(name) def test_manifest_casefold_and_parent_collisions_in_both_orders(self): original = copy.deepcopy(self.manifest['files']) for left, right in (('a', 'A'), ('a', 'a/b'), ('a/b', 'a'), ('A/b', 'a/c')): self.manifest['files'] = original + [ {'path': 'runtime-linux/results/' + p, 'size': 0, 'sha256': sha(b'')} for p in (left, right) ] with self.subTest(left=left, right=right), self.assertRaises(self.module.Failure): self.parsed() def test_manifest_rejects_live_authority_but_keeps_inert_archive(self): for suffix in ('scan_limiter.db', 'scan_limiter-old.db-wal', 'supervisor.instance.json', 'cluster_identity.json', 'anything.lock.backup', 'postmaster.pid', 'control/report.json', 'postgres/PG_VERSION', 'janitor.cursor.json'): with self.subTest(suffix=suffix), self.assertRaises(self.module.Failure): self.module._destination('runtime-linux/state/' + suffix) for name in ('windows-archive/.env.postgres', 'windows-archive/runtime/state/scan_limiter.db', 'windows-archive/runtime/control/report.json'): self.assertTrue(self.module._destination(name)) def test_manifest_rejects_invalid_counts_and_sequence_types(self): self.manifest['database']['table_counts']['target_queue'] = True with self.assertRaises(self.module.Failure): self.parsed() self.manifest['database']['table_counts']['target_queue'] = 2 self.manifest['database']['sequence_states']['public']['target_queue_id_seq']['is_called'] = 1 with self.assertRaises(self.module.Failure): self.parsed() self.manifest['database']['sequence_states'] = None with self.assertRaises(self.module.Failure): self.parsed() self.manifest['database'].pop('sequence_states') self.manifest['database'].pop('sequence_count') self.assertTrue(self.parsed()) def test_archive_streams_exact_bytes_and_replaces_only_placeholders(self): self.payloads['runtime-linux/results/deep/file.jsonl'] = b'\x00\xff\r\nexact bytes\n' self.payloads['windows-archive/inert/unicode-\u00e9.txt'] = b'archive bytes' self.snapshot() self.archive() self.assertFalse((self.runtime.DATA / 'windows-archive').exists()) self.archive(extract=True) for name, payload in self.payloads.items(): self.assertEqual((self.runtime.DATA / name).read_bytes(), payload) self.assertTrue(self.module._fsync_dir.called) def test_archive_rejects_links_specials_and_directory_members(self): for kind in (tarfile.SYMTYPE, tarfile.LNKTYPE, tarfile.FIFOTYPE, tarfile.CHRTYPE, tarfile.DIRTYPE): info = tarfile.TarInfo('runtime-linux/results/extra') info.type = kind if kind in (tarfile.SYMTYPE, tarfile.LNKTYPE): info.linkname = '/outside' self.snapshot(entries=list(self.payloads.items()) + [(info.name, info)]) with self.subTest(kind=kind), self.assertRaises(self.module.Failure): self.archive() def test_archive_rejects_duplicate_missing_and_unmanifested_files(self): entries = list(self.payloads.items()) for bad in (entries + [entries[0]], entries[:-1], entries + [('windows-archive/extra', b'x')]): self.snapshot(entries=bad) with self.subTest(length=len(bad)), self.assertRaises(self.module.Failure): self.archive() def test_archive_rejects_wrong_size_file_hash_and_whole_hash(self): self.manifest['files'][0]['size'] += 1 with self.assertRaises(self.module.Failure): self.archive() self.snapshot() self.manifest['files'][0]['sha256'] = '0' * 64 with self.assertRaises(self.module.Failure): self.archive() self.snapshot() self.manifest['archive']['sha256'] = '0' * 64 with self.assertRaises(self.module.Failure): self.archive() def test_archive_rejects_nonzero_trailer_even_with_matching_whole_hash(self): self.tar_bytes = self.tar_bytes[:-1] + b'x' self.manifest['archive']['sha256'] = sha(self.tar_bytes) with self.assertRaises(self.module.Failure): self.archive() def test_archive_rejects_oversized_pax_without_reading_payload(self): info = tarfile.TarInfo('pax') info.type, info.size = tarfile.XHDTYPE, 65537 bad = info.tobuf(format=tarfile.USTAR_FORMAT) + bytes(10240) self.tar_bytes = bad self.manifest['archive'] = {'bytes': len(bad), 'sha256': sha(bad)} with self.assertRaises(self.module.Failure): self.archive() def test_input_rejects_hardlink_and_nonregular_leaf(self): path = self.module.IMPORT / 'database.dump' alias = self.root / 'hardlink' os.link(path, alias) with self.assertRaises(self.module.Failure): with self.module._input(path): self.fail('a linked input was accepted') with self.assertRaises(self.module.Failure): with self.module._input(self.module.IMPORT): self.fail('a directory input was accepted') def test_input_rejects_symlink_components_by_metadata(self): path = mock.Mock() parent = mock.Mock() parent.lstat.return_value = SimpleNamespace(st_mode=stat.S_IFLNK, st_file_attributes=0) path.parents = (parent,) with self.assertRaises(self.module.Failure): self.module._chain(path) path.lstat.assert_not_called() def test_input_changed_after_open_is_rejected(self): path = self.module.IMPORT / 'database.dump' with self.assertRaises(self.module.Failure): with self.module._input(path) as (handle, _before): path.write_bytes(b'changed while opened') handle.seek(0) def test_dump_hash_and_magic_rechecked_on_same_retained_handle(self): path = self.module.IMPORT / 'database.dump' with self.module._input(path) as (handle, before): self.module._hash_input(self.runtime, path, handle, before, self.manifest['database'], dump=True) self.assertEqual(handle.tell(), 0) bad = dict(self.manifest['database'], sha256='f' * 64) with self.assertRaises(self.module.Failure): self.module._hash_input(self.runtime, path, handle, before, bad, dump=True) path.write_bytes(b'NOTPG-not-a-dump') with self.module._input(path) as (handle, before), self.assertRaises(self.module.Failure): self.module._hash_input(self.runtime, path, handle, before, {'bytes': path.stat().st_size, 'sha256': sha(path.read_bytes())}, dump=True) def test_freshness_refuses_markers_pgdata_identity_and_unexpected_files(self): m, r = self.module, self.runtime self.assertEqual(set(m._fresh(r)), set(m.PLACEHOLDERS)) for name in ('initialized.json', 'postgres-linux/PG_VERSION', 'runtime-linux/postgres/cluster_identity.json', 'runtime-linux/results/existing', 'config/windows-import-report.json', 'home/.history'): path = r.DATA / name path.write_bytes(b'preexisting') with self.subTest(name=name), self.assertRaises(m.Failure): m._fresh(r) path.unlink() (r.DATA / 'windows-archive').mkdir() with self.assertRaises(m.Failure): m._fresh(r) def test_freshness_validates_planned_paths_against_provisioned_directories(self): for name in ('runtime-linux/state/gharchive_cache', 'runtime-linux/state/GHARCHIVE_CACHE/file'): with self.subTest(name=name), self.assertRaises(self.module.Failure): self.module._fresh(self.runtime, [name]) def test_only_exact_placeholder_bytes_and_private_mode_are_accepted(self): m, r = self.module, self.runtime for value in (b'{}', b'{}\r\n', b'{ }\n', b'{}\n ', b'existing secret'): (r.DATA / 'config/secrets.yaml').write_bytes(value) with self.subTest(value=value), self.assertRaises(m.Failure): m._placeholder(r, 'config/secrets.yaml') (r.DATA / 'config/secrets.yaml').write_bytes(b'{}\n') r.private_path.side_effect = RuntimeError('fixture nonprivate mode') with self.assertRaises(RuntimeError): m._placeholder(r, 'runtime-linux/proxy.txt') def test_placeholder_is_rechecked_immediately_before_replace(self): m, r = self.module, self.runtime original = m._placeholder def changed(runtime, name, before=None): if before is not None: (r.DATA / name).write_bytes(b'changed by another writer') return original(runtime, name, before) m._placeholder = changed with self.assertRaises(m.Failure): self.archive(extract=True) self.assertEqual((r.DATA / 'config/secrets.yaml').read_bytes(), b'changed by another writer') self.assertTrue((r.DATA / 'config/secrets.yaml.windows-import-partial').is_file()) def test_space_estimate_reserve_and_upper_bound(self): m = self.module result = m._space(self.runtime, self.manifest) self.assertEqual(result['database_estimate_bytes'], 24 * m.GIB) self.assertFalse(result['database_estimate_from_metadata']) m.shutil.disk_usage.return_value.free = result['required_free_bytes'] - 1 with self.assertRaises(m.Failure): m._space(self.runtime, self.manifest) self.manifest['database']['database_bytes'] = 21 * m.GIB m.shutil.disk_usage.return_value.free = 200 * m.GIB self.assertEqual(m._space(self.runtime, self.manifest)['database_estimate_bytes'], 21 * m.GIB) self.manifest['database'].pop('database_bytes') self.manifest['database']['bytes'] = m.MAX_ESTIMATE // 4 + 1 with self.assertRaises(m.Failure): m._space(self.runtime, self.manifest) def test_mount_must_be_exact_readonly_independent_and_without_submounts(self): m = self.module class MountPath(PurePosixPath): def lstat(self): return SimpleNamespace(st_mode=stat.S_IFDIR, st_file_attributes=0) def stat(self): return SimpleNamespace(st_dev=2 if str(self) == '/import' else 3) m.IMPORT = MountPath('/import') runtime = SimpleNamespace(DATA=MountPath('/data')) m.os.listdir = mock.Mock(return_value=['manifest.json', 'files.tar', 'database.dump']) m.os.ST_RDONLY = 1 m.os.statvfs = mock.Mock(return_value=SimpleNamespace(f_flag=1)) good = (b'20 1 8:1 /snapshot /import ro - ext4 /dev/a rw\n' b'21 1 8:2 /volume /data rw - ext4 /dev/b rw\n') with mock.patch.object(m, 'open', mock.mock_open(read_data=good), create=True): m._mounts(runtime) for bad in (good.replace(b'/import ro', b'/import rw'), good.replace(b'/import ', b'/import/nested '), good + b'22 20 8:3 / /import/files.tar ro - ext4 /dev/c ro\n', good.replace(b'8:1 /snapshot', b'8:2 /volume')): with mock.patch.object(m, 'open', mock.mock_open(read_data=bad), create=True), self.assertRaises(m.Failure): m._mounts(runtime) m.os.statvfs.return_value.f_flag = 0 with mock.patch.object(m, 'open', mock.mock_open(read_data=good), create=True), self.assertRaises(m.Failure): m._mounts(runtime) def test_configuration_translates_archive_and_validates_new_private_path(self): m, r = self.module, self.runtime self.archive(extract=True) original = (r.DATA / 'windows-archive/app/config.yaml').read_bytes() translated = {'global': {}} translator = mock.Mock(return_value=(translated, ['global.root_dir'])) sys.modules['container_import_config'] = SimpleNamespace(translate_windows_config=translator) self.addCleanup(sys.modules.pop, 'container_import_config', None) yaml = SimpleNamespace(safe_load=json.loads, safe_dump=lambda value, **kwargs: json.dumps(value)) with mock.patch.dict(sys.modules, {'yaml': yaml}): expected = {key: str(r.DATA / 'runtime-linux' / folder) for key, folder in ( ('results_dir', 'results'), ('queue_dir', 'queues'), ('state_dir', 'state'), ('log_dir', 'logs'), ('keycheck_dir', 'keychecks'), ('postman_cache_dir', 'postman_cache'), ('result_spool_dir', 'result_spool'), ('legacy_result_spool_dir', 'result_spool'), ('scan_limiter_db', 'state/scan_limiter.db'), )} expected.update(proxy_file=str(r.DATA / 'runtime-linux/proxy.txt'), trufflehog_config=str(r.DATA / 'config/trufflehog-custom-detectors.yaml')) r.prepare_environment.side_effect = None r.prepare_environment.return_value = {'global': expected} path, config, adjusted, digest = m._configuration(r) translator.assert_called_once_with({'global': {}}, {}) r.prepare_environment.assert_called_once_with(path) self.assertEqual(path, r.DATA / 'config/windows-import.yaml') self.assertEqual(sha(path.read_bytes()), digest) self.assertEqual(adjusted, ['global.root_dir']) self.assertEqual((r.DATA / 'windows-archive/app/config.yaml').read_bytes(), original) def test_linux_credentials_must_match_generated_private_password(self): self.assertEqual(self.module._credentials(self.runtime)[1], PASSWORD) for value in ('postgresql://source:source@127.0.0.1:5432/truf', 'postgresql://truf:' + PASSWORD + '@provider.example:5432/truf'): self.module.os.environ['SCANNER_DB_URL'] = value with self.assertRaises(self.module.Failure): self.module._credentials(self.runtime) def test_command_cancellation_reaps_nonlifecycle_child_without_diagnostics(self): m = self.module child = mock.Mock() def wait(timeout): if not child.kill.called: self.runtime._shutdown_requested = True raise subprocess.TimeoutExpired(['private-argv'], timeout) return -9 child.wait.side_effect = wait m.subprocess.Popen.side_effect = None m.subprocess.Popen.return_value = child with self.assertRaises(m.Failure) as caught: m._command(self.runtime, ['safe-program'], 60, self.progress) self.assertEqual(caught.exception.code, 130) child.kill.assert_called_once() self.assertEqual(child.wait.call_count, 2) kwargs = m.subprocess.Popen.call_args.kwargs self.assertEqual(kwargs['stdout'], subprocess.DEVNULL) self.assertEqual(kwargs['stderr'], subprocess.DEVNULL) def test_lifecycle_child_is_retained_after_cancellation_until_it_exits(self): m = self.module child = mock.Mock() child.wait.side_effect = [subprocess.TimeoutExpired(['fixture'], 1), 0] original_wait = child.wait.side_effect def wait(timeout): self.runtime._shutdown_requested = True value = next(original_wait) if isinstance(value, Exception): raise value return value child.wait.side_effect = wait m.subprocess.Popen.side_effect = None m.subprocess.Popen.return_value = child with self.assertRaises(m.Failure) as caught: m._command(self.runtime, ['fixture-lifecycle'], 60, self.progress, lifecycle=True) self.assertEqual(caught.exception.code, 130) child.kill.assert_not_called() child.terminate.assert_not_called() self.assertEqual(child.wait.call_count, 2) self.assertEqual([call.kwargs['diagnostic'][:4] for call in self.progress.call_args_list], [(1, 130, 0, 1), (9, 130, 0, 1)]) def test_finite_restore_timeout_kills_and_reaps_only_client(self): m = self.module child = mock.Mock() child.wait.side_effect = [subprocess.TimeoutExpired(['fixture'], 1), -9] m.subprocess.Popen.side_effect = None m.subprocess.Popen.return_value = child m.time.monotonic.side_effect = [0, 2] with self.assertRaises(m.Failure) as caught: m._command(self.runtime, ['fixture-pg-restore'], 1, self.progress) self.assertEqual(caught.exception.code, 124) child.kill.assert_called_once() self.assertEqual(child.wait.call_count, 2) def test_broken_progress_pipe_cannot_abandon_lifecycle_child(self): m = self.module child = mock.Mock() child.wait.side_effect = [subprocess.TimeoutExpired(['fixture'], 1), subprocess.TimeoutExpired(['fixture'], 1), 0] m.subprocess.Popen.side_effect = None m.subprocess.Popen.return_value = child m.time.monotonic.side_effect = [0, 2] self.progress.side_effect = BrokenPipeError('synthetic closed output') with self.assertRaises(m.Failure) as caught: m._command(self.runtime, ['fixture-initialize-empty'], 1, self.progress, lifecycle=True) self.assertEqual(caught.exception.code, 124) self.assertEqual(child.wait.call_count, 3) child.kill.assert_not_called() child.terminate.assert_not_called() def test_killed_lifecycle_child_cannot_claim_its_cleanup_contract(self): m = self.module child = mock.Mock() child.wait.return_value = -9 m.subprocess.Popen.side_effect = None m.subprocess.Popen.return_value = child with self.assertRaises(m.Failure) as caught: m._command(self.runtime, ['fixture-initialize-empty'], 60, self.progress, lifecycle=True) self.assertTrue(caught.exception.uncertain) child.kill.assert_not_called() def test_uncertain_init_requires_stop_contract_before_return(self): self.fake_restore_dependencies() original = self.module._command.side_effect def killed(runtime, argv, *args, **kwargs): if 'initialize-empty' in argv: self.trace.append('killed-init') raise self.module.Failure(uncertain=True) return original(runtime, argv, *args, **kwargs) self.module._command.side_effect = killed with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace, ['killed-init', 'maintenance-stop']) self.module._connect.assert_not_called() def test_failed_maintenance_start_always_enters_confirmed_stop(self): self.fake_restore_dependencies('maintenance-start') with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace, ['initialize-empty', 'maintenance-start', 'maintenance-stop']) self.module._connect.assert_not_called() def test_failed_start_still_stops_when_cleanup_progress_raises(self): self.fake_restore_dependencies('maintenance-start') def progress(phase, *args, **kwargs): if phase == 12: raise BrokenPipeError('synthetic closed output') self.progress.side_effect = progress with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace[-1], 'maintenance-stop') def test_failed_restore_always_stops_without_reconciliation(self): self.fake_restore_dependencies('pg_restore') with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace[-1], 'maintenance-stop') self.module._rebase_postman.assert_not_called() self.assertFalse(self.runtime.INITIALIZED.exists()) def test_failed_migration_always_stops_and_leaves_data(self): self.fake_restore_dependencies('migrate') with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace[-1], 'maintenance-stop') self.assertTrue((self.runtime.DATA / 'config/windows-import-raw.json').exists()) self.db.require_final_cutover.assert_not_called() def test_raw_mismatch_stops_before_any_target_or_migration_change(self): self.fake_restore_dependencies() self.module._database_equality.side_effect = [{}, self.module.Failure()] with self.assertRaises(self.module.Failure): self.restore() self.assertEqual(self.trace[-1], 'maintenance-stop') self.assertNotIn('migrate', self.trace) self.module._rebase_postman.assert_not_called() def test_final_evidence_mismatch_stops_and_does_not_publish(self): self.fake_restore_dependencies() self.module._preserved.side_effect = [{}, self.module.Failure(2, 1)] with self.assertRaises(self.module.Failure) as caught: self.restore() self.assertEqual(caught.exception.code, 2) self.assertEqual(self.trace[-1], 'maintenance-stop') self.security.write_private_json_exclusive.assert_not_called() def test_restore_argv_credentials_order_and_readonly_final_contracts(self): self.fake_restore_dependencies() identity = self.restore() self.assertEqual(identity, self.identity) self.assertEqual(self.trace, ['initialize-empty', 'maintenance-start', 'pg_restore', 'migrate', 'maintenance-stop']) for action, argv, timeout, kwargs in self.commands: self.assertNotIn(PASSWORD, str(argv)) self.assertNotIn('postgresql://', str(argv)) self.assertNotIn('--initialize-base', argv) self.assertGreater(timeout, 0) if action == 'pg_restore': for flag in ('--single-transaction', '--exit-on-error', '--no-owner', '--no-acl', '--no-tablespaces'): self.assertIn(flag, argv) self.assertEqual(kwargs['env']['PGPASSWORD'], PASSWORD) self.assertIn('stdin', kwargs) self.assertGreaterEqual(timeout, 6 * 3600) self.assertEqual(self.module._hash_input.call_count, 2) self.assertTrue(self.module._database_equality.call_args_list[0].kwargs['empty']) self.db.conn.execute.assert_any_call('SET default_transaction_read_only = on') self.db.require_runtime_safety_schema.assert_called_once_with() self.db.require_final_cutover.assert_called_once_with() self.db.close.assert_called_once_with() def test_recovery_is_separate_bounded_fenced_cli_not_blanket_reset(self): self.fake_restore_dependencies() self.module._pipeline_counts.side_effect = [{'worker_leases': 1}, {'worker_leases': 0}] self.restore() recovery = [argv for _action, argv, _timeout, _kwargs in self.commands if '--recover-stale-result-pipeline' in argv] self.assertEqual(len(recovery), 1) for flag in ('--apply', '--sources-stopped', '--max-rows', '--max-seconds'): self.assertIn(flag, recovery[0]) self.assertNotIn('--initialize-base', recovery[0]) self.assertEqual(self.module._rebase_postman.call_count, 2) def test_recovery_rebases_inserted_derived_target_before_migration_and_preserves_history(self): m = self.module rebase = m._rebase_postman target, normalized, cache = self.postman_fixture() relative = next(iter(cache)) self.payloads[relative] = (self.runtime.DATA / relative).read_bytes() self.snapshot() self.fake_restore_dependencies() m._rebase_postman = rebase m._pipeline_counts.side_effect = [{'result_reservations': 1}, {'result_reservations': 0}] derived = {name: None for name in m.FENCES} derived.update(id=3, platform='postman', status='pending', resolver_state=None, target=json.dumps(target), normalized_target=normalized) rows = [dict(derived, id=1, status='done', current_result_reservation_id=41), dict(derived, id=2, status='quarantined', lease_token='historical-token')] history = copy.deepcopy(rows) writes = [] conn = mock.Mock() def execute(sql, params=()): if sql.startswith('SELECT * FROM public.target_queue'): self.assertIn("status IN ('pending','deferred','in_progress')", sql) return SimpleNamespace(fetchall=lambda: [row for row in rows if row['id'] > params[0] and row['status'] in ('pending', 'deferred', 'in_progress')]) if sql.startswith('UPDATE public.target_queue SET target = %s'): replacement, row_id, original, identity = params row = next(row for row in rows if row['id'] == row_id) self.assertEqual((row['target'], row['normalized_target']), (original, identity)) row['target'] = replacement writes.append(row_id) return SimpleNamespace(rowcount=1) self.assertTrue(sql.startswith('ALTER ROLE "truf" IN DATABASE "truf" ')) return mock.Mock() conn.execute.side_effect = execute @contextlib.contextmanager def connection(*args, **kwargs): yield conn m._connect.side_effect = connection original_command = m._command.side_effect def command(runtime, argv, *args, **kwargs): result = original_command(runtime, argv, *args, **kwargs) if '--recover-stale-result-pipeline' in argv: self.assertEqual(writes, []) rows.append(copy.deepcopy(derived)) elif argv[1:2] == ['migrate-runtime-safety']: self.assertEqual(writes, [3], 'derived target must be rebased before normal migration') return result m._command.side_effect = command self.restore() self.assertEqual(rows[:2], history) self.assertEqual(rows[2]['normalized_target'], normalized) self.assertEqual(json.loads(rows[2]['target'])['cache_path'], str(self.runtime.DATA / relative)) self.assertEqual((self.restore_report['postman_reviewed'], self.restore_report['postman_adjusted']), (1, 1)) self.assertEqual(self.trace[-1], 'maintenance-stop') def test_post_recovery_target_review_failure_stops_before_normal_migration(self): self.fake_restore_dependencies() m = self.module m._pipeline_counts.return_value = {'result_reservations': 1} m._rebase_postman.side_effect = [ {'postman_reviewed': 0, 'postman_adjusted': 0}, m.Failure(2, 1), ] with self.assertRaises(m.Failure) as caught: self.restore() self.assertEqual((caught.exception.code, caught.exception.count), (2, 1)) self.assertEqual(self.trace[-1], 'maintenance-stop') migrations = [argv for _action, argv, _timeout, _kwargs in self.commands if argv[1:2] == ['migrate-runtime-safety']] self.assertEqual(len(migrations), 1) self.assertIn('--recover-stale-result-pipeline', migrations[0]) self.assertFalse(self.runtime.INITIALIZED.exists()) def test_linux_role_diagnostics_cover_migration_and_are_reset_after_verification(self): self.fake_restore_dependencies() self.restore() settings = {setting for setting, _value in self.module.QUIET_PG_SETTINGS} self.assertEqual(len(self.sql_trace), 2 * len(settings)) for sql, trace in self.sql_trace: self.assertTrue(sql.startswith('ALTER ROLE "truf" IN DATABASE "truf" ')) self.assertNotIn(PASSWORD, sql) if ' RESET ' in sql: self.assertIn('migrate', trace) self.assertIn(sql.rsplit(' ', 1)[1], settings) else: self.assertNotIn('migrate', trace) self.assertIn(' SET ', sql) self.security.require_trusted_native_executable.assert_called_once_with('/usr/lib/postgresql/16/bin/pg_restore') def test_raw_equality_requires_exact_table_set_only_counts_and_sequences(self): m = self.module m._online = mock.Mock(return_value=160014) conn = mock.Mock() def execute(sql, *args): if "c.relkind IN ('r','p','f')" in sql: return SimpleNamespace(fetchall=lambda: [{'schema_name': 'public', 'name': 'target_queue', 'kind': 'r'}]) if 'count(*) AS count FROM ONLY' in sql: return SimpleNamespace(fetchone=lambda: {'count': 2}) if "c.relkind = 'S'" in sql: return SimpleNamespace(fetchall=lambda: [{'schema_name': 'public', 'name': 'target_queue_id_seq'}]) if 'SELECT last_value, is_called' in sql: return SimpleNamespace(fetchone=lambda: {'last_value': 7, 'is_called': True}) self.fail('unexpected fixture SQL') conn.execute.side_effect = execute result = m._database_equality(self.runtime, conn, self.manifest['database'], self.identity, self.progress) self.assertEqual((result['tables'], result['rows'], result['sequences_verified']), (1, 2, 1)) conn.execute.assert_any_call('SELECT count(*) AS count FROM ONLY "public"."target_queue"') for mutation in ('missing', 'count', 'sequence'): database = copy.deepcopy(self.manifest['database']) if mutation == 'missing': database['table_counts']['extra'] = 0 elif mutation == 'count': database['table_counts']['target_queue'] = 3 else: database['sequence_states']['public']['target_queue_id_seq']['is_called'] = False with self.subTest(mutation=mutation), self.assertRaises(m.Failure): m._database_equality(self.runtime, conn, database, self.identity, self.progress) def test_restore_rejects_nonempty_schema_and_foreign_linux_identity(self): m = self.module m._online = mock.Mock(return_value=160014) conn = mock.Mock() conn.execute.return_value.fetchone.return_value = {'count': 1} with self.assertRaises(m.Failure): m._database_equality(self.runtime, conn, self.manifest['database'], self.identity, self.progress, empty=True) for key, value in (('system_identifier', '1111111111'), ('pg_major', 15), ('data_directory', r'S:\postgres-data'), ('user', 'windows_user'), ('port', 15432)): with self.subTest(key=key), self.assertRaises(m.Failure): m._identity(self.runtime, dict(self.identity, **{key: value}), self.manifest['database']) def postman_fixture(self): content = b'{"synthetic": "local artifact"}\r\n' digest = sha(content) relative = 'runtime-linux/postman_cache/' + digest + '.json' (self.runtime.DATA / relative).write_bytes(content) cache = {relative: {'path': relative, 'size': len(content), 'sha256': digest}} target = {'sha256': digest, 'size': len(content), 'origin': {'do_not_rewrite': 'D:\\history'}, 'cache_path': 'D:\\truf\\runtime\\postman_cache\\' + digest + '.json'} return target, 'postman:sha256:' + digest, cache def test_postman_rebases_only_paths_preserving_semantic_identity_and_metadata(self): target, identity, cache = self.postman_fixture() result = json.loads(self.module._postman_target(self.runtime, json.dumps(target), identity, cache)) self.assertEqual(result['origin'], target['origin']) self.assertEqual(result['sha256'], target['sha256']) self.assertEqual(result['cache_path'], str(self.runtime.DATA / next(iter(cache)))) self.assertEqual(sys.modules['target_identity'].postman_target_identity(result), identity) def test_postman_path_identity_escape_size_and_hash_fail_closed(self): target, identity, cache = self.postman_fixture() for key, value in (('cache_path', 'D:\\truf\\runtime\\postman_cache\\..\\other.json'), ('cache_path', 'S:\\outside.json'), ('size', target['size'] + 1), ('sha256', '0' * 64)): with self.subTest(key=key), self.assertRaises(self.module.Failure): self.module._postman_target(self.runtime, json.dumps(dict(target, **{key: value})), identity, cache) with self.assertRaises(self.module.Failure): self.module._postman_target(self.runtime, json.dumps(target), 'postman:file:D:\\historical', cache) def test_postman_bound_pending_target_rolls_back_all_adjustments(self): m = self.module target, normalized, _cache = self.postman_fixture() row = {name: None for name in m.FENCES} row.update(id=1, platform='postman', target=json.dumps(target), normalized_target=normalized, status='pending', resolver_state=None) bad = dict(row, id=2, current_result_reservation_id=99) conn = mock.Mock() pages = iter(([row, bad], [])) def execute(sql, *args): if sql.startswith('SELECT'): return SimpleNamespace(fetchall=lambda: next(pages)) self.assertIn('SET target = %s', sql) self.assertNotIn('SET normalized_target', sql) self.assertNotIn('updated_at', sql) return SimpleNamespace(rowcount=1) conn.execute.side_effect = execute m._postman_target = mock.Mock(return_value='{"rebased":true}') with self.assertRaises(m.Failure) as caught: m._rebase_postman(self.runtime, conn, {}) self.assertEqual((caught.exception.code, caught.exception.count), (2, 1)) conn.rollback.assert_called_once() conn.commit.assert_not_called() self.assertIn("status IN ('pending','deferred','in_progress')", conn.execute.call_args_list[0].args[0]) def test_both_alternate_cached_platforms_reject_windows_locators_with_counts_only(self): m = self.module target, _normalized, cache = self.postman_fixture() text = json.dumps(target) rows = [] for index, platform in enumerate(('github_gists', 'github_archive_files'), 1): row = {name: None for name in m.FENCES} row.update(id=index, platform=platform, target=text, normalized_target=text.strip().lower(), status='pending', resolver_state=None) rows.append(row) before = copy.deepcopy(rows) conn = mock.Mock() pages = iter((rows, [])) conn.execute.side_effect = lambda sql, params: SimpleNamespace(fetchall=lambda: next(pages)) output = io.StringIO() with contextlib.redirect_stdout(output), contextlib.redirect_stderr(output), self.assertRaises(m.Failure) as caught: m._rebase_postman(self.runtime, conn, cache) self.assertEqual((caught.exception.code, caught.exception.count), (2, 2)) self.assertEqual(output.getvalue(), '') self.assertEqual(str(caught.exception), '2') self.assertEqual(rows, before) for call in conn.execute.call_args_list: self.assertTrue(call.args[0].startswith('SELECT')) self.assertIn("'github_gists','github_archive_files'", call.args[0]) conn.rollback.assert_called_once() conn.commit.assert_not_called() def test_portable_alternate_cached_targets_are_validated_without_rewriting_identity(self): m = self.module target, normalized, _cache = self.postman_fixture() target['cache_path'] = '/data/runtime-linux/postman_cache/' + target['sha256'] + '.json' text = json.dumps(target) m._postman_target = mock.Mock(return_value=text) for platform in ('github_gists', 'github_archive_files'): row = {name: None for name in m.FENCES} row.update(id=1, platform=platform, target=text, normalized_target=text.lower(), status='deferred', resolver_state=None) conn = mock.Mock() pages = iter(([row], [])) conn.execute.side_effect = lambda sql, params: SimpleNamespace(fetchall=lambda: next(pages)) with self.subTest(platform=platform): result = m._rebase_postman(self.runtime, conn, {}) self.assertEqual(result, {'postman_reviewed': 1, 'postman_adjusted': 0}) m._postman_target.assert_called_with(self.runtime, text, normalized, {}) self.assertEqual(row['normalized_target'], text.lower()) self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_alternate_cached_platform_cannot_use_postman_queue_identity(self): m = self.module target, normalized, _cache = self.postman_fixture() m._postman_target = mock.Mock() for platform in ('github_gists', 'github_archive_files'): row = {name: None for name in m.FENCES} row.update(id=1, platform=platform, target=json.dumps(target), normalized_target=normalized, status='pending', resolver_state=None) conn = mock.Mock() pages = iter(([row], [])) conn.execute.side_effect = lambda sql, params: SimpleNamespace(fetchall=lambda: next(pages)) with self.subTest(platform=platform), self.assertRaises(m.Failure) as caught: m._rebase_postman(self.runtime, conn, {}) self.assertEqual((caught.exception.code, caught.exception.count), (2, 1)) conn.rollback.assert_called_once() m._postman_target.assert_not_called() def projection_fixture(self): providers = ('openai', 'huggingface', 'gemini', 'gcp', 'anthropic', 'azure', 'openrouter', 'zai', 'qwen', 'kimi', 'groq', 'replicate', 'xai', 'deepseek', 'cohere', 'mistral', 'perplexity', 'together') streams = [ ('scan_results', 'results', 'scan_results.jsonl'), ('found_secrets', 'results', 'found_secrets.jsonl'), ('scan_errors', 'results', 'scan_errors.log'), ] streams.extend((f'keycheck:{provider}:results', 'keychecks', f'{provider}/{provider}Results.jsonl') for provider in providers) streams.extend((f'keycheck:{provider}:status', 'keychecks', f'{provider}/{provider}Checked.txt') for provider in providers[:13]) rows, files = [], {} for index, (stream, root, relative) in enumerate(streams, 1): name = 'runtime-linux/' + root + '/' + relative payload = ('synthetic projection ' + str(index) + '\r\n').encode('ascii') path = self.runtime.DATA / name path.parent.mkdir(mode=0o700, exist_ok=True) path.write_bytes(payload) files[name] = {'size': len(payload), 'sha256': sha(payload)} status = stream.endswith(':status') generation = 0 if status else index + 7 offset = len(payload) + (7 if index % 2 else -7) if status else len(payload) rows.append({'stream_name': stream, 'cursor_stream_name': stream, 'base_relative_path': relative, 'committed_offset': offset, 'generation': generation, 'current_generation': generation, 'last_append_id': index + 100}) return rows, files def test_projection_accepts_34_streams_without_changing_historical_cursors(self): rows, files = self.projection_fixture() self.assertEqual(len(rows), 34) self.assertEqual(sum(row['stream_name'].endswith(':results') for row in rows), 18) statuses = [row for row in rows if row['stream_name'].endswith(':status')] self.assertEqual(len(statuses), 13) differences = [row['committed_offset'] - files['runtime-linux/keychecks/' + row['base_relative_path']]['size'] for row in statuses] self.assertEqual(set(differences), {-7, 7}) before = copy.deepcopy(rows) before_files = copy.deepcopy(files) fingerprints = {name: self.module._regular(self.runtime.DATA / name, self.runtime) for name in files} conn = mock.Mock() conn.execute.return_value.fetchall.return_value = rows self.module._projection_files(self.runtime, conn, files) self.assertEqual(rows, before) self.assertEqual(files, before_files) self.assertEqual(fingerprints, {name: self.module._regular(self.runtime.DATA / name, self.runtime) for name in files}) conn.execute.assert_called_once() sql = ' '.join(conn.execute.call_args.args[0].split()) self.assertTrue(sql.startswith('SELECT')) self.assertIn('FULL JOIN public.projection_cursors c ON c.stream_name = s.stream_name', sql) self.assertIn('c.stream_name AS cursor_stream_name', sql) self.assertIn('s.current_generation, c.generation, c.committed_offset', sql) self.assertNotIn('*', sql) self.assertNotIn('stream_kind', sql) conn.commit.assert_called_once() self.runtime.private_path.assert_any_call(self.runtime.DATA / 'runtime-linux/keychecks/openai/openaiResults.jsonl') self.runtime.private_path.assert_any_call(self.runtime.DATA / 'runtime-linux/keychecks/huggingface/huggingfaceResults.jsonl') for row in statuses: self.runtime.private_path.assert_any_call(self.runtime.DATA / 'runtime-linux/keychecks' / row['base_relative_path']) def test_keycheck_projection_rejects_wrong_service_paths_and_unknown_streams(self): rows, files = self.projection_fixture() for key, value in ( ('base_relative_path', 'huggingface/huggingfaceResults.jsonl'), ('base_relative_path', 'openai/openairesults.jsonl'), ('base_relative_path', '../openai/openaiResults.jsonl'), ('stream_name', 'keycheck:openai:unknown'), ): invalid = copy.deepcopy(rows) invalid[3][key] = value conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(key=key, value=value), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_projection_rejects_noncanonical_scan_and_status_paths(self): rows, files = self.projection_fixture() for index, relative in ( (0, 'unknown.jsonl'), (1, 'scan_results.jsonl'), (2, 'scan_errors.jsonl'), (21, 'openai/openaiResults.jsonl'), (21, 'openai/openaiChecked.json'), (21, 'openai/openaichecked.txt'), (21, 'gemini/geminiChecked.txt'), (21, 'openai/unknown.txt'), (21, '../openai/openaiChecked.txt'), (21, 'openai\\openaiChecked.txt'), ): invalid = copy.deepcopy(rows) invalid[index]['base_relative_path'] = relative conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(index=index, relative=relative), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() def test_projection_extra_kind_columns_cannot_bypass_validation(self): rows, files = self.projection_fixture() for index, changes in ( (3, {'committed_offset': 0}), (21, {'stream_name': 'keycheck:openai:snapshot', 'cursor_stream_name': 'keycheck:openai:snapshot'}), (21, {'stream_name': 'keycheck:OpenAI:status', 'cursor_stream_name': 'keycheck:OpenAI:status'}), (21, {'cursor_stream_name': None}), ): invalid = copy.deepcopy(rows) invalid[index].update(changes, stream_kind='status_snapshot', kind='status') conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(index=index, changes=changes), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_projection_requires_complete_unique_stream_cursor_coverage(self): rows, files = self.projection_fixture() cases = [('empty', []), ('missing scan', rows[1:]), ('duplicate', rows + [rows[-1]])] for changes in ( {'cursor_stream_name': None, 'generation': None, 'committed_offset': None}, {'stream_name': None, 'base_relative_path': None, 'current_generation': None}, {'cursor_stream_name': 'keycheck:other:status'}, ): invalid = copy.deepcopy(rows) invalid[-1].update(changes) cases.append((str(changes), invalid)) missing_column = copy.deepcopy(rows) del missing_column[-1]['cursor_stream_name'] cases.append(('missing cursor identity column', missing_column)) for label, invalid in cases: conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(case=label), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_projection_requires_matching_generations_and_nonnegative_integers(self): rows, files = self.projection_fixture() for index in (0, 3, 21): for field in ('generation', 'current_generation', 'committed_offset'): for value in (-1, None, '0', False, 1.5, 2 ** 63): invalid = copy.deepcopy(rows) invalid[index][field] = value conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(index=index, field=field, value=value), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() invalid = copy.deepcopy(rows) invalid[index]['generation'] += 1 conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(index=index, mismatch=True), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() def test_status_projection_requires_manifested_file_size_and_hash(self): rows, files = self.projection_fixture() name = 'runtime-linux/keychecks/' + rows[21]['base_relative_path'] path = self.runtime.DATA / name payload = path.read_bytes() for case in ('unmanifested', 'unmanifested zero cursor', 'unmanifested absent zero cursor', 'missing', 'longer', 'shorter', 'same size changed bytes', 'wrong manifest hash', 'wrong manifest size'): path.write_bytes(payload) invalid_rows, invalid_files = copy.deepcopy(rows), copy.deepcopy(files) if case.startswith('unmanifested'): del invalid_files[name] if 'zero cursor' in case: invalid_rows[21]['committed_offset'] = 0 if 'absent' in case: path.unlink() elif case == 'missing': path.unlink() elif case == 'longer': path.write_bytes(payload + b'x') elif case == 'shorter': path.write_bytes(payload[:-1]) elif case == 'same size changed bytes': path.write_bytes(b'X' + payload[1:]) elif case == 'wrong manifest hash': invalid_files[name]['sha256'] = '0' * 64 else: invalid_files[name]['size'] += 1 before = copy.deepcopy(invalid_rows) conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid_rows with self.subTest(case=case), self.assertRaises((self.module.Failure, OSError)): self.module._projection_files(self.runtime, conn, invalid_files) conn.commit.assert_not_called() self.assertEqual(invalid_rows, before) self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_status_projection_accepts_empty_snapshot_or_zero_historical_cursor(self): rows, files = self.projection_fixture() name = 'runtime-linux/keychecks/' + rows[21]['base_relative_path'] path = self.runtime.DATA / name payload = path.read_bytes() for content, offset in ((b'', rows[21]['committed_offset']), (payload, 0)): path.write_bytes(content) files[name] = {'size': len(content), 'sha256': sha(content)} rows[21]['committed_offset'] = offset before = copy.deepcopy(rows) conn = mock.Mock() conn.execute.return_value.fetchall.return_value = rows with self.subTest(size=len(content), offset=offset): self.module._projection_files(self.runtime, conn, files) self.assertEqual(rows, before) conn.commit.assert_called_once() self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_status_projection_rejects_nonprivate_and_nonregular_files(self): rows, files = self.projection_fixture() path = self.runtime.DATA / 'runtime-linux/keychecks' / rows[21]['base_relative_path'] conn = mock.Mock() conn.execute.return_value.fetchall.return_value = rows private = self.runtime.private_path.side_effect def denied(value, **kwargs): if Path(value) == path: raise self.module.Failure() return private(value, **kwargs) with mock.patch.object(self.runtime, 'private_path', side_effect=denied), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() path.unlink() path.mkdir(mode=0o700) with mock.patch.object(self.runtime, 'private_path', side_effect=lambda value, **kwargs: value), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, files) conn.commit.assert_not_called() def test_keycheck_projection_requires_exact_manifest_file_and_offset(self): rows, files = self.projection_fixture() name = 'runtime-linux/keychecks/openai/openaiResults.jsonl' for present, size, offset in ((False, 0, rows[3]['committed_offset']), (True, files[name]['size'] + 1, rows[3]['committed_offset']), (False, 0, 0)): invalid_files, invalid_rows = copy.deepcopy(files), copy.deepcopy(rows) if present: invalid_files[name]['size'] = size else: del invalid_files[name] invalid_rows[3]['committed_offset'] = offset conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid_rows with self.subTest(present=present, offset=offset), self.assertRaises(self.module.Failure): self.module._projection_files(self.runtime, conn, invalid_files) conn.commit.assert_not_called() def test_unwritten_keycheck_stream_can_have_zero_cursor_and_no_file(self): rows, files = self.projection_fixture() rows.append({'stream_name': 'keycheck:new_service:results', 'cursor_stream_name': 'keycheck:new_service:results', 'base_relative_path': 'new_service/new_serviceResults.jsonl', 'committed_offset': 0, 'generation': 0, 'current_generation': 0}) conn = mock.Mock() conn.execute.return_value.fetchall.return_value = rows self.module._projection_files(self.runtime, conn, files) self.assertFalse((self.runtime.DATA / 'runtime-linux/keychecks/new_service').exists()) self.assertEqual(rows[-1]['generation'], 0) def test_projection_cursor_mismatch_never_gets_implicitly_reset(self): rows, files = self.projection_fixture() for index, row in enumerate(rows): if row['stream_name'].endswith(':status'): continue invalid = copy.deepcopy(rows) invalid[index]['committed_offset'] = 0 before = copy.deepcopy(invalid) conn = mock.Mock() conn.execute.return_value.fetchall.return_value = invalid with self.subTest(stream=row['stream_name']), self.assertRaises(self.module.Failure) as caught: self.module._projection_files(self.runtime, conn, files) self.assertEqual(caught.exception.code, 2) self.assertEqual(invalid, before) conn.commit.assert_not_called() self.assertTrue(all(call.args[0].startswith('SELECT') for call in conn.execute.call_args_list)) def test_preservation_hashes_existing_quarantine_rows_without_values(self): m = self.module m.PRESERVED = {'pipeline_quarantine': ('id', None)} conn, cursor = mock.MagicMock(), mock.MagicMock() conn.execute.return_value.fetchall.return_value = [{'name': 'id'}, {'name': 'review_status'}] conn.execute.return_value.fetchone.return_value = {'cutoff': 7} conn.cursor.return_value.__enter__.return_value = cursor cursor.__iter__.return_value = iter([{'digest': 'a' * 64}, {'digest': 'b' * 64}]) evidence = m._preserved(self.runtime, conn) self.assertEqual(evidence['pipeline_quarantine']['count'], 2) self.assertEqual(evidence['pipeline_quarantine']['sha256'], sha(('a' * 64 + 'b' * 64).encode('ascii'))) self.assertIn('pg_catalog.sha256', cursor.execute.call_args.args[0]) self.assertIn('"id" <= %s', cursor.execute.call_args.args[0]) self.assertEqual(cursor.execute.call_args.args[1], (7,)) cursor.__iter__.return_value = iter([{'digest': 'a' * 64}, {'digest': 'c' * 64}]) with self.assertRaises(m.Failure) as caught: m._preserved(self.runtime, conn, evidence) self.assertEqual((caught.exception.code, caught.exception.count), (2, 1)) def test_operations_authority_tables_are_preserved_across_import_migration(self): self.assertEqual(self.module.PRESERVED['runtime_operations'], ('operation_id', None)) self.assertEqual(self.module.PRESERVED['runtime_operations_control'], ('id', None)) self.assertEqual(self.module.PRESERVED['runtime_audit_events'], ('id', None)) def test_preservation_allows_additive_schema_without_inventing_old_evidence(self): m = self.module m.PRESERVED = {'new_audit_table': ('id', None)} conn = mock.Mock() conn.execute.return_value.fetchall.return_value = [] self.assertEqual(m._preserved(self.runtime, conn), {}) conn.execute.reset_mock() self.assertEqual(m._preserved(self.runtime, conn, {}), {}) conn.execute.assert_not_called() def test_uncertain_stop_retries_without_cancellation_unwind(self): m = self.module self.runtime._shutdown_requested = True pg = SimpleNamespace(ProbeKind=SimpleNamespace(STOPPED='stopped'), maintenance_stop=mock.Mock()) pg.maintenance_stop.side_effect = [RuntimeError(PRIVATE_VALUE), SimpleNamespace(completed=True, stopped=False), SimpleNamespace(completed=True, stopped=True)] backend = mock.Mock() backend.probe.return_value.kind = 'stopped' m._backend_stop(pg, {}, backend, self.progress) self.assertEqual(pg.maintenance_stop.call_count, 3) self.assertEqual(m.time.sleep.call_count, 2) backend.close.assert_not_called() self.assertEqual([call.kwargs['diagnostic'][:4] for call in self.progress.call_args_list], [(3, 1, 0, 6), (4, 1, 0, 1)]) self.assertEqual([call.args for call in self.progress.call_args_list], [(12, 1, 0), (12, 2, 0)]) def test_diagnostic_uses_only_known_type_ids_codes_counts_and_local_lines(self): m = self.module render = mock.Mock(side_effect=AssertionError('exception formatting is forbidden')) unknown = type(PRIVATE_VALUE, (Exception,), {'__str__': render, '__repr__': render})(PASSWORD) unknown.add_note(PRIVATE_VALUE) cases = ( (unknown, (1, 0, 0)), (m.Failure(2, 7), (2, 7, 1)), (m.Failure(PASSWORD, PRIVATE_VALUE), (1, 0, 1)), (m.Failure(True, False), (1, 0, 1)), (m.Failure(999, -1), (1, 0, 1)), (m.Failure(124, 2 ** 63), (124, 0, 1)), (KeyboardInterrupt(PASSWORD), (130, 0, 2)), (OSError(PASSWORD), (1, 0, 3)), (ValueError(PASSWORD), (1, 0, 4)), (TypeError(PASSWORD), (1, 0, 5)), (RuntimeError(PASSWORD), (1, 0, 6)), (SystemExit(PASSWORD), (1, 0, 0)), ) for index, (error, expected) in enumerate(cases): progress = mock.Mock() with self.subTest(case=index): m._diagnostic(progress, 3, error, 4) progress.assert_called_once_with(12, 4, 0, cleanup=True, diagnostic=(3, *expected, 0)) render.assert_not_called() try: m._integer(PRIVATE_VALUE) except m.Failure as error: m._diagnostic(self.progress, 1, error) line = self.progress.call_args.kwargs['diagnostic'][-1] self.assertGreater(line, m._integer.__code__.co_firstlineno) self.assertLess(line, m._sha.__code__.co_firstlineno) def test_public_first_error_is_durable_and_visible_before_repeated_stop_hold(self): m, r = self.module, self.runtime authority, restore = m._authority, m._restore self.fake_public_database() self.fake_restore_dependencies() m._authority, m._restore = authority, restore pg, backend, lock = self.fake_authority_dependencies() original = type(PRIVATE_VALUE, (Exception,), {})(PASSWORD, 'SELECT private_payload', str(self.root)) original.add_note('private-config: ' + PRIVATE_VALUE) m._database_equality.side_effect = original stdout, stderr = io.StringIO(), io.StringIO() stderr.flush = mock.Mock(wraps=stderr.flush) failure_path, hold_path = (r.DATA / ('config/windows-import-' + name + '.json') for name in ('failure', 'hold')) observations = [] stopped = pg.maintenance_stop.side_effect def stop(config, backend): observations.append({ 'failure': failure_path.read_bytes() if failure_path.exists() else None, 'hold': hold_path.read_bytes() if hold_path.exists() else None, 'stderr': stderr.getvalue(), 'flushes': stderr.flush.call_count, 'initialize_locked': self.locked, 'released': lock.release.call_count, 'closed': backend.close.call_count, 'constructed': pg.PostgresBackend.call_count, 'reported': (r.DATA / 'config/windows-import-report.json').exists(), 'marked': r.INITIALIZED.exists(), }) r._shutdown_requested = True if len(observations) == 1: raise OSError(PRIVATE_VALUE + PASSWORD) if len(observations) == 2: raise KeyboardInterrupt(PRIVATE_VALUE) if len(observations) == 3: return SimpleNamespace(completed=True, stopped=False, detail=PRIVATE_VALUE + PASSWORD) return stopped(config, backend) pg.maintenance_stop.side_effect = stop with mock.patch.object(m, '_write', wraps=m._write) as write, \ mock.patch.object(m.os, 'open', wraps=m.os.open) as opened, \ contextlib.redirect_stdout(stdout), contextlib.redirect_stderr(stderr): code = m.import_snapshot(r, self.manifest_sha) self.assertEqual(code, 1) fields = ('phase', 'stage', 'code', 'review_count', 'type_id', 'line', 'attempt') events = [] for text in stderr.getvalue().splitlines(): if text.startswith('import-snapshot-diagnostic '): self.assertRegex(text, r'^import-snapshot-diagnostic(?: [0-9]+){7}$') events.append(dict(zip(fields, map(int, text.split()[1:])))) else: self.assertEqual(text, 'import-snapshot 7 1 0') self.assertEqual([event['stage'] for event in events], [1, 3, 3, 4]) self.assertEqual([event['type_id'] for event in events], [0, 3, 2, 1]) self.assertEqual([event['code'] for event in events], [1, 1, 130, 1]) self.assertEqual([event['attempt'] for event in events], [0, 1, 2, 3]) self.assertTrue(all(event['phase'] == 7 and event['review_count'] == 0 for event in events)) self.assertGreater(events[0]['line'], restore.__code__.co_firstlineno) self.assertLess(events[0]['line'], m.import_snapshot.__code__.co_firstlineno) self.assertEqual(len(observations), 4) for index, observed in enumerate(observations): self.assertEqual(json.loads(observed['failure']), events[0]) self.assertIn('import-snapshot-diagnostic 7 1 1 0 0 ', observed['stderr']) self.assertGreater(observed['flushes'], 0) self.assertEqual((observed['initialize_locked'], observed['released'], observed['closed'], observed['constructed'], observed['reported'], observed['marked']), (True, 0, 0, 1, False, False)) if index: self.assertEqual(json.loads(observed['hold']), events[1]) else: self.assertIsNone(observed['hold']) for path, event in ((failure_path, events[0]), (hold_path, events[1])): writes = [call for call in write.call_args_list if call.args[1] == path] self.assertEqual(len(writes), 1) self.assertEqual(json.loads(path.read_bytes()), event) opened.assert_any_call(path, m.os.O_WRONLY | m.os.O_CREAT | m.os.O_EXCL | m.os.O_NOFOLLOW | getattr(m.os, 'O_BINARY', 0), 0o600) r.private_path.assert_any_call(path) m._fsync_dir.assert_any_call(r.DATA / 'config') report_bytes = (r.DATA / 'config/windows-import-report.json').read_bytes() report = json.loads(report_bytes) self.assertEqual((report['status'], report['phase'], report['code']), ('failed-unmarked', 7, 1)) self.assertEqual((report['failure'], report['hold']), (events[0], events[1])) output = stdout.getvalue() + stderr.getvalue() + report_bytes.decode('ascii') output += failure_path.read_text(encoding='ascii') + hold_path.read_text(encoding='ascii') for value in (PRIVATE_VALUE, PASSWORD, 'SELECT private_payload', str(self.root), 'private-config'): self.assertNotIn(value, output) for text in stdout.getvalue().splitlines(): self.assertRegex(text, r'^import-snapshot(?: [0-9]+){3}$') pg.PostgresBackend.assert_called_once_with({}) self.assertEqual(pg.maintenance_stop.call_args_list, [mock.call({}, backend=backend)] * 4) backend.close.assert_called_once_with() lock.release.assert_called_once_with() self.assertLess(self.trace.index('confirmed-stop'), self.trace.index('backend-close')) self.assertLess(self.trace.index('backend-close'), self.trace.index('cluster-release')) self.assertLess(self.trace.index('cluster-release'), self.trace.index('initialize-unlock')) self.security.write_private_json_exclusive.assert_not_called() m.subprocess.Popen.assert_not_called() def test_successful_stop_result_does_not_hide_final_probe_failure_or_release_authority(self): m = self.module pg, backend, lock = self.fake_authority_dependencies() probes = [] def probe(): probes.append((lock.release.call_count, backend.close.call_count, pg.maintenance_stop.call_count)) if len(probes) == 1: return SimpleNamespace(kind='ready') if len(probes) == 2: raise OSError(PRIVATE_VALUE + PASSWORD) return SimpleNamespace(kind='foreign' if len(probes) == 3 else 'stopped', detail=PRIVATE_VALUE) backend.probe.side_effect = probe original = m.Failure(2, 9) with self.assertRaises(m.Failure) as caught: with m._authority(self.runtime, {}, self.manifest['database'], self.progress): self.runtime._shutdown_requested = True raise original self.assertIs(caught.exception, original) self.assertEqual(probes, [(0, 0, 0), (0, 0, 1), (0, 0, 2), (0, 0, 3)]) diagnostics = [call.kwargs['diagnostic'] for call in self.progress.call_args_list] self.assertEqual([value[:4] for value in diagnostics], [(1, 2, 9, 1), (5, 1, 0, 3), (6, 1, 0, 1)]) self.assertEqual([call.args for call in self.progress.call_args_list], [(12, 0, 0), (12, 1, 0), (12, 2, 0)]) pg.PostgresBackend.assert_called_once_with({}) self.assertEqual(pg.maintenance_stop.call_args_list, [mock.call({}, backend=backend)] * 3) backend.close.assert_called_once_with() lock.release.assert_called_once_with() self.assertEqual(self.trace[-2:], ['backend-close', 'cluster-release']) def test_close_hold_and_broken_diagnostics_cannot_replace_failure_or_abandon_backend(self): m = self.module pg, backend, lock = self.fake_authority_dependencies() original = m.Failure(2, 9) stopped, stops, closes = pg.maintenance_stop.side_effect, [], [] self.progress.side_effect = BrokenPipeError(PRIVATE_VALUE) m.time.sleep.side_effect = KeyboardInterrupt(PRIVATE_VALUE) def stop(config, backend): stops.append((lock.release.call_count, backend.close.call_count)) if len(stops) == 1: raise KeyboardInterrupt(PRIVATE_VALUE) return stopped(config, backend) def close(): closes.append((lock.release.call_count, self.stopped)) if len(closes) == 1: raise OSError(PRIVATE_VALUE) if len(closes) == 2: raise KeyboardInterrupt(PRIVATE_VALUE) self.trace.append('backend-close') pg.maintenance_stop.side_effect, backend.close.side_effect = stop, close with self.assertRaises(m.Failure) as caught: with m._authority(self.runtime, {}, self.manifest['database'], self.progress): self.runtime._shutdown_requested = True raise original self.assertIs(caught.exception, original) self.assertEqual(stops, [(0, 0), (0, 0), (0, 1), (0, 2)]) self.assertEqual(closes, [(0, True)] * 3) self.assertEqual([call.kwargs['diagnostic'][:4] for call in self.progress.call_args_list], [(1, 2, 9, 1), (3, 130, 0, 2), (7, 1, 0, 3), (7, 130, 0, 2)]) pg.PostgresBackend.assert_called_once_with({}) self.assertEqual(pg.maintenance_stop.call_args_list, [mock.call({}, backend=backend)] * 4) lock.release.assert_called_once_with() self.assertEqual(self.trace[-2:], ['backend-close', 'cluster-release']) def test_backend_construction_hold_preserves_initial_failure_and_cluster_lock(self): m = self.module pg, backend, lock = self.fake_authority_dependencies() original, constructions = m.Failure(2, 3), [] def construct(config): constructions.append((lock.release.call_count, len(self.progress.call_args_list))) if len(constructions) == 1: raise original if len(constructions) < 4: raise OSError(PRIVATE_VALUE) return backend pg.PostgresBackend.side_effect = construct with self.assertRaises(m.Failure) as caught: with m._authority(self.runtime, {}, self.manifest['database'], self.progress): self.fail('construction failure must not enter the operation') self.assertIs(caught.exception, original) self.assertEqual(constructions, [(0, 0), (0, 1), (0, 2), (0, 3)]) self.assertEqual([call.kwargs['diagnostic'][:4] for call in self.progress.call_args_list], [(1, 2, 3, 1), (2, 1, 0, 3), (2, 1, 0, 3)]) pg.maintenance_stop.assert_called_once_with({}, backend=backend) backend.close.assert_called_once_with() lock.release.assert_called_once_with() def test_failed_private_and_stream_diagnostics_still_retain_authority_and_first_code(self): m, r = self.module, self.runtime authority = m._authority self.fake_public_database() m._authority = authority pg, backend, lock = self.fake_authority_dependencies() writes, stops = [], [] write, stopped = m._write, pg.maintenance_stop.side_effect def failing_write(runtime, path, payload): if path.name in ('windows-import-failure.json', 'windows-import-hold.json'): writes.append((path.name, json.loads(payload), lock.release.call_count)) raise OSError(PRIVATE_VALUE + PASSWORD) return write(runtime, path, payload) def restore(runtime, path, config, manifest, files, dump, fingerprint, progress, report): progress(7) with m._authority(runtime, config, manifest['database'], progress): r._shutdown_requested = True raise m.Failure(2, 17) def stop(config, backend): stops.append((self.locked, lock.release.call_count, backend.close.call_count)) if len(stops) == 1: raise OSError(PRIVATE_VALUE + PASSWORD) return stopped(config, backend) class BrokenOutput(io.StringIO): def write(self, text): if r._shutdown_requested: raise KeyboardInterrupt(PRIVATE_VALUE) return super().write(text) class BrokenDiagnostics(io.StringIO): def write(self, text): if text == 'import-snapshot-diagnostic': raise BrokenPipeError(PRIVATE_VALUE + PASSWORD) return super().write(text) m._write, m._restore.side_effect, pg.maintenance_stop.side_effect = failing_write, restore, stop stdout, stderr = BrokenOutput(), BrokenDiagnostics() with contextlib.redirect_stdout(stdout), contextlib.redirect_stderr(stderr): code = m.import_snapshot(r, self.manifest_sha) self.assertEqual(code, 2) self.assertEqual(stops, [(True, 0, 0)] * 2) self.assertEqual([name for name, _event, _released in writes], ['windows-import-failure.json', 'windows-import-hold.json']) self.assertTrue(all(released == 0 for _name, _event, released in writes)) report_bytes = (r.DATA / 'config/windows-import-report.json').read_bytes() report = json.loads(report_bytes) self.assertEqual((report['code'], report['review_count']), (2, 17)) self.assertEqual((report['failure'], report['hold']), (writes[0][1], writes[1][1])) self.assertEqual(stderr.getvalue(), 'import-snapshot 7 2 17\n') for value in (PRIVATE_VALUE, PASSWORD): self.assertNotIn(value, stdout.getvalue() + stderr.getvalue() + report_bytes.decode('ascii')) pg.PostgresBackend.assert_called_once_with({}) backend.close.assert_called_once_with() lock.release.assert_called_once_with() self.assertEqual(self.trace[-3:], ['backend-close', 'cluster-release', 'initialize-unlock']) self.security.write_private_json_exclusive.assert_not_called() m.subprocess.Popen.assert_not_called() def test_authority_failure_stops_before_backend_close_and_unlock(self): m = self.module lock = mock.Mock() lock.acquire.side_effect = lambda: self.trace.append('cluster-lock') lock.release.side_effect = lambda: self.trace.append('cluster-release') self.security.ClusterAuthorityLock = mock.Mock(return_value=lock) backend = mock.Mock() backend.probe.return_value.kind = 'ready' backend.close.side_effect = lambda: self.trace.append('backend-close') def stop(config, backend): self.trace.append('confirmed-stop') backend.probe.return_value.kind = 'stopped' return SimpleNamespace(completed=True, stopped=True) pg = SimpleNamespace(PostgresBackend=mock.Mock(return_value=backend), ProbeKind=SimpleNamespace(READY='ready', STOPPED='stopped'), verify_cluster_identity=mock.Mock(return_value=self.identity), maintenance_stop=stop) sys.modules['postgres_runtime'] = pg with self.assertRaises(m.Failure): with m._authority(self.runtime, {}, self.manifest['database'], self.progress): raise m.Failure() self.assertEqual(self.trace, ['cluster-lock', 'confirmed-stop', 'backend-close', 'cluster-release']) def test_stop_cli_retries_even_when_shutdown_is_requested(self): m = self.module self.runtime._shutdown_requested = True m._command = mock.Mock(side_effect=[m.Failure(), None]) m._stop_cli(self.runtime, self.runtime.DATA / 'config/windows-import.yaml', self.progress) self.assertEqual(m._command.call_count, 2) for call in m._command.call_args_list: self.assertTrue(call.kwargs['stopping']) self.assertTrue(call.kwargs['lifecycle']) self.assertEqual(self.progress.call_args.kwargs['diagnostic'][:4], (8, 1, 0, 1)) def test_public_success_marker_last_after_confirmed_stop_with_private_evidence(self): self.fake_public_database() code, stdout, stderr = self.public() self.assertEqual((code, stderr), (0, '')) self.assertNotIn(PRIVATE_VALUE, stdout) self.assertNotIn(PASSWORD, stdout) self.assertLess(self.trace.index('maintenance-stop'), self.trace.index('marker')) self.assertLess(self.trace.index('stop-confirmed-under-authority'), self.trace.index('marker')) self.assertLess(self.trace.index('marker'), self.trace.index('initialize-unlock')) marker = json.loads(self.runtime.INITIALIZED.read_bytes()) self.assertEqual(marker['system_identifier'], self.identity['system_identifier']) self.assertEqual(marker['manifest_sha256'], self.manifest_sha) self.assertEqual(marker['format'], self.runtime.FORMAT) report = self.runtime.DATA / 'config/windows-import-report.json' self.assertEqual(marker['import_report_sha256'], sha(report.read_bytes())) self.assertEqual(json.loads(report.read_bytes())['status'], 'verified-stopped') self.assertEqual((self.runtime.DATA / 'config/windows-import-manifest.json').read_bytes(), self.manifest_bytes) self.module.subprocess.Popen.assert_not_called() def test_public_bad_archive_refuses_before_snapshot_mutations(self): self.fake_public_database() self.manifest['archive']['sha256'] = '0' * 64 self.save_manifest() code, _stdout, _stderr = self.public() self.assertEqual(code, 1) self.module._restore.assert_not_called() self.assertFalse((self.runtime.DATA / 'config/windows-import-manifest.json').exists()) self.assertEqual((self.runtime.DATA / 'config/secrets.yaml').read_bytes(), b'{}\n') def test_partial_import_is_unmarked_not_deleted_and_never_reused(self): self.fake_public_database() self.module._configuration.side_effect = RuntimeError(PRIVATE_VALUE + PASSWORD) code, stdout, stderr = self.public() self.assertEqual(code, 1) self.assertNotIn(PRIVATE_VALUE, stdout + stderr) self.assertNotIn(PASSWORD, stdout + stderr) self.assertFalse(self.runtime.INITIALIZED.exists()) self.assertEqual((self.runtime.DATA / 'config/secrets.yaml').read_bytes(), self.payloads['config/secrets.yaml']) report = self.runtime.DATA / 'config/windows-import-report.json' before = report.read_bytes() self.assertEqual(json.loads(before)['status'], 'failed-unmarked') code, _stdout, _stderr = self.public() self.assertEqual(code, 1) self.assertEqual(report.read_bytes(), before) self.module._restore.assert_not_called() def test_copied_quarantine_bytes_are_checked_again_before_marker(self): name = 'scanner-result-bundles/quarantine/held.bundle' self.payloads[name] = b'synthetic quarantined artifact' self.snapshot() self.fake_public_database() original = self.module._restore.side_effect def changed(*args, **kwargs): identity = original(*args, **kwargs) (self.runtime.DATA / name).write_bytes(b'changed quarantine bytes') return identity self.module._restore.side_effect = changed code, _stdout, _stderr = self.public() self.assertEqual(code, 1) self.assertFalse(self.runtime.INITIALIZED.exists()) self.assertEqual((self.runtime.DATA / name).read_bytes(), b'changed quarantine bytes') def test_cancellation_before_import_and_after_stop_never_publishes_marker(self): self.fake_public_database() self.runtime._shutdown_requested = True code, _stdout, _stderr = self.public() self.assertEqual(code, 130) self.assertEqual(self.trace, []) self.runtime._shutdown_requested = False @contextlib.contextmanager def stopped(*args, **kwargs): self.runtime._shutdown_requested = True yield self.identity self.module._authority.side_effect = stopped code, _stdout, _stderr = self.public() self.assertEqual(code, 130) self.assertFalse(self.runtime.INITIALIZED.exists()) self.security.write_private_json_exclusive.assert_not_called() def test_interrupt_handler_updates_callers_module_and_is_restored(self): self.fake_public_database() original = self.module._restore.side_effect def interrupt(*args, **kwargs): handler = self.module.signal.signal.call_args_list[0].args[1] handler(self.module.signal.SIGINT, None) return original(*args, **kwargs) self.module._restore.side_effect = interrupt code, _stdout, _stderr = self.public() self.assertEqual(code, 130) self.assertTrue(self.runtime._shutdown_requested) self.assertEqual(self.module.signal.signal.call_args.args, (2, 'fixture-prior-handler')) self.assertFalse(self.runtime.INITIALIZED.exists()) def test_changed_input_during_final_close_cannot_leave_initialized_marker(self): self.fake_public_database() original = self.module._unchanged visits = 0 def changed(path, handle, before, runtime=None): nonlocal visits if self.stopped and path == self.module.IMPORT / 'manifest.json': visits += 1 if visits == 2: raise self.module.Failure() return original(path, handle, before, runtime) self.module._unchanged = changed code, _stdout, _stderr = self.public() self.assertEqual(code, 1) self.assertFalse(self.runtime.INITIALIZED.exists()) @unittest.skipUnless(os.name == 'posix', 'native directory fsync and POSIX mode semantics') def test_native_private_file_modes_and_directory_fsync(self): self.module._fsync_dir = self.native_fsync_dir path = self.runtime.DATA / 'config/native-private-test' self.module._write(self.runtime, path, b'synthetic') self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o600) self.assertEqual(path.stat().st_nlink, 1) if __name__ == '__main__': unittest.main()