117 lines
5.4 KiB
Python
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 = {}
|