Files
ARR-2.0-0918/arr_web/local_replay_database.py

117 lines
5.4 KiB
Python

"""Persistent, newly owned local replay PostgreSQL. No DSN or existing cluster input."""
from __future__ import annotations
import os
from pathlib import Path
import shlex
import shutil
import subprocess
import tempfile
from arr_web.local_xml_replay import PROJECT, document
from integrations.ohip.capture_job import atomic_json, private_directory
class ReplayDatabase:
def __init__(self, root: Path, *, schema_version: int = 18):
if schema_version not in (18, 19):
raise ValueError("unsupported_local_schema")
self.schema_version = schema_version
self.root = root
self.data = root / "postgres"
self.socket = None
self.started = False
self._saved_pg = {}
def start(self):
import psycopg
self.driver = psycopg
private_directory(self.root)
self.initdb, self.pg_ctl = shutil.which("initdb"), shutil.which("pg_ctl")
if not self.initdb or not self.pg_ctl:
raise RuntimeError("local_postgres_binaries_required")
identity = {"version": "arr-owned-replay-postgres/v1", "data_directory": str(self.data.resolve())}
marker = self.root / "database-owner.json"
if not marker.exists():
if self.data.exists():
raise RuntimeError("refusing_unowned_postgres_directory")
atomic_json(marker, identity, replace=False)
if document(marker) != identity or self.data.is_symlink():
raise RuntimeError("replay_database_owner_mismatch")
self._saved_pg = {k: v for k, v in os.environ.items() if k.startswith("PG")}
for key in self._saved_pg:
os.environ.pop(key)
self.socket = Path(tempfile.mkdtemp(prefix="arr-rpg-", dir="/tmp"))
try:
if not self.data.exists():
self._command([self.initdb, "-D", str(self.data), "-U", "arr_replay", "--auth-local=trust",
"--auth-host=reject", "--encoding=UTF8", "--locale=C"])
private_directory(self.data)
# Never stop an already running instance or recover it implicitly.
if (self.data / "postmaster.pid").exists():
raise RuntimeError("replay_database_already_running")
options = shlex.join(["-k", str(self.socket), "-p", "55433", "-c", "listen_addresses=",
"-c", "unix_socket_permissions=0700", "-c", "fsync=on"])
self._command([self.pg_ctl, "-D", str(self.data), "-l", str(self.root / "postgres.log"),
"-o", options, "-w", "-t", "20", "start"])
self.started = True
self._bootstrap()
return self
except BaseException:
self.close()
raise
@staticmethod
def _command(arguments):
result = subprocess.run(arguments, capture_output=True, timeout=35)
if result.returncode:
raise RuntimeError("owned_postgres_command_failed")
def connect(self, _dsn=None, *, database="booking_test", autocommit=False):
connection = self.driver.connect(host=str(self.socket), port=55433, dbname=database,
user="arr_replay", connect_timeout=5, autocommit=True)
try:
if (Path(connection.execute("SHOW data_directory").fetchone()[0]).resolve() != self.data.resolve()
or connection.execute("SHOW listen_addresses").fetchone()[0] != ""):
raise RuntimeError("refusing_non_replay_database")
connection.autocommit = autocommit
return connection
except BaseException:
connection.close()
raise
def _bootstrap(self):
marker = self.root / "database-ready.json"
if marker.exists():
if document(marker) != {"schema": f"008-{self.schema_version:03d}", "database": "booking_test"}:
raise RuntimeError("replay_schema_marker_invalid")
with self.connect() as connection:
connection.execute("SELECT 1 FROM finance.daily_versions LIMIT 1")
return
with self.connect(database="postgres", autocommit=True) as connection:
# An unexpected existing DB is evidence of interrupted initialization; do not erase it.
# Existing migrations explicitly guard this database name. It is a NEW
# database inside this owned Unix-socket cluster, never the user's DSN.
connection.execute("CREATE DATABASE booking_test TEMPLATE template0 ENCODING 'UTF8'")
paths = sorted(p for p in (PROJECT / "database").glob("[0-9][0-9][0-9]_*.sql")
if ".down." not in p.name and 8 <= int(p.name[:3]) <= self.schema_version)
if len(paths) != self.schema_version - 7:
raise RuntimeError("unexpected_replay_migration_set")
with self.connect() as connection:
for path in paths:
connection.execute(path.read_text(), prepare=False)
atomic_json(marker, {"schema": f"008-{self.schema_version:03d}", "database": "booking_test"}, replace=False)
def close(self):
try:
if self.started:
self._command([self.pg_ctl, "-D", str(self.data), "-m", "fast", "-w", "-t", "20", "stop"])
self.started = False
if self.socket and not self.started:
shutil.rmtree(self.socket)
self.socket = None
finally:
os.environ.update(self._saved_pg)
self._saved_pg = {}