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

1421 lines
66 KiB
Python

"""One-shot, offline Windows logical import into a freshly provisioned container.
The caller passes its executing container_runtime module, after container and
default-environment preflight and installation of a flag-only SIGTERM handler.
SIGINT is temporarily made flag-only as well, including the child launch window.
Dispatch action import-snapshot, with required --manifest-sha256 HEX, to
import_snapshot(sys.modules[__name__], args.manifest_sha256) BEFORE initialize().
Only /import/{manifest.json,files.tar,database.dump} is read. Nothing is resumed,
repaired, removed, or started under a supervisor. A failed import needs review.
CLI contracts: postgres-runtime initialize-empty/maintenance-start/maintenance-stop
--config PATH; migrate-runtime-safety --config PATH --apply --sources-stopped
(never --initialize-base). If needed, the separate, fenced recovery action adds
--recover-stale-result-pipeline --max-rows 10000 --max-seconds 3600.
READY CLI handoffs are intentional: each CLI takes its own ClusterAuthorityLock.
initialize.lock remains held across those handoffs. Lifecycle children are NEVER
killed on a deadline/cancellation; uncertain stop retains ownership and retries.
Phases: 1 preflight, 2 archive verification, 3 extraction, 4 config, 5 initdb,
6 maintenance start, 7 raw equality, 8 Postman, 9 recovery, 10 migration,
11 final checks, 12 confirmed stop, 13 publication. Codes: 1 rejected/failed,
2 review required (count only), 124 timeout, 130 cancellation. Helpers' diagnostic
text is suppressed, including migration SQL diagnostics via temporary defaults
on the new Linux role (reset after successful final verification; left suppressed
on a failed, stopped, unmarked import). Evidence is exclusive,
private, fsynced, and never reused:
config/windows-import-manifest.json, windows-import-raw.json, and
windows-import-report.json. initialized.json is the LAST publication, contains
the Linux system identifier, runtime FORMAT, and manifest_sha256.
Safe diagnostics precede cleanup: import-snapshot-diagnostic followed by numeric
phase, stage, code, review_count, type_id, importer line (0 if unavailable), attempt.
Stages: 1 operation failure, 2 backend construction, 3 stop call, 4 stop result,
5 final probe call, 6 final probe state, 7 backend close, 8 stop CLI, 9 child wait.
Stages 5/6 mean the stop result already confirmed completed AND stopped.
Type IDs: 0 other, 1 Failure, 2 KeyboardInterrupt, 3 OSError, 4 ValueError,
5 TypeError, 6 RuntimeError (including subclasses, never their names or text).
After mutations, the first operation failure and first HOLD are also attempted
once as private, exclusive, fsynced config/windows-import-failure.json and
windows-import-hold.json. They are not completion proofs and are never replaced.
Later HOLD stages remain numeric events; failed diagnostic I/O cannot release
ownership or replace the original failure. Ordinary progress lines are unchanged.
database.database_bytes supplies the actual physical source size. Older v1
snapshots without it use max(24 GiB, dump * 4), rejecting estimates over 1 TiB,
plus all file bytes and a 20 GiB free reserve.
This does not measure physical free space on the Windows host's S: drive.
Integration still needs the image's native PG16/psycopg/PyYAML and real mount
tests. Use /data/config/windows-import.yaml for subsequent container commands;
this module neither changes the default profile nor starts the supervisor.
database.sequence_states is the exporter's optional {schema: {sequence:
{last_value: int, is_called: bool}}} evidence. Partitioned/foreign/non-public
tables are refused because v1 cannot prove their complete row inventory.
Ordinary initialize-empty failure relies on that CLI's confirmed temporary-child
cleanup; abnormal child exit instead enters maintenance-stop's indefinite HOLD.
Do not force-kill a retained importer to bypass unconfirmed stop.
"""
import contextlib
import hashlib
import json
import logging
import ntpath
import os
from pathlib import Path
import re
import shutil
import signal
import stat
import subprocess
import sys
import tarfile
import time
import unicodedata
from urllib.parse import unquote, urlsplit
IMPORT = Path('/import')
BLOCK = 1024 * 1024
GIB = 1024 ** 3
MAX_MANIFEST = 32 * BLOCK
MAX_CONFIG = 4 * BLOCK
MAX_ESTIMATE = 1024 * GIB
COUNT_TIMEOUT = 3600
RESTORE_TIMEOUT = 12 * 3600
MIGRATE_TIMEOUT = 6 * 3600
LIFECYCLE_TIMEOUT = 600
SNAPSHOT_FORMAT = 'truf-windows-snapshot-v1'
ACTIVE = {'results', 'queues', 'state', 'keychecks', 'postman_cache', 'result_spool'}
REQUIRED_FILES = {
'windows-archive/app/config.yaml', 'config/secrets.yaml',
'config/trufflehog-custom-detectors.yaml', 'runtime-linux/proxy.txt',
}
PLACEHOLDERS = {'runtime-linux/proxy.txt': b'', 'config/secrets.yaml': b'{}\n'}
FENCES = (
'lease_owner', 'lease_token', 'claim_batch', 'leased_at', 'lease_expires_at',
'current_result_reservation_id', 'claim_event_id', 'resolver_token',
)
QUIET_PG_SETTINGS = (
('log_min_error_statement', 'panic'), ('log_min_messages', 'panic'),
('log_statement', 'none'), ('log_min_duration_statement', '-1'),
)
# Fixed evidence only, not a general backup/repair facility. Recovery may append
# history and advance its own leases, but cannot rewrite these existing facts.
PRESERVED = {
'pipeline_quarantine': ('id', None),
'projection_streams': ('stream_name', None),
'projection_cursors': ('stream_name', None),
'projection_appends': ('id', None),
'projection_append_audit': ('id', None),
'projection_rotations': ('id', None),
'target_scans': ('id', (
'id', 'scan_event_id', 'scan_event_hash', 'queue_id', 'claim_lease_token',
'target', 'normalized_target', 'result_reservation_id',
)),
'result_reservations': ('id', (
'id', 'reservation_token', 'bundle_id', 'scan_event_id', 'queue_id',
'ready_relative_path', 'producer_instance_id', 'producer_pid',
'producer_creation_time', 'producer_executable',
)),
'result_bundles': ('reservation_id', (
'reservation_id', 'bundle_id', 'scan_event_id', 'scan_event_hash',
'relative_path', 'actual_bytes',
)),
'worker_progress_events': ('id', None),
'worker_diagnostics': ('id', None),
'runtime_operations': ('operation_id', None),
'runtime_operations_control': ('id', None),
'runtime_audit_events': ('id', None),
}
class Failure(Exception):
def __init__(self, code=1, count=0, *, uncertain=False):
self.code, self.count, self.uncertain = code, count, uncertain
super().__init__(code)
def _diagnostic(progress, stage, exc, attempt=0):
# Only local categories and importer line numbers cross the quiet boundary.
try:
code, count = (exc.code, exc.count) if isinstance(exc, Failure) else (
130 if isinstance(exc, KeyboardInterrupt) else 1, 0)
if type(code) is not int or code not in (1, 2, 124, 130):
code = 1
if type(count) is not int or not 0 <= count <= 2 ** 63 - 1:
count = 0
type_id = next((index for index, kind in enumerate(
(Failure, KeyboardInterrupt, OSError, ValueError, TypeError, RuntimeError), 1)
if isinstance(exc, kind)), 0)
line, traceback = 0, exc.__traceback__
while traceback is not None:
if traceback.tb_frame.f_code.co_filename == __file__:
line = traceback.tb_lineno
traceback = traceback.tb_next
progress(12, attempt, 0, cleanup=True, diagnostic=(stage, code, count, type_id, line))
except BaseException:
pass
def _checkpoint(runtime):
if runtime._shutdown_requested:
raise Failure(130)
@contextlib.contextmanager
def _quiet():
previous = logging.root.manager.disable
with open(os.devnull, 'w', encoding='ascii') as sink:
try:
logging.disable(sys.maxsize)
with contextlib.redirect_stdout(sink), contextlib.redirect_stderr(sink):
yield
finally:
logging.disable(previous)
def _integer(value, minimum=0, maximum=2 ** 63 - 1):
if type(value) is not int or not minimum <= value <= maximum:
raise Failure()
return value
def _sha(value):
if not isinstance(value, str) or not re.fullmatch(r'[0-9a-f]{64}', value):
raise Failure()
return value
def _object(pairs):
result = {}
for key, value in pairs:
if key in result:
raise Failure()
result[key] = value
return result
def _json(payload):
def invalid(_value):
raise Failure()
return json.loads(payload, object_pairs_hook=_object, parse_constant=invalid)
def _relative(name):
if (not isinstance(name, str) or not name or len(name.encode('utf-8')) > 4096
or unicodedata.normalize('NFC', name) != name
or '\\' in name or ':' in name or any(ord(c) < 32 or ord(c) == 127 for c in name)):
raise Failure()
parts = name.split('/')
if (len(parts) > 64 or any(not p or p in ('.', '..') or p.endswith((' ', '.'))
or len(p.encode('utf-8')) > 255 for p in parts)):
raise Failure()
return parts
def _destination(name):
parts = _relative(name)
archived = parts[0] == 'windows-archive' and len(parts) > 1
allowed = (archived or name in REQUIRED_FILES or
len(parts) > 2 and parts[0] == 'runtime-linux' and parts[1] in ACTIVE or
len(parts) > 2 and parts[0] == 'scanner-result-bundles'
and parts[1] in {'tmp', 'ready', 'quarantine'})
if not allowed:
raise Failure()
folded = [p.casefold() for p in parts]
if any('.lock' in p or p.endswith('.pid') for p in folded):
raise Failure()
if not archived and any(
p in {'control', 'postgres', 'postgres-linux', 'postgres-data', 'pg_wal',
'pg_xlog', 'pg_version', 'postmaster.opts', 'cluster_identity.json',
'supervisor.instance.json', 'janitor.cursor.json'}
or p.startswith('scan_limiter') for p in folded
):
raise Failure()
return parts
def _manifest(payload, expected):
if (not isinstance(expected, str) or not re.fullmatch(r'[0-9A-Fa-f]{64}', expected)
or not 0 < len(payload) <= MAX_MANIFEST
or hashlib.sha256(payload).hexdigest() != expected.lower()):
raise Failure()
value = _json(payload)
if not isinstance(value, dict) or value.get('format') != SNAPSHOT_FORMAT:
raise Failure()
source, database, archive = value['source'], value['database'], value['archive']
if (source.get('supervisor_stopped') is not True or source.get('postgres_stopped') is not True
or ntpath.normcase(source.get('root', '')) != ntpath.normcase(r'D:\truf')
or ntpath.normcase(source.get('postgres_data_dir', '')) != ntpath.normcase(r'S:\postgres-data')
or ntpath.normcase(database.get('data_directory', '')) != ntpath.normcase(r'S:\postgres-data')
or _integer(database['version_num']) // 10000 != 16
or not re.fullmatch(r'[1-9][0-9]{0,19}', database['system_identifier'])):
raise Failure()
_integer(database['port'], 1, 65535)
for name in (database['database_name'], database['user_name']):
_identifier(name)
for item in (archive, database):
_sha(item['sha256'])
_integer(item['bytes'], 6, 8 * MAX_ESTIMATE)
tables = database['table_counts']
if not isinstance(tables, dict) or not tables:
raise Failure()
for name, count in tables.items():
_identifier(name)
_integer(count)
sequences = database.get('sequence_states')
sequence_count = 0
if 'sequence_states' in database:
if not isinstance(sequences, dict):
raise Failure()
for schema, states in sequences.items():
_identifier(schema)
if schema == 'information_schema' or schema.startswith('pg_') or not isinstance(states, dict):
raise Failure()
for name, state in states.items():
_identifier(name)
if (not isinstance(state, dict) or set(state) != {'last_value', 'is_called'}
or type(state['is_called']) is not bool):
raise Failure()
_integer(state['last_value'], -(2 ** 63))
sequence_count += 1
if 'sequence_count' in database and _integer(database['sequence_count']) != sequence_count:
raise Failure()
if 'database_bytes' in database:
_integer(database['database_bytes'], 1, MAX_ESTIMATE)
files, names, parents = {}, {}, {}
if not isinstance(value['files'], list):
raise Failure()
for item in value['files']:
if not isinstance(item, dict) or set(item) != {'path', 'size', 'sha256'}:
raise Failure()
name = item['path']
parts = _destination(name)
folded = name.casefold()
if folded in names or folded in parents:
raise Failure()
for index in range(1, len(parts)):
parent = '/'.join(parts[:index])
key = parent.casefold()
if key in names or key in parents and parents[key] != parent:
raise Failure()
parents[key] = parent
names[folded] = name
files[name] = {'size': _integer(item['size'], maximum=8 * MAX_ESTIMATE),
'sha256': _sha(item['sha256'])}
if not REQUIRED_FILES.issubset(files):
raise Failure()
file_bytes = sum(item['size'] for item in files.values())
if file_bytes > archive['bytes'] or files['windows-archive/app/config.yaml']['size'] > MAX_CONFIG:
raise Failure()
if 'counts' in value:
counts = {'files': len(files), 'file_bytes': file_bytes,
'active_files': sum(not n.startswith('windows-archive/') for n in files),
'archival_files': sum(n.startswith('windows-archive/') for n in files),
'public_tables': len(tables)}
if sequences is not None:
counts['sequences'] = sequence_count
if any(_integer(value['counts'][k]) != v for k, v in counts.items()):
raise Failure()
return value, files
def _fingerprint(info):
return (info.st_dev, info.st_ino, info.st_mode, info.st_nlink, info.st_size,
info.st_mtime_ns, info.st_ctime_ns)
def _chain(path):
for parent in (*reversed(path.parents), path):
info = parent.lstat()
if stat.S_ISLNK(info.st_mode) or getattr(info, 'st_file_attributes', 0) & 0x400:
raise Failure()
def _regular(path, runtime=None):
_chain(path)
if runtime is not None:
runtime.private_path(path)
info = path.lstat()
if not stat.S_ISREG(info.st_mode) or info.st_nlink != 1:
raise Failure()
return _fingerprint(info)
@contextlib.contextmanager
def _input(path, runtime=None):
before = _regular(path, runtime)
descriptor = os.open(path, os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK | getattr(os, 'O_BINARY', 0))
with os.fdopen(descriptor, 'rb', buffering=0) as handle:
_unchanged(path, handle, before, runtime)
yield handle, before
_unchanged(path, handle, before, runtime)
def _unchanged(path, handle, before, runtime=None):
if _fingerprint(os.fstat(handle.fileno())) != before or _regular(path, runtime) != before:
raise Failure()
def _mounts(runtime):
_chain(IMPORT)
if not stat.S_ISDIR(IMPORT.lstat().st_mode):
raise Failure()
with open('/proc/self/mountinfo', 'rb') as handle:
payload = handle.read(BLOCK + 1)
if len(payload) > BLOCK:
raise Failure()
mounts = {}
for line in payload.splitlines():
fields = line.split()
if len(fields) < 10 or b'-' not in fields:
raise Failure()
# Mountinfo escapes cannot disguise an exact fixed mount or a submount.
target = re.sub(rb'\\([0-7]{3})', lambda m: bytes([int(m[1], 8)]), fields[4])
if target in mounts:
raise Failure()
mounts[target] = fields
if target.startswith((b'/import/', b'/data/')):
raise Failure()
incoming, destination = mounts.get(b'/import'), mounts.get(b'/data')
if (incoming is None or destination is None or b'ro' not in incoming[5].split(b',')
or incoming[0] == destination[0] or incoming[2:4] == destination[2:4]
or IMPORT.stat().st_dev == runtime.DATA.stat().st_dev
or not os.statvfs(IMPORT).f_flag & os.ST_RDONLY):
raise Failure()
if set(os.listdir(IMPORT)) != {'manifest.json', 'files.tar', 'database.dump'}:
raise Failure()
def _placeholder(runtime, name, before=None):
path = runtime.DATA / name
with _input(path, runtime) as (handle, fingerprint):
expected = PLACEHOLDERS[name]
if handle.read(len(expected) + 1) != expected or before is not None and fingerprint != before:
raise Failure()
return fingerprint
def _fresh(runtime, files=()):
allowed_dirs = set(runtime.DIRECTORIES)
allowed_files = {'.provisioned.json', '.provision.lock', 'initialize.lock',
'postgres-password', *PLACEHOLDERS}
seen_dirs, seen_files = set(), set()
def visit(path):
runtime.private_path(path, directory=True)
if path.stat().st_dev != runtime.DATA.stat().st_dev:
raise Failure()
with os.scandir(path) as entries:
for entry in entries:
child = path / entry.name
name = child.relative_to(runtime.DATA).as_posix()
if entry.is_dir(follow_symlinks=False):
if name not in allowed_dirs:
raise Failure()
seen_dirs.add(name)
visit(child)
else:
if name not in allowed_files:
raise Failure()
_regular(child, runtime)
seen_files.add(name)
visit(runtime.DATA)
if seen_dirs != allowed_dirs or seen_files != allowed_files:
raise Failure()
existing = {name.casefold(): name for name in seen_dirs | seen_files}
for name in files:
if name.casefold() in existing and name not in PLACEHOLDERS:
raise Failure()
parts = name.split('/')
for index in range(1, len(parts)):
parent = '/'.join(parts[:index])
found = existing.get(parent.casefold())
if found is not None and (found != parent or found not in seen_dirs):
raise Failure()
marker = runtime._read_json(runtime.PROVISIONED)
if marker != {'format': runtime.FORMAT, 'uid': 10001, 'gid': 10001}:
raise Failure()
for name in ('.provision.lock', 'initialize.lock'):
if (runtime.DATA / name).stat().st_size != 0:
raise Failure()
return {name: _placeholder(runtime, name) for name in PLACEHOLDERS}
def _space(runtime, manifest):
database = manifest['database']
estimate = database.get('database_bytes', max(24 * GIB, database['bytes'] * 4))
_integer(estimate, 1, MAX_ESTIMATE)
needed = sum(item['size'] for item in manifest['files']) + estimate + 20 * GIB
free = shutil.disk_usage(runtime.DATA).free
if free < needed:
raise Failure()
return {'database_estimate_bytes': estimate, 'required_free_bytes': needed,
'available_free_bytes': free, 'database_estimate_from_metadata': 'database_bytes' in database}
def _fsync_dir(path):
descriptor = os.open(path, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
os.fsync(descriptor)
finally:
os.close(descriptor)
@contextlib.contextmanager
def _new_file(runtime, path):
runtime.private_path(path.parent, directory=True)
descriptor = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW
| getattr(os, 'O_BINARY', 0), 0o600)
with os.fdopen(descriptor, 'wb') as handle:
_regular(path, runtime)
if not os.path.samestat(os.fstat(handle.fileno()), path.lstat()):
raise Failure()
yield handle
handle.flush()
os.fsync(handle.fileno())
_regular(path, runtime)
_fsync_dir(path.parent)
def _write(runtime, path, payload):
with _new_file(runtime, path) as handle:
if handle.write(payload) != len(payload):
raise Failure()
return hashlib.sha256(payload).hexdigest()
def _encoded(value):
return (json.dumps(value, ensure_ascii=True, sort_keys=True, separators=(',', ':')) + '\n').encode('ascii')
class _TarInfo(tarfile.TarInfo):
def _proc_pax(self, archive):
if self.type != tarfile.XHDTYPE or not 0 < self.size <= 65536:
raise Failure()
return super()._proc_pax(archive)
def _proc_gnulong(self, archive):
raise Failure()
def _proc_sparse(self, archive):
raise Failure()
def _proc_gnusparse_00(self, *args):
raise Failure()
def _proc_gnusparse_01(self, *args):
raise Failure()
def _proc_gnusparse_10(self, *args):
raise Failure()
class _HashReader:
def __init__(self, handle, size):
self.handle, self.limit = handle, size
self.digest, self.size = hashlib.sha256(), 0
def read(self, size):
block = self.handle.read(min(size, BLOCK))
self.size += len(block)
if self.size > self.limit:
raise Failure()
self.digest.update(block)
return block
def _archive(runtime, handle, metadata, files, progress, placeholders=None):
handle.seek(0)
reader = _HashReader(handle, metadata['bytes'])
seen, size, copied = set(), 0, {}
with tarfile.open(fileobj=reader, mode='r|', tarinfo=_TarInfo) as archive:
for member in archive:
_checkpoint(runtime)
name = member.name
_destination(name)
if (not member.isreg() or member.type not in (tarfile.REGTYPE, tarfile.AREGTYPE)
or member.linkname or member.sparse is not None
or set(member.pax_headers) - {'path', 'size'}
or name not in files or name in seen or member.size != files[name]['size']):
raise Failure()
seen.add(name)
destination = runtime.DATA / name
replacement = placeholders is not None and name in placeholders
output = destination.with_name(destination.name + '.windows-import-partial') if replacement else destination
if placeholders is not None:
for parent in reversed(destination.parents):
if parent == runtime.DATA or runtime.DATA in parent.parents:
if not os.path.lexists(parent):
runtime.private_path(parent.parent, directory=True)
parent.mkdir(mode=0o700)
_fsync_dir(parent.parent)
runtime.private_path(parent, directory=True)
context = _new_file(runtime, output) if placeholders is not None else contextlib.nullcontext(None)
with archive.extractfile(member) as source, context as target:
digest, received = hashlib.sha256(), 0
while True:
_checkpoint(runtime)
block = source.read(BLOCK)
if not block:
break
received += len(block)
digest.update(block)
if target is not None and target.write(block) != len(block):
raise Failure()
if received != member.size or digest.hexdigest() != files[name]['sha256']:
raise Failure()
if replacement:
_placeholder(runtime, name, placeholders[name])
os.replace(output, destination)
_regular(destination, runtime)
_fsync_dir(destination.parent)
if placeholders is not None:
copied[name] = _regular(destination, runtime)
size += received
if len(seen) % 128 == 0:
progress(3 if placeholders is not None else 2, len(seen), size)
trailer = archive.offset
while reader.read(BLOCK):
_checkpoint(runtime)
if (seen != set(files) or reader.size != metadata['bytes']
or reader.digest.hexdigest() != metadata['sha256']
or not 1024 <= reader.size - trailer <= 2 * tarfile.RECORDSIZE):
raise Failure()
# tarfile stops at its first zero header. Do not accept a second hidden tar
# or nonzero data in its buffered trailer, even if the whole hash was signed.
handle.seek(trailer)
if any(handle.read(2 * tarfile.RECORDSIZE + 1)):
raise Failure()
progress(3 if placeholders is not None else 2, len(seen), size)
return copied
def _hash_input(runtime, path, handle, before, metadata, private=False, dump=False):
_unchanged(path, handle, before, runtime if private else None)
handle.seek(0)
digest, size, prefix = hashlib.sha256(), 0, b''
while True:
_checkpoint(runtime)
block = handle.read(BLOCK)
if not block:
break
if not prefix:
prefix = block[:5]
size += len(block)
if size > metadata.get('bytes', metadata.get('size')):
raise Failure()
digest.update(block)
if (size != metadata.get('bytes', metadata.get('size')) or digest.hexdigest() != metadata['sha256']
or dump and prefix != b'PGDMP'):
raise Failure()
_unchanged(path, handle, before, runtime if private else None)
handle.seek(0)
def _configuration(runtime):
import yaml
from container_import_config import translate_windows_config
values = []
for path in (runtime.DATA / 'windows-archive/app/config.yaml', runtime.DEFAULT_CONFIG):
with _input(path, runtime) as (handle, _before):
payload = handle.read(MAX_CONFIG + 1)
if len(payload) > MAX_CONFIG:
raise Failure()
values.append(yaml.safe_load(payload))
translated, adjusted = translate_windows_config(*values)
global_config = translated['global']
if global_config.get('database_url') or global_config.get('dashboard_db_url'):
raise Failure()
path = runtime.DATA / 'config/windows-import.yaml'
config_hash = _write(runtime, path, yaml.safe_dump(translated, allow_unicode=False).encode('utf-8'))
config = runtime.prepare_environment(path)
expected = {key: str(runtime.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(runtime.DATA / 'runtime-linux/proxy.txt'),
trufflehog_config=str(runtime.DATA / 'config/trufflehog-custom-detectors.yaml'))
if any(config['global'].get(key) != value for key, value in expected.items()):
raise Failure()
return path, config, adjusted, config_hash
def _credentials(runtime):
with _input(runtime.PASSWORD, runtime) as (handle, _before):
password = handle.read(130).decode('ascii').rstrip('\n')
if not re.fullmatch(r'[A-Za-z0-9_-]{32,128}', password):
raise Failure()
dsn = os.environ['SCANNER_DB_URL']
parsed = urlsplit(dsn)
if (parsed.scheme != 'postgresql' or parsed.hostname != '127.0.0.1' or parsed.port != 5432
or unquote(parsed.username or '') != 'truf' or unquote(parsed.path) != '/truf'
or unquote(parsed.password or '') != password or parsed.query or parsed.fragment
or any(os.environ.get(k) != dsn for k in ('DATABASE_URL', 'TRUF_MANAGED_POSTGRES_DSN'))):
raise Failure()
return dsn, password
def _command(runtime, command, timeout, progress, *, env=None, stdin=None, lifecycle=False, stopping=False):
if not stopping:
_checkpoint(runtime)
process = subprocess.Popen(command, stdin=subprocess.DEVNULL if stdin is None else stdin,
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
env=env, close_fds=True, start_new_session=True)
exited, failure, held = False, None, 0
try:
deadline = time.monotonic() + timeout
while True:
try:
result = process.wait(timeout=1)
exited = True
break
except subprocess.TimeoutExpired:
pass
except BaseException:
failure = failure or Failure(130)
if not stopping and runtime._shutdown_requested:
failure = failure or Failure(130)
if time.monotonic() >= deadline:
failure = failure or Failure(124)
if failure is not None:
if lifecycle:
# A CLI can be retaining a not-yet-bookkept postmaster. Its
# own positive stop contract, not the parent, owns exit.
held += 1
if held == 1 or held % 60 == 0:
if held == 1 and not stopping:
_diagnostic(progress, 1, failure)
_diagnostic(progress, 9, failure, held)
else:
try:
process.kill()
except OSError:
pass
except BaseException as exc:
failure = failure or Failure(130 if isinstance(exc, KeyboardInterrupt) else 1)
_diagnostic(progress, 9 if stopping else 1, exc)
finally:
# Even a broken progress pipe or an interrupt outside wait() cannot
# abandon a lifecycle child, nor a client with an open restore session.
while not exited:
if not lifecycle:
try:
process.kill()
except BaseException:
pass
try:
result = process.wait(timeout=1)
exited = True
except BaseException:
pass
if lifecycle and result < 0:
raise Failure(failure.code if failure is not None else 1, uncertain=True)
if failure is not None and not stopping:
raise failure
if result != 0:
raise Failure()
if not stopping:
_checkpoint(runtime)
def _identifier(name):
if not isinstance(name, str) or not name or '\x00' in name or len(name.encode('utf-8')) > 63:
raise Failure()
return '"' + name.replace('"', '""') + '"'
@contextlib.contextmanager
def _connect(runtime, *, readonly=True):
import psycopg
from psycopg.rows import dict_row
dsn, _password = _credentials(runtime)
options = (
f'-c search_path=public -c statement_timeout={COUNT_TIMEOUT * 1000} '
'-c lock_timeout=10000 -c idle_in_transaction_session_timeout=3600000 '
'-c row_security=off -c log_min_error_statement=panic -c log_statement=none '
'-c log_min_messages=panic -c log_min_duration_statement=-1 '
'-c default_transaction_read_only=' + ('on' if readonly else 'off')
)
connection = psycopg.connect(dsn, connect_timeout=10, row_factory=dict_row,
options=options, application_name='truf-container-import',
tcp_user_timeout=60000)
try:
yield connection
finally:
connection.close()
def _identity(runtime, identity, source):
if (identity['pg_major'] != 16 or identity['data_directory'] != str(runtime.DATA / 'postgres-linux')
or identity['database'] != 'truf' or identity['user'] != 'truf' or identity['port'] != 5432
or not re.fullmatch(r'[1-9][0-9]{0,19}', identity['system_identifier'])
or identity['system_identifier'] == source['system_identifier']):
raise Failure()
def _online(connection, identity):
row = connection.execute("""SELECT current_database() AS database, current_user AS user_name,
current_setting('data_directory') AS data_directory, current_setting('port')::int AS port,
current_setting('server_version_num')::int AS version_num, pg_is_in_recovery() AS in_recovery,
(SELECT system_identifier::text FROM pg_catalog.pg_control_system()) AS system_identifier,
(SELECT rolsuper FROM pg_catalog.pg_roles WHERE rolname = current_user) AS superuser,
(SELECT count(*) FROM pg_catalog.pg_stat_activity
WHERE backend_type = 'client backend' AND pid <> pg_backend_pid()) AS other_clients""").fetchone()
expected = {'database': identity['database'], 'user_name': identity['user'],
'data_directory': identity['data_directory'], 'port': identity['port'],
'system_identifier': identity['system_identifier'], 'in_recovery': False,
'superuser': True, 'other_clients': 0}
if (not row or any(row.get(k) != v for k, v in expected.items())
or _integer(row['version_num']) // 10000 != 16):
raise Failure()
return row['version_num']
def _database_equality(runtime, connection, database, identity, progress, *, empty=False):
version = _online(connection, identity)
if empty:
row = connection.execute("""SELECT
(SELECT count(*) FROM pg_catalog.pg_class c JOIN pg_catalog.pg_namespace n
ON n.oid = c.relnamespace WHERE n.nspname <> 'information_schema' AND n.nspname !~ '^pg_')
+ (SELECT count(*) FROM pg_catalog.pg_proc p JOIN pg_catalog.pg_namespace n
ON n.oid = p.pronamespace WHERE n.nspname <> 'information_schema' AND n.nspname !~ '^pg_')
+ (SELECT count(*) FROM pg_catalog.pg_type t JOIN pg_catalog.pg_namespace n
ON n.oid = t.typnamespace WHERE n.nspname <> 'information_schema' AND n.nspname !~ '^pg_')
+ (SELECT count(*) FROM pg_catalog.pg_namespace
WHERE nspname NOT IN ('public', 'information_schema') AND nspname !~ '^pg_') AS count""").fetchone()
if row['count'] != 0:
raise Failure()
return {}
tables = connection.execute("""SELECT n.nspname AS schema_name, c.relname AS name, c.relkind AS kind
FROM pg_catalog.pg_class c JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE c.relkind IN ('r','p','f') AND n.nspname <> 'information_schema' AND n.nspname !~ '^pg_'
ORDER BY n.nspname, c.relname""").fetchall()
if (any(t['schema_name'] != 'public' or t['kind'] != 'r' for t in tables)
or {t['name'] for t in tables} != set(database['table_counts'])):
raise Failure()
counts = {}
for table in tables:
_checkpoint(runtime)
name = table['name']
count = connection.execute('SELECT count(*) AS count FROM ONLY "public".' + _identifier(name)).fetchone()['count']
if _integer(count) != database['table_counts'][name]:
raise Failure()
counts[name] = count
progress(7, len(counts), 0)
sequence_count = 0
if 'sequence_states' in database:
sequences = connection.execute("""SELECT n.nspname AS schema_name, c.relname AS name
FROM pg_catalog.pg_class c JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE c.relkind = 'S' AND n.nspname <> 'information_schema' AND n.nspname !~ '^pg_'
ORDER BY n.nspname, c.relname""").fetchall()
expected = {(schema, name): state for schema, states in database['sequence_states'].items()
for name, state in states.items()}
if {(row['schema_name'], row['name']) for row in sequences} != set(expected):
raise Failure()
for (schema, name), state in expected.items():
_checkpoint(runtime)
row = connection.execute('SELECT last_value, is_called FROM ' + _identifier(schema)
+ '.' + _identifier(name)).fetchone()
if dict(row) != state:
raise Failure()
sequence_count += 1
connection.commit()
_online(connection, identity)
connection.commit()
return {'table_counts': counts, 'tables': len(counts), 'rows': sum(counts.values()),
'sequences_verified': sequence_count, 'sequences_provided': 'sequence_states' in database,
'version_num': version}
def _postman_target(runtime, target, normalized, files):
from target_identity import postman_target_identity
if not isinstance(target, str) or len(target.encode('utf-8')) > BLOCK:
raise Failure(2, 1)
value = _json(target)
if not isinstance(value, dict):
raise Failure(2, 1)
digest = _sha(str(value.get('sha256', '')).strip().lower())
if normalized != 'postman:sha256:' + digest or postman_target_identity(value) != normalized:
raise Failure(2, 1)
changed, offered = False, False
for key in ('cache_path', 'local_path'):
if key not in value:
continue
offered = True
path = value[key]
if not isinstance(path, str) or not path:
raise Failure(2, 1)
windows = path.replace('\\', '/')
prefix = 'd:/truf/runtime/postman_cache/'
linux = str(runtime.DATA / 'runtime-linux/postman_cache').replace('\\', '/') + '/'
if windows.casefold().startswith(prefix):
relative = windows[len(prefix):]
elif path.startswith(linux):
relative = path[len(linux):]
else:
raise Failure(2, 1)
_relative(relative)
name = 'runtime-linux/postman_cache/' + relative
# The manifest rejects casefold collisions. Windows path case may differ
# from the preserved filename; never rename the copied artifact itself.
entry = files.get(name.casefold())
if entry is None or entry['sha256'] != digest or entry['size'] <= 0:
raise Failure(2, 1)
for size_key in ('size', 'bytes'):
if size_key in value and (type(value[size_key]) not in (int, str)
or int(value[size_key]) != entry['size']):
raise Failure(2, 1)
newpath = runtime.DATA / entry['path']
with _input(newpath, runtime) as (handle, before):
_hash_input(runtime, newpath, handle, before, entry, private=True)
newvalue = str(newpath)
changed |= newvalue != value[key]
value[key] = newvalue
if not offered or postman_target_identity(value) != normalized:
raise Failure(2, 1)
return json.dumps(value, ensure_ascii=True, sort_keys=True, separators=(',', ':')) if changed else target
def _rebase_postman(runtime, connection, files):
from target_identity import postman_target_identity
cache = {name.casefold(): dict(entry, path=name) for name, entry in files.items()
if name.startswith('runtime-linux/postman_cache/')}
adjusted, reviewed, unsafe, after = 0, 0, 0, 0
while True:
_checkpoint(runtime)
rows = connection.execute("""SELECT * FROM public.target_queue
WHERE platform IN ('postman','github_gists','github_archive_files')
AND status IN ('pending','deferred','in_progress') AND id > %s
ORDER BY id LIMIT 128 FOR UPDATE""", (after,)).fetchall()
if not rows:
break
for row in rows:
after = _integer(row['id'], 1)
reviewed += 1
try:
if (row['status'] not in ('pending', 'deferred')
or any(row[name] is not None for name in FENCES)
or row['resolver_state'] not in (None, 'resolved')):
raise Failure(2, 1)
normalized = row['normalized_target']
if row['platform'] != 'postman':
# ScannerDB normalizes these platforms as the entire lower-
# cased JSON, not the artifact digest. Changing a locator
# would change queue identity, even for the same artifact.
if (not isinstance(row['target'], str) or len(row['target'].encode('utf-8')) > BLOCK
or normalized != row['target'].strip().lower()):
raise Failure(2, 1)
normalized = postman_target_identity(row['target'])
replacement = _postman_target(runtime, row['target'], normalized, cache)
if row['platform'] != 'postman' and replacement != row['target']:
raise Failure(2, 1)
except (Failure, ValueError, TypeError, OSError):
_checkpoint(runtime)
unsafe += 1
continue
if replacement != row['target']:
# All rows are locked in this transaction. Only the target JSON
# changes: no timestamp, identity, attempt, token, or history edit.
result = connection.execute("""UPDATE public.target_queue SET target = %s
WHERE id = %s AND target = %s AND normalized_target = %s
AND status IN ('pending','deferred') AND current_result_reservation_id IS NULL
AND lease_token IS NULL AND claim_event_id IS NULL""",
(replacement, row['id'], row['target'], row['normalized_target']))
if result.rowcount != 1:
raise Failure()
adjusted += 1
if unsafe:
connection.rollback()
raise Failure(2, unsafe)
_checkpoint(runtime)
connection.commit()
return {'postman_reviewed': reviewed, 'postman_adjusted': adjusted}
def _preserved(runtime, connection, previous=None):
evidence = {}
for table, (key, selected) in PRESERVED.items():
_checkpoint(runtime)
if previous is not None and table not in previous:
continue
old = previous.get(table) if previous is not None else None
if old is None:
columns = [row['name'] for row in connection.execute("""SELECT a.attname AS name
FROM pg_catalog.pg_attribute a JOIN pg_catalog.pg_class c ON c.oid = a.attrelid
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = 'public' AND c.relname = %s AND a.attnum > 0 AND NOT a.attisdropped
ORDER BY a.attname""", (table,)).fetchall()]
if not columns:
continue
if key not in columns:
raise Failure(2, 1)
# Additive migration may introduce a table/column not in this
# snapshot. The preservation proof covers every preexisting fact.
columns = [name for name in selected if name in columns] if selected else columns
cutoff = connection.execute('SELECT max(' + _identifier(key) + ') AS cutoff FROM public.'
+ _identifier(table)).fetchone()['cutoff']
else:
columns, cutoff = old['columns'], old['cutoff']
digest, count = hashlib.sha256(), 0
if cutoff is not None:
selection = ','.join(_identifier(column) for column in columns)
# Named cursor keeps millions of historical identifiers off both
# the client's buffered result set and the diagnostic channels.
with connection.cursor(name='windows_import_evidence') as cursor:
cursor.execute('SELECT pg_catalog.encode(pg_catalog.sha256(pg_catalog.convert_to('
'pg_catalog.row_to_json(e)::text, \'UTF8\')), \'hex\') AS digest FROM '
'(SELECT ' + selection + ' FROM public.' + _identifier(table)
+ ' WHERE ' + _identifier(key) + ' <= %s ORDER BY ' + _identifier(key) + ') e', (cutoff,))
for row in cursor:
_checkpoint(runtime)
digest.update(_sha(row['digest']).encode('ascii'))
count += 1
item = {'columns': columns, 'cutoff': cutoff, 'count': count, 'sha256': digest.hexdigest()}
if old is not None and item != old:
raise Failure(2, 1)
evidence[table] = item
connection.commit()
return evidence
def _pipeline_counts(connection):
queries = {
'worker_leases': "SELECT count(*) AS count FROM public.pipeline_leases WHERE state NOT IN ('released','failed')",
'result_reservations': "SELECT count(*) AS count FROM public.result_reservations WHERE state IN ('scanning','ready','ingesting','db_committed')",
'queue_leases': "SELECT count(*) AS count FROM public.target_queue q LEFT JOIN public.result_reservations r ON r.id=q.current_result_reservation_id WHERE q.status='in_progress' OR r.state IN ('scanning','ready','ingesting','db_committed')",
'blob_leases': "SELECT count(*) AS count FROM public.docker_content_blobs WHERE state IN ('leased','submitted') OR lease_reservation_id IS NOT NULL",
}
counts = {name: _integer(connection.execute(sql).fetchone()['count']) for name, sql in queries.items()}
connection.commit()
return counts
def _projection_files(runtime, connection, files):
rows = connection.execute("""SELECT s.stream_name, c.stream_name AS cursor_stream_name,
s.base_relative_path, s.current_generation, c.generation, c.committed_offset
FROM public.projection_streams s
FULL JOIN public.projection_cursors c ON c.stream_name = s.stream_name""").fetchall()
scan_streams = {'scan_results': 'scan_results.jsonl', 'found_secrets': 'found_secrets.jsonl',
'scan_errors': 'scan_errors.log'}
stream_names = {row.get('stream_name') for row in rows}
if not scan_streams.keys() <= stream_names or len(stream_names) != len(rows):
raise Failure(2, len(rows))
for row in rows:
stream_name = row.get('stream_name')
if (not isinstance(stream_name, str) or row.get('cursor_stream_name') != stream_name
or _integer(row.get('current_generation')) != _integer(row.get('generation'))):
raise Failure(2, 1)
offset = _integer(row.get('committed_offset'))
relative = row.get('base_relative_path')
_relative(relative)
root = 'runtime-linux/results/'
expected = scan_streams.get(stream_name)
status = False
if expected is None:
match = re.fullmatch(r'keycheck:([a-z0-9][a-z0-9_.-]{0,63}):(results|status)', stream_name)
if not match:
raise Failure(2, 1)
status = match[2] == 'status'
suffix = 'Checked.txt' if status else 'Results.jsonl'
expected = f'{match[1]}/{match[1]}{suffix}'
root = 'runtime-linux/keychecks/'
if relative != expected:
raise Failure(2, 1)
name = root + relative
path = runtime.DATA / name
metadata = files.get(name)
if status:
# Replacement status snapshots do not advance historical append cursors.
if metadata is None:
raise Failure(2, 1)
with _input(path, runtime) as (handle, before):
_hash_input(runtime, path, handle, before, metadata, private=True)
else:
if offset != (metadata['size'] if metadata is not None else 0):
raise Failure(2, 1)
if metadata is not None:
if _regular(path, runtime)[4] != metadata['size']:
raise Failure()
elif os.path.lexists(path):
raise Failure(2, 1)
connection.commit()
def _backend_stop(pg, config, backend, progress):
retries = 0
while True:
stage = 3
try:
result = pg.maintenance_stop(config, backend=backend)
stage = 4
if result.completed is not True or result.stopped is not True:
raise Failure()
stage = 5
probe = backend.probe()
stage = 6
if probe.kind != pg.ProbeKind.STOPPED:
raise Failure()
return
except BaseException as exc:
retries += 1
_diagnostic(progress, stage, exc, retries)
try:
time.sleep(2)
except BaseException:
pass
@contextlib.contextmanager
def _authority(runtime, config, source, progress, *, stopped=False):
import postgres_runtime as pg
from runtime_security import ClusterAuthorityLock
dsn, _password = _credentials(runtime)
lock = ClusterAuthorityLock(config, endpoint_dsn=dsn)
lock.acquire()
backend = None
try:
backend = pg.PostgresBackend(config)
if stopped:
_backend_stop(pg, config, backend, progress)
elif backend.probe().kind != pg.ProbeKind.READY:
raise Failure()
identity = pg.verify_cluster_identity(config)
_identity(runtime, identity, source)
yield identity
if not stopped and backend.probe().kind != pg.ProbeKind.READY:
raise Failure()
except BaseException as exc:
_diagnostic(progress, 1, exc)
retries = 0
while backend is None:
try:
backend = pg.PostgresBackend(config)
except BaseException as stop_exc:
retries += 1
_diagnostic(progress, 2, stop_exc, retries)
try:
time.sleep(2)
except BaseException:
pass
_backend_stop(pg, config, backend, progress)
raise
finally:
if backend is not None:
retries = 0
while True:
try:
backend.close()
break
except BaseException as exc:
retries += 1
_diagnostic(progress, 7, exc, retries)
_backend_stop(pg, config, backend, progress)
lock.release()
def _stop_cli(runtime, path, progress):
retries = 0
while True:
try:
_command(runtime, runtime._bootstrap_command('postgres-runtime', 'maintenance-stop', '--config', str(path)),
LIFECYCLE_TIMEOUT, progress, lifecycle=True, stopping=True)
return
except BaseException as exc:
retries += 1
_diagnostic(progress, 8, exc, retries)
try:
time.sleep(2)
except BaseException:
pass
def _restore(runtime, path, config, manifest, files, dump, fingerprint, progress, report):
try:
_command(runtime, runtime._bootstrap_command('postgres-runtime', 'initialize-empty', '--config', str(path)),
LIFECYCLE_TIMEOUT, progress, lifecycle=True)
except Failure as exc:
_diagnostic(progress, 1, exc)
if exc.uncertain:
# A killed init CLI did not fulfill its ownership contract. Without
# an identity the stop CLI must HOLD, not adopt or release this data.
_stop_cli(runtime, path, progress)
report['maintenance_stopped'] = True
raise
# initialize-empty itself retains its temporary postmaster until stopped,
# including failure. A failed unbound initdb is not adopted by maintenance.
started = False
try:
_checkpoint(runtime)
progress(6)
started = True
_command(runtime, runtime._bootstrap_command('postgres-runtime', 'maintenance-start', '--config', str(path)),
LIFECYCLE_TIMEOUT, progress, lifecycle=True)
with _authority(runtime, config, manifest['database'], progress) as identity:
report['linux_system_identifier'] = identity['system_identifier']
progress(7)
with _connect(runtime) as connection:
_database_equality(runtime, connection, manifest['database'], identity, progress, empty=True)
_hash_input(runtime, IMPORT / 'database.dump', dump, fingerprint, manifest['database'], dump=True)
_dsn, password = _credentials(runtime)
from runtime_security import require_trusted_native_executable
restore_binary = require_trusted_native_executable('/usr/lib/postgresql/16/bin/pg_restore')
environment = {k: v for k, v in os.environ.items() if not k.upper().startswith('PG')}
environment.update(
PGHOST='127.0.0.1', PGHOSTADDR='127.0.0.1', PGPORT='5432', PGUSER='truf',
PGDATABASE='truf', PGPASSWORD=password, PGPASSFILE=os.devnull, PGSERVICEFILE=os.devnull,
PGSSLMODE='disable', PGGSSENCMODE='disable', PGCONNECT_TIMEOUT='10',
PGAPPNAME='truf-container-import-restore', PGCLIENTENCODING='UTF8', LC_ALL='C', LANG='C',
PGOPTIONS=f'-c statement_timeout={RESTORE_TIMEOUT * 1000} -c lock_timeout=60000 '
'-c idle_in_transaction_session_timeout=3600000 -c row_security=off '
'-c log_min_error_statement=panic -c log_min_messages=panic '
'-c log_statement=none -c log_min_duration_statement=-1',
)
_command(runtime, [restore_binary, '--dbname=truf', '--no-password',
'--single-transaction', '--exit-on-error', '--no-owner', '--no-acl', '--no-tablespaces'],
RESTORE_TIMEOUT, progress, env=environment, stdin=dump)
_hash_input(runtime, IMPORT / 'database.dump', dump, fingerprint, manifest['database'], dump=True)
with _connect(runtime) as connection:
raw = _database_equality(runtime, connection, manifest['database'], identity, progress)
report['raw'] = raw
preserved = _preserved(runtime, connection)
_projection_files(runtime, connection, files)
report['pipeline_before'] = _pipeline_counts(connection)
report['raw_evidence_sha256'] = _write(runtime, runtime.DATA / 'config/windows-import-raw.json',
_encoded({'manifest_sha256': report['manifest_sha256'], **raw}))
progress(8)
with _connect(runtime, readonly=False) as connection:
report.update(_rebase_postman(runtime, connection, files))
# The existing migration CLI clears PGOPTIONS. Database-scoped
# defaults on this freshly generated Linux role also suppress
# server log DETAIL/STATEMENT values, not just child stderr.
for setting, value in QUIET_PG_SETTINGS:
connection.execute('ALTER ROLE "truf" IN DATABASE "truf" SET '
+ setting + " TO '" + value + "'")
connection.commit()
report['recovery_applied'] = any(report['pipeline_before'].values())
if report['recovery_applied']:
progress(9)
_command(runtime, runtime._bootstrap_command(
'migrate-runtime-safety', '--config', str(path), '--apply', '--sources-stopped',
'--recover-stale-result-pipeline', '--max-rows', '10000', '--max-seconds', '3600',
), COUNT_TIMEOUT + 600, progress)
# Bundle ingestion can insert derived_postman_targets carrying the
# original Windows locators. Recheck after that intentional handoff.
with _authority(runtime, config, manifest['database'], progress):
with _connect(runtime, readonly=False) as connection:
_online(connection, identity)
for name, count in _rebase_postman(runtime, connection, files).items():
# Reviewed counts are visits; adjusted counts are writes.
report[name] += count
# Both CLIs own their own authority; no importer SQL sessions survive
# the handoff (the migration explicitly rejects other client sessions).
progress(10)
_command(runtime, runtime._bootstrap_command(
'migrate-runtime-safety', '--config', str(path), '--apply', '--sources-stopped',
), MIGRATE_TIMEOUT, progress)
progress(11)
with _authority(runtime, config, manifest['database'], progress):
with _connect(runtime) as connection:
_online(connection, identity)
_preserved(runtime, connection, preserved)
report['preserved'] = {name: {k: item[k] for k in ('count', 'sha256')}
for name, item in preserved.items()}
report['pipeline_after'] = _pipeline_counts(connection)
if any(report['pipeline_after'].values()):
raise Failure(2, sum(report['pipeline_after'].values()))
from scanner_db import ScannerDB
db = ScannerDB(db_url=_credentials(runtime)[0], initialize=False)
try:
if not db.enabled or not db.conn.is_postgres:
raise Failure()
db.conn.execute('SET default_transaction_read_only = on')
db.conn.execute(f'SET statement_timeout = {COUNT_TIMEOUT * 1000}')
db.conn.commit()
db.require_runtime_safety_schema()
db.require_final_cutover()
cutover = db.final_cutover_status()
report['cutover_sha256'] = _sha(cutover['evidence_sha256'])
report['migration_rows'] = _integer(db.conn.execute(
'SELECT count(*) AS count FROM runtime_schema_migrations').fetchone()['count'])
db.conn.commit()
finally:
db.close()
with _connect(runtime, readonly=False) as connection:
for setting, _value in QUIET_PG_SETTINGS:
connection.execute('ALTER ROLE "truf" IN DATABASE "truf" RESET ' + setting)
connection.commit()
return identity
except BaseException as exc:
_diagnostic(progress, 1, exc)
raise
finally:
if started:
try:
progress(12, cleanup=True)
except BaseException:
pass
_stop_cli(runtime, path, progress)
report['maintenance_stopped'] = True
def import_snapshot(runtime, manifest_sha256) -> int:
"""Import once, returning only a safe code; never let diagnostics escape."""
output, errors = sys.stdout, sys.stderr
phase, report, mutations = 1, {}, False
def progress(number, count=0, size=0, *, cleanup=False, diagnostic=None):
nonlocal phase
if not cleanup:
_checkpoint(runtime)
phase = number
if diagnostic is not None:
stage, failure_code, review_count, type_id, line = diagnostic
event = dict(phase=phase, stage=stage, code=failure_code, review_count=review_count,
type_id=type_id, line=line, attempt=count)
key = 'failure' if stage == 1 else 'hold'
first = key not in report
if first:
report[key] = event
if first or stage != 1:
try:
print('import-snapshot-diagnostic', *event.values(), file=errors, flush=True)
except BaseException:
pass
if first and mutations:
try:
_write(runtime, runtime.DATA / ('config/windows-import-' + key + '.json'), _encoded(event))
except BaseException:
pass
try:
print('import-snapshot', number, count, size, file=output, flush=True)
except BaseException:
if not cleanup:
raise
code, review_count, signal_installed = 1, 0, False
try:
previous_interrupt = signal.signal(signal.SIGINT, lambda *_: setattr(runtime, '_shutdown_requested', True))
signal_installed = True
with _quiet():
from runtime_security import PrivateFileLock, write_private_json_exclusive
_checkpoint(runtime)
lock_path = runtime.INITIALIZE_LOCK
runtime.private_path(lock_path.parent, directory=True)
if os.path.lexists(lock_path):
_regular(lock_path, runtime)
with PrivateFileLock(str(lock_path)), contextlib.ExitStack() as stack:
try:
progress(1)
_mounts(runtime)
opened = {name: stack.enter_context(_input(IMPORT / name))
for name in ('manifest.json', 'files.tar', 'database.dump')}
manifest_handle, manifest_fingerprint = opened['manifest.json']
payload = manifest_handle.read(MAX_MANIFEST + 1)
manifest, files = _manifest(payload, manifest_sha256)
placeholders = _fresh(runtime, files)
report = {'manifest_sha256': manifest_sha256.lower(),
'archive_sha256': manifest['archive']['sha256'],
'database_sha256': manifest['database']['sha256'],
'files': len(files), 'file_bytes': sum(f['size'] for f in files.values()),
**_space(runtime, manifest)}
archive, archive_fingerprint = opened['files.tar']
dump, dump_fingerprint = opened['database.dump']
progress(2)
_archive(runtime, archive, manifest['archive'], files, progress)
_unchanged(IMPORT / 'files.tar', archive, archive_fingerprint)
_hash_input(runtime, IMPORT / 'database.dump', dump, dump_fingerprint, manifest['database'], dump=True)
_unchanged(IMPORT / 'manifest.json', manifest_handle, manifest_fingerprint)
if _fresh(runtime, files) != placeholders:
raise Failure()
_space(runtime, manifest)
progress(3)
mutations = True
_write(runtime, runtime.DATA / 'config/windows-import-manifest.json', payload)
copied = _archive(runtime, archive, manifest['archive'], files, progress, placeholders)
_unchanged(IMPORT / 'files.tar', archive, archive_fingerprint)
progress(4)
path, config, adjusted, config_hash = _configuration(runtime)
report.update(adjusted_keys=adjusted, config_sha256=config_hash)
os.environ.update(TRUF_DB_STATEMENT_TIMEOUT_MS=str(COUNT_TIMEOUT * 1000),
TRUF_DB_LOCK_TIMEOUT_MS='10000', TRUF_DB_IDLE_TRANSACTION_TIMEOUT_MS='3600000')
progress(5)
identity = _restore(runtime, path, config, manifest, files, dump, dump_fingerprint, progress, report)
# Reacquire the stopped endpoint and retain it through the
# last publication, rather than treating a CLI exit as a marker.
with _authority(runtime, config, manifest['database'], progress, stopped=True) as stopped_identity:
if identity != stopped_identity:
raise Failure()
checked = 0
for name, before in copied.items():
_checkpoint(runtime)
# Only the explicitly permitted fenced recovery can
# consume/move incoming tmp/ready bundle artifacts.
if (report.get('recovery_applied') and name.startswith(
('scanner-result-bundles/tmp/', 'scanner-result-bundles/ready/'))):
continue
destination = runtime.DATA / name
if _regular(destination, runtime) != before:
with _input(destination, runtime) as (handle, current):
_hash_input(runtime, destination, handle, current, files[name], private=True)
checked += 1
if checked % 1024 == 0:
progress(11, checked, 0)
report['preserved_files'] = checked
report['remaining_free_bytes'] = shutil.disk_usage(runtime.DATA).free
if report['remaining_free_bytes'] < 20 * GIB:
raise Failure()
for name, (handle, before) in opened.items():
_unchanged(IMPORT / name, handle, before)
stack.close()
progress(13)
report.update(status='verified-stopped', maintenance_stopped=True)
report_hash = _write(runtime, runtime.DATA / 'config/windows-import-report.json', _encoded(report))
_checkpoint(runtime)
write_private_json_exclusive(str(runtime.INITIALIZED), {
'format': runtime.FORMAT, 'system_identifier': identity['system_identifier'],
'pg_major': identity['pg_major'], 'manifest_sha256': manifest_sha256.lower(),
'import_report_sha256': report_hash,
})
code = 0
except BaseException as exc:
_diagnostic(progress, 1, exc)
failure = report.get('failure', {})
code, review_count = failure.get('code', 1), failure.get('review_count', 0)
if mutations and not os.path.lexists(runtime.DATA / 'config/windows-import-report.json'):
report.update(status='failed-unmarked', phase=phase, code=code, review_count=review_count)
_write(runtime, runtime.DATA / 'config/windows-import-report.json', _encoded(report))
raise
except BaseException as exc:
_diagnostic(progress, 1, exc)
failure = report.get('failure', {})
code, review_count = failure.get('code', 1), failure.get('review_count', 0)
finally:
if signal_installed:
signal.signal(signal.SIGINT, previous_interrupt)
if code:
print('import-snapshot', phase, code, review_count, file=errors, flush=True)
else:
summary = {name: report[name] for name in (
'manifest_sha256', 'archive_sha256', 'database_sha256', 'files', 'file_bytes',
'config_sha256', 'adjusted_keys', 'cutover_sha256', 'migration_rows',
'postman_reviewed', 'postman_adjusted',
)}
summary.update({key: report['raw'][key] for key in ('tables', 'rows', 'sequences_verified')})
print(json.dumps(summary, ensure_ascii=True, sort_keys=True), file=output, flush=True)
return code