1815 lines
96 KiB
Python
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()
|