Files
truf-server/docker/test_windows_snapshot.py
T
2026-09-30 20:30:56 +03:00

1098 lines
55 KiB
Python

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