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

881 lines
38 KiB
Python

"""Offline Windows export, never a supervisor launcher or a source migrator.
Run only after the canonical Windows supervisor stop:
python -I -S -B docker/windows_snapshot.py capture --output D:\\truf-docker\\docker\\imports\\NAME
Output is numeric: phase, count, bytes; failures are phase, exit code. Phases:
1 arguments, 2 source loading, 3 private output, 4 offline inventory, 5 maintenance
start, 6 database, 7 confirmed stop (count = retries), 8 tar, 9 revalidation,
10 publication. Exit 1 is a guarded failure; 124 is a client timeout.
An unconfirmed maintenance stop deliberately RETAINS authority and retries. Do
not terminate this process to bypass that guard. There is no valid manifest
until cleanup confirms STOPPED. Killing Windows/processes can defeat any lock.
SIGINT/SIGBREAK only request cancellation while capture owns the source. They
cannot unwind the original backend's spawned-but-not-yet-bookkept launch window.
Assumes intact original helper APIs/layout and existing credentials permitted to
dump all data, read pg_control_system(), and observe all client sessions.
"""
import argparse
import contextlib
import ctypes
from dataclasses import dataclass
from datetime import datetime, timezone
import fnmatch
import hashlib
import importlib
import importlib.util
import json
import ntpath
import os
from pathlib import Path, PureWindowsPath
import re
import shutil
import signal
import stat
import subprocess
import sys
import tarfile
import threading
import time
from types import SimpleNamespace
from urllib.parse import unquote, urlsplit
SOURCE_ROOT = Path(r'D:\truf')
POSTGRES_DATA = Path(r'S:\postgres-data')
BUNDLE_ROOT = Path(r'S:\scanner-result-bundles')
IMPORTS_ROOT = Path(r'D:\truf-docker\docker\imports')
KEYCHECK_COPY = 'keychecks \u2014 \u043a\u043e\u043f\u0438\u044f'
BLOCK = 1024 * 1024
QUERY_TIMEOUT = 30
COUNT_TIMEOUT = 1800
DUMP_TIMEOUT = 6 * 3600
SESSION_TIMEOUT = 12 * 3600
MAX_METADATA = 4 * BLOCK
ACTIVE_DIRS = ('queues', 'state', 'keychecks', 'results', 'postman_cache', 'result_spool')
EXCLUSION_POLICY = [
'Only explicitly reviewed source roots and mappings are selected.',
'No physical PGDATA/WAL, PostgreSQL binaries/logs, or active control authority.',
'No runtime/downloads, git, traces, freeze-diagnostics, or S: scanner-work.',
'No gharchive_cache, .git, .opencode, tests, code caches, *.lock*, or *.pid.',
'Ordinary *.log* excluded outside result/keycheck projections; scan_errors.log* retained.',
'Windows scratch databases, temporary state and janitor cursor are archival only.',
]
class Failure(Exception):
def __init__(self, code=1):
self.code = code
super().__init__(code)
@contextlib.contextmanager
def _defer_signals():
pending = False
previous = {}
def request(_number, _frame):
nonlocal pending
pending = True
def checkpoint():
if pending:
raise Failure()
try:
for number in (signal.SIGINT, getattr(signal, 'SIGBREAK', None)):
if number is not None:
previous[number] = signal.signal(number, request)
yield checkpoint
finally:
for number, handler in previous.items():
signal.signal(number, handler)
@dataclass(frozen=True)
class File:
source: Path
size: int
fingerprint: tuple
@dataclass
class Inventory:
files: dict
directories: dict
exclusions: dict
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, getattr(info, 'st_file_attributes', 0))
def _check_type(info, directory=False):
if (getattr(info, 'st_file_attributes', 0) & 0x400
or stat.S_ISLNK(info.st_mode)
or not (stat.S_ISDIR(info.st_mode) if directory else stat.S_ISREG(info.st_mode))
or (not directory and info.st_nlink != 1)):
raise Failure()
def _check_chain(path):
# Inspect ancestors first: even lstat(child) otherwise traverses a junction.
path = Path(path).absolute()
for part in (*reversed(path.parents), path):
info = part.lstat()
if stat.S_ISLNK(info.st_mode) or getattr(info, 'st_file_attributes', 0) & 0x400:
raise Failure()
def _file_info(path):
_check_chain(path)
info = path.lstat()
_check_type(info)
return info
def _same_windows_path(left, right):
return ntpath.normcase(ntpath.normpath(str(left))) == ntpath.normcase(ntpath.normpath(str(right)))
def _output_path(value):
path = PureWindowsPath(value)
name = path.name
if (not path.is_absolute() or not _same_windows_path(path.parent, IMPORTS_ROOT)
or any(part in ('.', '..') for part in re.split(r'[\\/]', value))
or not re.fullmatch(r'[A-Za-z0-9][A-Za-z0-9_.-]{0,95}', name)
or name.endswith('.')
or re.fullmatch(r'CON|PRN|AUX|NUL|COM[1-9]|LPT[1-9]', name.split('.')[0], re.I)):
raise Failure()
return Path(str(path))
def _verify_private_acl(path, security, directory):
security.reject_reparse_components(str(path))
sddl = security._windows_private_sddl(str(path)).upper()
alias = lambda sid: 'SY' if sid == 'S-1-5-18' else sid
sid = alias(security._windows_current_user_sid().upper())
owner = re.search(r'O:([^:()]+?)(?=[GDS]:|$)', sddl)
aces = [ace.split(';') for ace in re.findall(r'\(([^()]*)\)', sddl)]
expected = {sid, 'SY'}
if (not owner or alias(owner.group(1)) != sid or 'D:P' not in sddl
or len(aces) != len(expected)):
raise Failure()
trustees = set()
for ace in aces:
if (len(ace) != 6 or ace[:5] != ['A', 'OICI' if directory else '', 'FA', '', '']):
raise Failure()
trustees.add(alias(ace[5]))
if trustees != expected:
raise Failure()
def _secure_path(path, security, directory=False):
# The original public helper includes BA. Narrow its verified, empty path
# using the same no-reparse Win32 primitives, before any sensitive write.
harden = security.harden_private_directory if directory else security.harden_private_file
harden(str(path))
sid = security._windows_current_user_sid()
trustees = ('SY',) if sid.upper() == 'S-1-5-18' else (sid, 'SY')
flags = 'OICI' if directory else ''
sddl = 'D:P' + ''.join(f'(A;{flags};FA;;;{trustee})' for trustee in trustees)
descriptor = ctypes.c_void_p()
if not security._CONVERT_SDDL(sddl, 1, ctypes.byref(descriptor), None):
raise Failure()
handle = None
try:
handle = security._CREATE_FILE(str(path), 0x60000, 7, None, 3, 0x2200000, None)
if handle == ctypes.c_void_p(-1).value:
raise Failure()
info = security._BY_HANDLE_FILE_INFORMATION()
if (not security._GET_FILE_INFORMATION(handle, ctypes.byref(info))
or info.dwFileAttributes & 0x400
or not security._SET_KERNEL_OBJECT_SECURITY(handle, 0x80000004, descriptor)):
raise Failure()
finally:
if handle is not None and handle != ctypes.c_void_p(-1).value:
security._CLOSE_HANDLE(handle)
security._LOCAL_FREE(descriptor)
_verify_private_acl(path, security, directory)
def _prepare_output(output, security):
if output.parent != IMPORTS_ROOT:
raise Failure()
_check_chain(IMPORTS_ROOT.parent)
_check_type(IMPORTS_ROOT.parent.lstat(), directory=True)
try:
IMPORTS_ROOT.mkdir()
except FileExistsError:
_check_chain(IMPORTS_ROOT)
_check_type(IMPORTS_ROOT.lstat(), directory=True)
else:
_secure_path(IMPORTS_ROOT, security, directory=True)
output.mkdir() # Never reuse, overwrite or repair an existing snapshot.
_secure_path(output, security, directory=True)
@contextlib.contextmanager
def _output_file(path, security):
_check_chain(path.parent)
with path.open('xb', buffering=0) as handle:
_secure_path(path, security)
info = _file_info(path)
opened = os.fstat(handle.fileno())
if (info.st_dev, info.st_ino) != (opened.st_dev, opened.st_ino):
raise Failure()
yield handle
handle.flush()
os.fsync(handle.fileno())
def _destination(path, used, directories):
parts = path.split('/')
if (not parts or any(not p or p in ('.', '..') for p in parts)
or any(ord(c) < 32 or ord(c) == 127 for c in path)
or '\\' in path or ':' in path):
raise Failure()
folded = path.casefold()
parents = {'/'.join(parts[:i]).casefold() for i in range(1, len(parts))}
if folded in used or folded in directories or parents.intersection(used):
raise Failure()
used.add(folded)
directories.update(parents)
def _excluded(relative, control=False):
parts = relative.casefold().split('/')
name = parts[-1]
if any(p in {'.git', '.opencode', 'tests', '__pycache__', '.pytest_cache',
'.mypy_cache', '.ruff_cache', 'node_modules'} for p in parts):
return 'code_cache_or_unreviewed_code'
if 'gharchive_cache' in parts:
return 'downloaded_archives'
if any(fnmatch.fnmatchcase(p, '*.lock*') or p.endswith('.pid') for p in parts):
return 'locks_or_pids'
if name.endswith(('.pyc', '.pyo')):
return 'code_cache_or_unreviewed_code'
if control and any(re.search(
r'supervisor|instance|token|authority|manifest|handshake|capability|(?:^|[._-])(?:pid|lock)(?:[._-]|$)',
p) for p in parts[2:]):
return 'control_authority'
projection = any(p in ('results', 'keychecks', KEYCHECK_COPY.casefold())
or p.startswith(('found_secrets.jsonl', 'scan_results.jsonl', 'scan_errors.log')) for p in parts)
config_backup = len(parts) == 2 and parts[0] == 'app' and name.startswith(('config.yaml.', 'secrets.yaml.'))
if (fnmatch.fnmatchcase(name, '*.log*') and not projection
and not config_backup
and not fnmatch.fnmatchcase(name, 'scan_errors.log*')
and 'publication-ledger.sqlite3' not in name):
return 'ordinary_logs'
return None
def _inventory():
files, watched, excluded, used, parents = {}, {}, {}, set(), set()
def visit(path, relative, target, control=False, expect_directory=None):
reason = _excluded(relative, control)
if reason:
excluded[reason] = excluded.get(reason, 0) + 1
return
try:
_check_chain(path)
info = path.lstat()
except FileNotFoundError:
watched[str(path)] = None
return
directory = stat.S_ISDIR(info.st_mode)
_check_type(info, directory=directory)
if expect_directory is not None and directory != expect_directory:
raise Failure()
if directory:
watched[str(path)] = _fingerprint(info)
with os.scandir(path) as entries:
names = sorted(entry.name for entry in entries)
for name in names:
visit(path / name, relative + '/' + name, target + '/' + name, control)
if _fingerprint(path.lstat()) != watched[str(path)]:
raise Failure()
return
if control and not path.name.casefold().endswith('.json'):
excluded['non_report_control'] = excluded.get('non_report_control', 0) + 1
return
if relative.casefold().startswith('runtime/state/'):
scratch = relative.casefold().split('/')[2:]
if any(fnmatch.fnmatchcase(p, 'scan_limiter*.db*')
or fnmatch.fnmatchcase(p, '*.tmp*') or p == 'janitor.cursor.json' for p in scratch):
target = 'windows-archive/' + relative
_destination(target, used, parents)
files[target] = File(path, info.st_size, _fingerprint(info))
for name in ACTIVE_DIRS:
visit(SOURCE_ROOT / 'runtime' / name, 'runtime/' + name, 'runtime-linux/' + name,
expect_directory=True)
visit(SOURCE_ROOT / 'runtime/proxy.txt', 'runtime/proxy.txt', 'runtime-linux/proxy.txt',
expect_directory=False)
for name in ('secrets.yaml', 'trufflehog-custom-detectors.yaml'):
visit(SOURCE_ROOT / 'app' / name, 'app/' + name, 'config/' + name, expect_directory=False)
for relative in ('state', 'runtime/backups', 'runtime/imports', 'runtime/' + KEYCHECK_COPY):
visit(SOURCE_ROOT / relative, relative, 'windows-archive/' + relative, expect_directory=True)
for relative in (
'app/config.yaml', 'app/.streamlit/config.toml', '.env.postgres', 'docker-compose.postgres.yml',
'runner_state.json',
'runtime/keychecks.7z', 'runtime/orkey.txt', 'runtime/check-openrouter-keys.ps1',
):
visit(SOURCE_ROOT / relative, relative, 'windows-archive/' + relative, expect_directory=False)
for folder, patterns in (
('', ('checked_*.txt', 'todo_*.txt', 'scanner.db*', 'found_secrets.jsonl*',
'scan_results.jsonl*', 'scan_errors.log*', '*.publication-ledger.sqlite3*')),
('app', ('config.yaml.*', 'secrets.yaml.*', 'scanner.db*')),
('runtime', ('*.md',)),
):
path = SOURCE_ROOT / folder
_check_chain(path)
info = path.lstat()
_check_type(info, directory=True)
watched[str(path)] = _fingerprint(info)
with os.scandir(path) as entries:
names = sorted(entry.name for entry in entries
if any(fnmatch.fnmatchcase(entry.name.casefold(), p) for p in patterns))
for name in names:
relative = folder + '/' + name if folder else name
projection_family = not folder and name.casefold().startswith(
('found_secrets.jsonl', 'scan_results.jsonl', 'scan_errors.log'))
visit(path / name, relative, 'windows-archive/' + relative,
expect_directory=None if projection_family else False)
visit(SOURCE_ROOT / 'runtime/control', 'runtime/control', 'windows-archive/runtime/control',
control=True, expect_directory=True)
_check_chain(BUNDLE_ROOT)
visit(BUNDLE_ROOT, 'scanner-result-bundles', 'scanner-result-bundles', expect_directory=True)
return Inventory(files, watched, excluded)
class HashWriter:
def __init__(self, handle):
self.handle = handle
self.digest = hashlib.sha256()
self.size = 0
def write(self, block):
if self.handle.write(block) != len(block):
raise Failure()
self.digest.update(block)
self.size += len(block)
return len(block)
def metadata(self):
return {'bytes': self.size, 'sha256': self.digest.hexdigest()}
class HashReader:
def __init__(self, handle):
self.handle = handle
self.digest = hashlib.sha256()
self.size = 0
def read(self, size):
block = self.handle.read(size)
self.digest.update(block)
self.size += len(block)
return block
def _write_tar(output, security, inventory, report):
manifest_files = []
with _output_file(output / 'files.tar', security) as handle:
writer = HashWriter(handle)
with tarfile.open(fileobj=writer, mode='w|', format=tarfile.PAX_FORMAT, copybufsize=BLOCK) as archive:
for name, entry in sorted(inventory.files.items()):
if _fingerprint(_file_info(entry.source)) != entry.fingerprint:
raise Failure()
with entry.source.open('rb', buffering=0) as source:
before = _fingerprint(os.fstat(source.fileno()))
# Windows Python 3.12 stat/fstat use different ctime bases.
# Compare ctime within each API, not across the two APIs.
if before[:6] + before[7:] != entry.fingerprint[:6] + entry.fingerprint[7:]:
raise Failure()
reader = HashReader(source)
info = tarfile.TarInfo(name)
info.size, info.mode, info.mtime = entry.size, 0o600, 0
archive.addfile(info, reader)
if (reader.size != entry.size
or _fingerprint(os.fstat(source.fileno())) != before):
raise Failure()
if _fingerprint(_file_info(entry.source)) != entry.fingerprint:
raise Failure()
manifest_files.append({'path': name, 'size': entry.size, 'sha256': reader.digest.hexdigest()})
report(8, len(manifest_files), writer.size)
metadata = writer.metadata()
return manifest_files, metadata
@contextlib.contextmanager
def _silence():
with open(os.devnull, 'w', encoding='utf-8') as sink:
with contextlib.redirect_stdout(sink), contextlib.redirect_stderr(sink):
yield
def _load_source():
app = SOURCE_ROOT / 'app'
for name in ('child_bootstrap.py', 'postgres_runtime.py', 'runtime_security.py'):
_file_info(app / name)
spec = importlib.util.spec_from_file_location('_snapshot_child_bootstrap', app / 'child_bootstrap.py')
bootstrap = importlib.util.module_from_spec(spec)
spec.loader.exec_module(bootstrap)
bootstrap._enable_dependency_paths('postgres-runtime')
sys.path.insert(0, str(app))
pg = importlib.import_module('postgres_runtime')
security = importlib.import_module('runtime_security')
for module in (pg, security):
if not _same_windows_path(module.__file__, app / (module.__name__ + '.py')):
raise Failure()
# The original loader does not overwrite inherited credentials. Remove all
# connection/path overrides first, so only the original .env can choose them.
for key in list(os.environ):
if key.upper().startswith(('PG', 'TRUF_', 'SCANNER_', 'TRUFFLEHOG_')) or key.upper() == 'DATABASE_URL':
os.environ.pop(key, None)
config_path, env_path = app / 'config.yaml', SOURCE_ROOT / '.env.postgres'
inputs = {p: _fingerprint(_file_info(p)) for p in (config_path, env_path)}
config = pg._load_config(str(config_path))
expected = {'root_dir': SOURCE_ROOT, 'project_dir': app, 'runtime_dir': SOURCE_ROOT / 'runtime',
'postgres_data_dir': POSTGRES_DATA, 'result_bundle_dir': BUNDLE_ROOT}
for key, path in expected.items():
if not _same_windows_path(config.get('global', {}).get(key, ''), path):
raise Failure()
for key, name in (('queue_dir', 'queues'), ('state_dir', 'state'), ('keycheck_dir', 'keychecks'),
('results_dir', 'results'), ('postman_cache_dir', 'postman_cache'),
('result_spool_dir', 'result_spool'), ('control_dir', 'control'), ('log_dir', 'logs')):
value = config.get('global', {}).get(key)
if value and not _same_windows_path(value, SOURCE_ROOT / 'runtime' / name):
raise Failure()
paths = pg.postgres_runtime_paths(config)
if (not _same_windows_path(paths['data_dir'], POSTGRES_DATA)
or not _same_windows_path(paths['postgres_dir'], SOURCE_ROOT / 'runtime/postgres')):
raise Failure()
security.preflight_lifecycle_paths(str(config_path), config)
loaded = pg.load_postgres_environment(str(config_path), config)
if not _same_windows_path(loaded or '', env_path):
raise Failure()
dsn = pg.canonical_database_url()
if not dsn:
raise Failure()
identity_path = SOURCE_ROOT / 'runtime/postgres/cluster_identity.json'
inputs[identity_path] = _fingerprint(_file_info(identity_path))
return SimpleNamespace(pg=pg, security=security, config=config, dsn=dsn, inputs=inputs)
def _supervisor_absent(source):
supervisor = source.config.get('supervisor', {})
paths = {SOURCE_ROOT / 'runtime/control/supervisor.instance.json',
SOURCE_ROOT / 'runtime/logs/supervisor.instance.json',
SOURCE_ROOT / 'runtime/logs/supervisor.pid'}
for key in ('instance_file',):
if supervisor.get(key):
paths.add(Path(supervisor[key]))
for key in ('control_dir', 'log_dir'):
if supervisor.get(key):
paths.add(Path(supervisor[key]) / 'supervisor.instance.json')
if any(os.path.lexists(path) for path in paths):
raise Failure()
def _validate_identity(source, identity):
values = source.pg.configured_cluster_values()
parsed = urlsplit(source.dsn)
if (identity['pg_major'] != 16 or not str(identity['system_identifier']).isdigit()
or not _same_windows_path(identity['data_directory'], POSTGRES_DATA)
or any(identity[k] != values[k] for k in ('database', 'user', 'port'))
or parsed.scheme not in ('postgresql', 'postgres') or parsed.hostname != '127.0.0.1'
or parsed.port != identity['port'] or unquote(parsed.username or '') != identity['user']
or unquote(parsed.path[1:]) != identity['database'] or parsed.password is None
or parsed.query or parsed.fragment):
raise Failure()
def _client_environment(dsn):
parsed = urlsplit(dsn)
env = {key: value for key, value in os.environ.items()
if key.upper() in {'SYSTEMROOT', 'WINDIR', 'SYSTEMDRIVE', 'TEMP', 'TMP'}}
env.update(PGHOST='127.0.0.1', PGHOSTADDR='127.0.0.1', PGPORT=str(parsed.port),
PGDATABASE=unquote(parsed.path[1:]), PGUSER=unquote(parsed.username or ''),
PGPASSWORD=unquote(parsed.password or ''), PGPASSFILE=os.devnull,
PGSERVICEFILE=os.devnull, PGSSLMODE='disable', PGGSSENCMODE='disable',
PGCONNECT_TIMEOUT='5', PGCLIENTENCODING='UTF8', PGAPPNAME='truf-windows-snapshot',
PGOPTIONS=f'-c default_transaction_read_only=on -c statement_timeout={COUNT_TIMEOUT * 1000} '
'-c lock_timeout=10000 -c idle_in_transaction_session_timeout=0 '
'-c search_path=pg_catalog -c row_security=off',
LC_ALL='C', LANG='C')
return env
@contextlib.contextmanager
def _deadline(process, seconds):
expired = threading.Event()
def expire():
expired.set()
try:
process.kill()
except OSError:
pass
timer = threading.Timer(seconds, expire)
timer.daemon = True
timer.start()
try:
yield
except BaseException:
if expired.is_set():
raise Failure(124) from None
raise
else:
if expired.is_set():
raise Failure(124)
finally:
timer.cancel()
timer.join()
@contextlib.contextmanager
def _client(command, env, timeout, interactive=False):
process = subprocess.Popen(command, stdin=subprocess.PIPE if interactive else subprocess.DEVNULL,
stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, env=env,
creationflags=0x08000000, close_fds=True, bufsize=0)
try:
with _deadline(process, timeout):
yield process
if process.stdin is not None:
process.stdin.close()
if process.wait(timeout=10) != 0:
raise Failure()
finally:
if process.poll() is None:
process.kill()
process.wait(timeout=10)
if process.stdin is not None:
process.stdin.close()
process.stdout.close()
def _query(process, sql, timeout=QUERY_TIMEOUT):
with _deadline(process, timeout):
payload = (sql + '\n').encode('utf-8')
if process.stdin.write(payload) != len(payload):
raise Failure()
process.stdin.flush()
line = process.stdout.readline(MAX_METADATA + 1)
if not line.endswith(b'\n') or len(line) > MAX_METADATA:
raise Failure()
try:
return json.loads(line)
except (ValueError, UnicodeError):
raise Failure() from None
OTHER_CLIENTS = """(SELECT count(*) FROM pg_catalog.pg_stat_activity
WHERE backend_type = 'client backend' AND pid <> pg_catalog.pg_backend_pid())"""
DATABASE_METADATA = """BEGIN ISOLATION LEVEL REPEATABLE READ READ ONLY;
SELECT pg_catalog.json_build_object(
'version_num', current_setting('server_version_num')::integer,
'system_identifier', (SELECT system_identifier::text FROM pg_catalog.pg_control_system()),
'database_name', current_database(), 'user_name', current_user,
'port', current_setting('port')::integer, 'data_directory', current_setting('data_directory'),
'in_recovery', pg_is_in_recovery(), 'read_only', current_setting('transaction_read_only'),
'all_sessions_visible', (SELECT rolsuper FROM pg_catalog.pg_roles WHERE rolname = current_user)
OR pg_has_role(current_user, 'pg_read_all_stats', 'MEMBER'),
'snapshot', pg_export_snapshot(), 'database_bytes', pg_database_size(current_database()),
'tables', (SELECT COALESCE(json_agg(c.relname ORDER BY c.relname), '[]'::json)
FROM pg_catalog.pg_class c JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = 'public' AND c.relkind = 'r'),
'sequences', (SELECT COALESCE(jsonb_agg(jsonb_build_array(n.nspname, c.relname)
ORDER BY n.nspname, c.relname), '[]'::jsonb)
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_'),
'other_clients', """ + OTHER_CLIENTS + ');'
def _validate_database(metadata, identity):
if (type(metadata.get('version_num')) is not int or metadata['version_num'] // 10000 != 16
or metadata.get('system_identifier') != identity['system_identifier']
or metadata.get('database_name') != identity['database']
or metadata.get('user_name') != identity['user'] or type(metadata.get('port')) is not int
or metadata['port'] != identity['port']
or not _same_windows_path(metadata.get('data_directory', ''), POSTGRES_DATA)
or metadata.get('in_recovery') is not False or metadata.get('read_only') != 'on'
or metadata.get('all_sessions_visible') is not True
or type(metadata.get('other_clients')) is not int or metadata['other_clients'] != 0
or not re.fullmatch(r'[0-9A-Fa-f]+-[0-9A-Fa-f]+-[0-9]+', metadata.get('snapshot', ''))
or type(metadata.get('database_bytes')) is not int or metadata['database_bytes'] < 0):
raise Failure()
tables = metadata.get('tables')
if (not isinstance(tables, list) or any(not isinstance(t, str) or not t or '\0' in t for t in tables)
or len(tables) != len(set(tables))):
raise Failure()
sequences = metadata.get('sequences')
if (not isinstance(sequences, list)
or any(not isinstance(pair, list) or len(pair) != 2
or any(not isinstance(name, str) or not name or '\0' in name for name in pair)
for pair in sequences)
or len(sequences) != len({tuple(pair) for pair in sequences})):
raise Failure()
def _sequence_states(process, sequences):
states = {}
for schema, name in sequences:
quoted = '.'.join('"' + part.replace('"', '""') + '"' for part in (schema, name))
value = _query(process, "SELECT pg_catalog.json_build_object('last_value', last_value, "
"'is_called', is_called) FROM " + quoted + ';')
if (not isinstance(value, dict) or set(value) != {'last_value', 'is_called'}
or type(value['last_value']) is not int or type(value['is_called']) is not bool):
raise Failure()
states.setdefault(schema, {})[name] = value
return states
def _no_other_clients(process):
# Activity views cache within a transaction; explicitly refresh before checking.
_query(process, "SELECT json_build_object('cleared', pg_stat_clear_snapshot() IS NULL);")
count = _query(process, 'SELECT ' + OTHER_CLIENTS + ';')
if type(count) is not int or count != 0:
raise Failure()
def _capture_database(source, identity, output, inventory, report):
report(6, 0, 0)
env = _client_environment(source.dsn)
binaries = SOURCE_ROOT / 'runtime/postgres/pgsql/bin'
for name in ('psql', 'pg_dump'):
path = binaries / (name + '.exe')
before = _fingerprint(_file_info(path))
result = subprocess.run([str(path), '--version'], stdin=subprocess.DEVNULL, stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL, timeout=15, env=env, creationflags=0x08000000,
close_fds=True)
if (result.returncode != 0 or not re.fullmatch(
rb'(?:psql|pg_dump) \(PostgreSQL\) 16(?:\.[0-9]+)*(?: \([^\r\n]*\))?\s*', result.stdout)
or _fingerprint(_file_info(path)) != before):
raise Failure()
source.inputs[path] = before
command = [str(binaries / 'psql.exe'), '-X', '-q', '-A', '-t', '-w', '-v', 'ON_ERROR_STOP=1', '-f', '-']
with _client(command, env, SESSION_TIMEOUT, interactive=True) as process:
metadata = _query(process, DATABASE_METADATA)
_validate_database(metadata, identity)
file_bytes = sum(entry.size for entry in inventory.files.values())
if shutil.disk_usage(output).free < file_bytes + metadata['database_bytes'] + 512 * BLOCK:
raise Failure()
counts = {}
for name in metadata['tables']:
quoted = '"' + name.replace('"', '""') + '"'
count = _query(process, 'SELECT count(*) FROM ONLY "public".' + quoted + ';', COUNT_TIMEOUT)
if type(count) is not int or count < 0:
raise Failure()
counts[name] = count
report(6, len(counts), 0)
_no_other_clients(process)
# Sequences are not MVCC-isolated, even in this exported-snapshot session.
# With the source stopped, require their values to stay fixed across dump.
sequences = _sequence_states(process, metadata['sequences'])
dump = [str(binaries / 'pg_dump.exe'), '--format=custom', '--no-owner', '--no-acl',
'--no-tablespaces', '--compress=1', '--no-password', '--lock-wait-timeout=10s',
'--snapshot=' + metadata['snapshot']]
dump_env = dict(env, PGOPTIONS=env['PGOPTIONS'].replace(
f'statement_timeout={COUNT_TIMEOUT * 1000}', f'statement_timeout={DUMP_TIMEOUT * 1000}'))
with _output_file(output / 'database.dump', source.security) as handle:
writer = HashWriter(handle)
prefix = b''
with _client(dump, dump_env, DUMP_TIMEOUT) as dumping:
while True:
block = dumping.stdout.read(BLOCK)
if not block:
break
if len(prefix) < 5:
prefix = (prefix + block)[:5]
writer.write(block)
if prefix != b'PGDMP' or writer.size <= 5:
raise Failure()
dump_metadata = writer.metadata()
_no_other_clients(process)
if _sequence_states(process, metadata['sequences']) != sequences:
raise Failure()
report(6, len(counts), dump_metadata['bytes'])
return {key: metadata[key] for key in ('version_num', 'system_identifier', 'database_name',
'user_name', 'port', 'data_directory', 'database_bytes')} | {
'table_counts': counts, 'table_count_mode': 'ONLY public ordinary tables; shared exported snapshot',
'sequence_states': sequences, 'sequence_count': len(metadata['sequences']),
'sequence_state_mode': 'Non-system schemas; non-MVCC values checked unchanged across dump in export session',
'schema_migration_counts': {name: counts[name] for name in ('runtime_schema_migrations', 'schema_migrations')
if name in counts}, **dump_metadata}
def _stop_confirmed(source, backend, report):
retries = 0
while True:
try:
with _silence():
result = source.pg.maintenance_stop(source.config, backend=backend)
if not result.completed or not result.stopped or backend.probe().kind != source.pg.ProbeKind.STOPPED:
raise Failure()
return
except BaseException:
# Even Ctrl-C must not release authority over a possibly live source.
retries += 1
try:
report(7, retries, 0)
time.sleep(2)
except BaseException:
pass
def _unchanged_inputs(source):
for path, fingerprint in source.inputs.items():
if _fingerprint(_file_info(path)) != fingerprint:
raise Failure()
def _publish_manifest(output, security, manifest):
temporary = output / 'manifest.json.partial'
with _output_file(temporary, security) as handle:
writer = HashWriter(handle)
for chunk in json.JSONEncoder(ensure_ascii=True, sort_keys=True, indent=2).iterencode(manifest):
writer.write(chunk.encode('utf-8'))
writer.write(b'\n')
# Windows rename refuses an existing destination. A partial JSON is not valid
# snapshot authority, even when all preceding large files were completed.
os.rename(temporary, output / 'manifest.json')
def capture(output, report):
with _defer_signals() as checkpoint:
def progress(number, count, size):
if number != 7:
checkpoint()
report(number, count, size)
_capture(output, progress, checkpoint)
def _capture(output, report, checkpoint):
report(2, 0, 0)
with _silence():
source = _load_source()
report(3, 0, 0)
with _silence():
_prepare_output(output, source.security)
authority = source.security.ClusterAuthorityLock(source.config, endpoint_dsn=source.dsn)
authority.acquire()
backend, attempted, manifest, published = None, False, None, False
try:
report(4, 0, 0)
with _silence():
_supervisor_absent(source)
identity = source.pg.verify_cluster_identity(source.config)
_validate_identity(source, identity)
backend = source.pg.PostgresBackend(source.config)
if backend.probe().kind != source.pg.ProbeKind.STOPPED:
raise Failure()
inventory = _inventory()
_unchanged_inputs(source)
report(4, len(inventory.files), sum(entry.size for entry in inventory.files.values()))
try:
report(5, 0, 0)
with _silence():
_supervisor_absent(source)
if backend.probe().kind != source.pg.ProbeKind.STOPPED:
raise Failure()
checkpoint()
attempted = True
# Never check cancellation inside the original start helper:
# Popen precedes _accepted_start_at_monotonic. Its non-raising
# signal fence lets that bookkeeping and close() finish first.
if source.pg.maintenance_start(source.config, backend=backend).kind != source.pg.ProbeKind.READY:
raise Failure()
checkpoint()
database = _capture_database(source, identity, output, inventory, report)
finally:
if attempted:
report(7, 0, 0)
_stop_confirmed(source, backend, report)
checkpoint()
files, archive = _write_tar(output, source.security, inventory, report)
report(9, len(files), archive['bytes'])
manifest = {
'format': 'truf-windows-snapshot-v1', 'created_at': datetime.now(timezone.utc).isoformat(),
'database': database, 'files': files, 'archive': archive,
'source': {'root': str(SOURCE_ROOT), 'postgres_data_dir': str(POSTGRES_DATA),
'supervisor_stopped': True, 'postgres_stopped': True},
'exclusions': {'policy': EXCLUSION_POLICY, 'observed_entries': inventory.exclusions},
'counts': {'files': len(files), 'file_bytes': sum(entry['size'] for entry in files),
'active_files': sum(not entry['path'].startswith('windows-archive/') for entry in files),
'archival_files': sum(entry['path'].startswith('windows-archive/') for entry in files),
'public_tables': len(database['table_counts']), 'sequences': database['sequence_count']},
}
finally:
# Do not use the lock's __exit__: an unconfirmed stop must retain it.
if attempted:
_stop_confirmed(source, backend, report)
try:
try:
if manifest is not None:
checkpoint()
with _silence():
_supervisor_absent(source)
_unchanged_inputs(source)
_verify_private_acl(output, source.security, directory=True)
if _inventory() != inventory:
raise Failure()
report(10, len(manifest['files']), manifest['archive']['bytes'])
_publish_manifest(output, source.security, manifest)
published = True
finally:
with _silence():
try:
if backend is not None:
backend.close()
finally:
authority.release()
checkpoint()
except BaseException:
if published:
(output / 'manifest.json').unlink()
raise
def main(argv=None):
phase = 1
def report(number, count=0, size=0):
nonlocal phase
phase = number
print(number, count, size, flush=True)
environment, paths = os.environ.copy(), sys.path[:]
try:
if os.name != 'nt' or not (sys.flags.isolated and sys.flags.no_site and sys.dont_write_bytecode):
raise Failure()
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('action', choices=('capture',))
parser.add_argument('--output', required=True)
with _silence():
args = parser.parse_args(argv)
output = _output_path(args.output)
capture(output, report)
return 0
except BaseException as exc:
timed_out = isinstance(exc, subprocess.TimeoutExpired) or isinstance(exc, Failure) and exc.code == 124
code = 124 if timed_out else 1
print(phase, code, file=sys.stderr, flush=True)
return code
finally:
os.environ.clear()
os.environ.update(environment)
sys.path[:] = paths
if __name__ == '__main__':
raise SystemExit(main())