Files
truf-server/tests/test_container_import.py
2026-09-30 20:30:56 +03:00

1815 lines
96 KiB
Python

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