diff --git a/src/baudolo/restore/__main__.py b/src/baudolo/restore/__main__.py index d95f635..a2a3d6b 100644 --- a/src/baudolo/restore/__main__.py +++ b/src/baudolo/restore/__main__.py @@ -27,6 +27,20 @@ def _add_common_backup_args(p: argparse.ArgumentParser) -> None: ) +def _add_common_engine_args(p: argparse.ArgumentParser) -> None: + p.add_argument("--container", required=True) + p.add_argument("--db-password", required=True) + p.add_argument("--empty", action="store_true") + p.add_argument( + "--no-version-check", + action="store_true", + help=( + "Replay even if the dump comes from a newer engine than the target. " + "With --empty this can leave an emptied database behind." + ), + ) + + def main(argv: list[str] | None = None) -> int: parser = argparse.ArgumentParser( prog="baudolo-restore", @@ -48,17 +62,15 @@ def main(argv: list[str] | None = None) -> int: p_pg = sub.add_parser("postgres", help="Restore a single PostgreSQL database dump") _add_common_backup_args(p_pg) - p_pg.add_argument("--container", required=True) + _add_common_engine_args(p_pg) p_pg.add_argument("--db-name", required=True) p_pg.add_argument("--db-user", default=None, help="Defaults to db-name if omitted") - p_pg.add_argument("--db-password", required=True) - p_pg.add_argument("--empty", action="store_true") p_cluster = sub.add_parser( "cluster", help="Restore a full PostgreSQL cluster dump (pg_dumpall)" ) _add_common_backup_args(p_cluster) - p_cluster.add_argument("--container", required=True) + _add_common_engine_args(p_cluster) p_cluster.add_argument( "--instance", required=True, @@ -69,18 +81,14 @@ def main(argv: list[str] | None = None) -> int: required=True, help="Superuser of the instance; the dump creates roles and databases", ) - p_cluster.add_argument("--db-password", required=True) - p_cluster.add_argument("--empty", action="store_true") p_mdb = sub.add_parser( "mariadb", help="Restore a single MariaDB/MySQL-compatible dump" ) _add_common_backup_args(p_mdb) - p_mdb.add_argument("--container", required=True) + _add_common_engine_args(p_mdb) p_mdb.add_argument("--db-name", required=True) p_mdb.add_argument("--db-user", default=None, help="Defaults to db-name if omitted") - p_mdb.add_argument("--db-password", required=True) - p_mdb.add_argument("--empty", action="store_true") args = parser.parse_args(argv) @@ -116,6 +124,7 @@ def main(argv: list[str] | None = None) -> int: backups_dir=args.backups_dir, ).sql_file(args.db_name), empty=args.empty, + check_version=not args.no_version_check, ) return 0 @@ -132,6 +141,7 @@ def main(argv: list[str] | None = None) -> int: backups_dir=args.backups_dir, ).cluster_file(args.instance), empty=args.empty, + check_version=not args.no_version_check, ) return 0 @@ -150,6 +160,7 @@ def main(argv: list[str] | None = None) -> int: backups_dir=args.backups_dir, ).sql_file(args.db_name), empty=args.empty, + check_version=not args.no_version_check, ) return 0 diff --git a/src/baudolo/restore/db/cluster.py b/src/baudolo/restore/db/cluster.py index b3c7f59..cb6bc7e 100644 --- a/src/baudolo/restore/db/cluster.py +++ b/src/baudolo/restore/db/cluster.py @@ -25,6 +25,7 @@ import tempfile from collections.abc import Iterable, Iterator from ..run import docker_exec +from .version import guard CONTROL_DB = "postgres" _CLUSTER_PRECLEAN_SQL = os.path.join(os.path.dirname(__file__), "cluster_preclean.sql") @@ -65,6 +66,7 @@ def restore_cluster_sql( password: str, sql_path: str, empty: bool, + check_version: bool = True, ) -> None: """Replay a pg_dumpall stream into a running instance. @@ -78,10 +80,21 @@ def restore_cluster_sql( replay stops at the first object that already exists, which is the honest outcome: recreating a cluster over a populated one is a decision, not a default. + check_version: refuse a dump from a newer major version than the + running engine before anything is dropped. """ if not os.path.isfile(sql_path): raise FileNotFoundError(sql_path) + if check_version: + guard( + sql_path=sql_path, + engine="postgres", + container=container, + user=user, + password=password, + ) + docker_env = {"PGPASSWORD": password} if empty: diff --git a/src/baudolo/restore/db/mariadb.py b/src/baudolo/restore/db/mariadb.py index d04d07c..ce7baed 100644 --- a/src/baudolo/restore/db/mariadb.py +++ b/src/baudolo/restore/db/mariadb.py @@ -4,6 +4,7 @@ import os import sys from ..run import docker_exec, docker_exec_sh +from .version import guard def _pick_client(container: str) -> str: @@ -37,12 +38,23 @@ def restore_mariadb_sql( password: str, sql_path: str, empty: bool, + check_version: bool = True, ) -> None: client = _pick_client(container) if not os.path.isfile(sql_path): raise FileNotFoundError(sql_path) + if check_version: + guard( + sql_path=sql_path, + engine="mariadb", + container=container, + user=user, + password=password, + client=client, + ) + if empty: # Do not hardcode 'mysql': MariaDB 11 images may not ship that binary. result = docker_exec( diff --git a/src/baudolo/restore/db/postgres.py b/src/baudolo/restore/db/postgres.py index 1759ae4..27ff59a 100644 --- a/src/baudolo/restore/db/postgres.py +++ b/src/baudolo/restore/db/postgres.py @@ -5,6 +5,7 @@ import tempfile from collections.abc import Iterable, Iterator from ..run import docker_exec +from .version import guard _SUPERUSER_ONLY_PREFIXES = (b"COMMENT ON EXTENSION", b"ALTER DEFAULT PRIVILEGES") _EMPTY_PRECLEAN_SQL = os.path.join(os.path.dirname(__file__), "empty_preclean.sql") @@ -46,10 +47,20 @@ def restore_postgres_sql( password: str, sql_path: str, empty: bool, + check_version: bool = True, ) -> None: if not os.path.isfile(sql_path): raise FileNotFoundError(sql_path) + if check_version: + guard( + sql_path=sql_path, + engine="postgres", + container=container, + user=user, + password=password, + ) + docker_env = {"PGPASSWORD": password} if empty: diff --git a/src/baudolo/restore/db/version.py b/src/baudolo/restore/db/version.py new file mode 100644 index 0000000..0db51eb --- /dev/null +++ b/src/baudolo/restore/db/version.py @@ -0,0 +1,146 @@ +"""Refuse a dump the target engine is too old to read. + +A restore with ``--empty`` destroys before it replays: the pre-clean drops the +schema in one session and the dump goes in the next, with no rollback across +the two. A dump the engine cannot parse therefore does not fail harmlessly - +it leaves an emptied database behind. Comparing the two versions first turns +that into a refusal. + +Both engines state their origin in the dump's own header, and they do not +state it the same way. Postgres writes ``-- Dumped from database version`` +around line seven. MariaDB opens line two with ``-- MariaDB dump 10.19-11.8.8``, +where the first number is mariadb-dump's own version, and names the server only +further down on the tab-separated ``-- Server version`` line. Matching the first +number in the header would read the tool on one engine and the server on the +other, so each engine gets its own pattern. + +A ``pg_dumpall`` cluster dump has no version line of its own: its header opens +with the cluster banner and the roles section, and the first +``-- Dumped from database version`` belongs to the first database's embedded +``pg_dump`` output, arbitrarily far down. Hence the scan runs to +``SCAN_LINES`` rather than to a header-sized handful. +""" + +from __future__ import annotations + +import re + +from ..run import docker_exec, stdout_of + +SCAN_LINES = 2000 +DUMP_VERSION = { + "postgres": re.compile(r"^-- Dumped from database version (\S+)"), + "mariadb": re.compile(r"^-- Server version\s+(\S+)"), +} + + +class VersionMismatch(Exception): + """The dump cannot be replayed into this engine.""" + + +def major_of(version: str) -> int: + """The major number of an engine version string. + + Args: + version: as the engine spells it, e.g. ``17.11`` or + ``11.8.8-MariaDB-ubu2404``. + + Raises: + VersionMismatch: the string does not start with a number. + """ + leading = re.match(r"(\d+)", version) + if not leading: + raise VersionMismatch(f"cannot read a major version from '{version}'") + return int(leading.group(1)) + + +def dump_version(sql_path: str, engine: str) -> str: + """Read the engine version a dump was taken from, out of its own header. + + Args: + sql_path: the dump to read. + engine: ``postgres`` or ``mariadb``. + + Returns: + The version string as the dump spells it. + + Raises: + VersionMismatch: no version line within the first ``SCAN_LINES``. + """ + pattern = DUMP_VERSION[engine] + with open(sql_path, encoding="utf-8", errors="replace") as handle: + for _ in range(SCAN_LINES): + line = handle.readline() + if not line: + break + found = pattern.search(line) + if found: + return found.group(1) + raise VersionMismatch( + f"{sql_path} carries no {engine} version header in its first {SCAN_LINES} lines" + ) + + +def server_version( + container: str, engine: str, user: str, password: str, client: str = "" +) -> str: + """Ask the running engine which version it is.""" + if engine == "postgres": + return stdout_of( + docker_exec( + container, + ["psql", "-U", user, "-tAc", "SHOW server_version"], + capture=True, + docker_env={"PGPASSWORD": password}, + ) + ) + return stdout_of( + docker_exec( + container, + [ + client or "mariadb", + "-u", + user, + f"--password={password}", + "-N", + "-B", + "-e", + "SELECT VERSION()", + ], + capture=True, + ) + ) + + +def assert_replayable(sql_path: str, engine: str, dumped: str, serving: str) -> None: + """Refuse a dump from a newer major version than the target engine. + + Restoring forward across a major version is the upgrade path and stays + allowed; backward is refused, because a newer dump uses syntax an older + server rejects and the pre-clean would already have dropped the schema. + + Raises: + VersionMismatch: the dump is newer than the engine. + """ + if major_of(dumped) > major_of(serving): + raise VersionMismatch( + f"{sql_path} came from {engine} {dumped} but {serving} is running; " + "a newer dump does not replay into an older engine, and --empty " + "would drop the schema before finding out" + ) + + +def guard( + *, + sql_path: str, + engine: str, + container: str, + user: str, + password: str, + client: str = "", +) -> None: + """Compare the dump's origin against the running engine before replaying.""" + dumped = dump_version(sql_path, engine) + serving = server_version(container, engine, user, password, client) + assert_replayable(sql_path, engine, dumped, serving) + print(f"OK: dump is from {engine} {dumped}, {serving} is serving.") diff --git a/src/baudolo/restore/run.py b/src/baudolo/restore/run.py index 6c08416..dcc2ac5 100644 --- a/src/baudolo/restore/run.py +++ b/src/baudolo/restore/run.py @@ -40,6 +40,12 @@ def run( raise +def stdout_of(completed: subprocess.CompletedProcess) -> str: + """The captured stdout as stripped text, whether it came back bytes or str.""" + raw = completed.stdout or b"" + return (raw.decode() if isinstance(raw, bytes) else raw).strip() + + def docker_exec( container: str, argv: list[str], diff --git a/tests/e2e/test_e2e_version_gate.py b/tests/e2e/test_e2e_version_gate.py new file mode 100644 index 0000000..5211b3f --- /dev/null +++ b/tests/e2e/test_e2e_version_gate.py @@ -0,0 +1,283 @@ +"""A dump from a newer engine must be refused before --empty destroys anything. + +The pre-clean and the replay are two separate sessions with no rollback across +them, so a dump the engine cannot parse leaves an emptied database behind. The +decisive assertion here is not the non-zero exit - it is that the payload is +still readable afterwards. +""" + +import re +import unittest +from pathlib import Path + +from .helpers import ( + MARIADB_DATA_DIR, + MARIADB_IMAGE, + POSTGRES_DATA_DIR, + POSTGRES_IMAGE, + backup_path, + backup_run, + cleanup_docker, + create_minimal_compose_dir, + ensure_empty_dir, + latest_version_dir, + require_docker, + run, + unique, + wait_for_mariadb, + wait_for_mariadb_sql, + wait_for_postgres, + write_databases_csv, +) + +PAYLOAD = "gate-payload" +FUTURE = "99.0" + + +def rewrite_version(dump: Path, pattern: str, version: str) -> str: + """Make the dump claim ``version``; return what it claimed before.""" + text = dump.read_text(encoding="utf-8", errors="replace") + found = re.search(pattern, text) + if not found: + raise AssertionError(f"{dump} carries no version header matching {pattern}") + claimed = found.group(1) + dump.write_text( + text.replace(found.group(0), found.group(0).replace(claimed, version), 1), + encoding="utf-8", + ) + return claimed + + +class GateCase: + """Drive one engine through refusal, escape hatch and truthful replay.""" + + engine = "" + pattern = "" + + @classmethod + def restore(cls, *extra: str): + return run( + [ + "baudolo-restore", + cls.engine, + cls.volume, + cls.hash, + cls.version, + "--backups-dir", + cls.backups_dir, + "--repo-name", + cls.repo_name, + "--container", + cls.container, + "--db-name", + cls.db_name, + "--db-user", + cls.db_user, + "--db-password", + cls.db_password, + "--empty", + *extra, + ], + check=False, + ) + + @classmethod + def prepare(cls) -> None: + cls.backups_dir = f"/tmp/{cls.prefix}/Backups" + ensure_empty_dir(cls.backups_dir) + cls.compose_dir = create_minimal_compose_dir(f"/tmp/{cls.prefix}") + cls.repo_name = cls.prefix + cls.databases_csv = f"/tmp/{cls.prefix}/databases.csv" + write_databases_csv( + cls.databases_csv, + [(cls.container, cls.db_name, cls.db_user, cls.db_password)], + ) + backup_run( + backups_dir=cls.backups_dir, + repo_name=cls.repo_name, + compose_dir=cls.compose_dir, + databases_csv=cls.databases_csv, + database_containers=[cls.container], + images_no_stop_required=[cls.image], + ) + cls.hash, cls.version = latest_version_dir(cls.backups_dir, cls.repo_name) + cls.dump = ( + backup_path(cls.backups_dir, cls.repo_name, cls.version, cls.volume) + / "sql" + / f"{cls.db_name}.backup.sql" + ) + + cls.truthful_version = rewrite_version(cls.dump, cls.pattern, FUTURE) + cls.refused = cls.restore() + cls.payload_after_refusal = cls.read_payload() + + cls.forced = cls.restore("--no-version-check") + cls.payload_after_force = cls.read_payload() + + rewrite_version(cls.dump, cls.pattern, cls.truthful_version) + cls.replayed = cls.restore() + cls.payload_after_replay = cls.read_payload() + + def test_the_dump_states_the_engine_it_came_from(self) -> None: + self.assertRegex(self.truthful_version, r"^\d+") + + def test_a_newer_dump_is_refused(self) -> None: + self.assertNotEqual(self.refused.returncode, 0, self.refused.stdout) + + def test_the_refusal_names_the_version_it_refused(self) -> None: + self.assertIn(FUTURE, self.refused.stderr) + self.assertIn("older engine", self.refused.stderr) + + def test_the_refusal_left_the_data_untouched(self) -> None: + self.assertEqual( + self.payload_after_refusal, + PAYLOAD, + "--empty pre-cleaned before the version was checked", + ) + + def test_the_escape_hatch_replays_anyway(self) -> None: + self.assertEqual(self.forced.returncode, 0, self.forced.stderr) + self.assertEqual(self.payload_after_force, PAYLOAD) + + def test_a_truthful_dump_replays(self) -> None: + self.assertEqual(self.replayed.returncode, 0, self.replayed.stderr) + self.assertEqual(self.payload_after_replay, PAYLOAD) + + +class TestE2EPostgresVersionGate(GateCase, unittest.TestCase): + engine = "postgres" + pattern = r"-- Dumped from database version (\S+)" + + @classmethod + def setUpClass(cls) -> None: + require_docker() + cls.prefix = unique("baudolo-e2e-pg-gate") + cls.container = f"{cls.prefix}-pg" + cls.volume = f"{cls.prefix}-pg-vol" + cls.image = POSTGRES_IMAGE + cls.db_name = "appdb" + cls.db_user = "postgres" + cls.db_password = "pgpw" + + run(["docker", "volume", "create", cls.volume]) + run( + [ + "docker", + "run", + "-d", + "--name", + cls.container, + "-e", + f"POSTGRES_PASSWORD={cls.db_password}", + "-v", + f"{cls.volume}:{POSTGRES_DATA_DIR}", + POSTGRES_IMAGE, + ] + ) + wait_for_postgres(cls.container, user=cls.db_user) + cls.sql("postgres", f"CREATE DATABASE {cls.db_name}") + cls.sql( + cls.db_name, + f"CREATE TABLE t (v text); INSERT INTO t VALUES ('{PAYLOAD}');", + ) + cls.prepare() + + @classmethod + def tearDownClass(cls) -> None: + cleanup_docker(containers=[cls.container], volumes=[cls.volume]) + + @classmethod + def sql(cls, database: str, statement: str) -> str: + p = run( + [ + "docker", + "exec", + cls.container, + "sh", + "-lc", + f'psql -U {cls.db_user} -d {database} -t -A -c "{statement}"', + ], + check=False, + ) + return (p.stdout or "").strip() + + @classmethod + def read_payload(cls) -> str: + return cls.sql(cls.db_name, "SELECT v FROM t") + + +class TestE2EMariadbVersionGate(GateCase, unittest.TestCase): + engine = "mariadb" + pattern = r"-- Server version\s+(\S+)" + + @classmethod + def setUpClass(cls) -> None: + require_docker() + cls.prefix = unique("baudolo-e2e-mdb-gate") + cls.container = f"{cls.prefix}-mdb" + cls.volume = f"{cls.prefix}-mdb-vol" + cls.image = MARIADB_IMAGE + cls.db_name = "appdb" + cls.db_user = "test" + cls.db_password = "testpw" + + run(["docker", "volume", "create", cls.volume]) + run( + [ + "docker", + "run", + "-d", + "--name", + cls.container, + "-e", + "MARIADB_ROOT_PASSWORD=rootpw", + "-e", + f"MARIADB_DATABASE={cls.db_name}", + "-e", + f"MARIADB_USER={cls.db_user}", + "-e", + f"MARIADB_PASSWORD={cls.db_password}", + "-v", + f"{cls.volume}:{MARIADB_DATA_DIR}", + MARIADB_IMAGE, + ] + ) + wait_for_mariadb(cls.container, root_password="rootpw", timeout_s=90) + wait_for_mariadb_sql( + cls.container, user=cls.db_user, password=cls.db_password, timeout_s=90 + ) + cls.sql( + f"CREATE TABLE {cls.db_name}.t (v VARCHAR(50)); " + f"INSERT INTO {cls.db_name}.t VALUES ('{PAYLOAD}');" + ) + cls.prepare() + + @classmethod + def tearDownClass(cls) -> None: + cleanup_docker(containers=[cls.container], volumes=[cls.volume]) + + @classmethod + def sql(cls, statement: str) -> str: + p = run( + [ + "docker", + "exec", + cls.container, + "sh", + "-lc", + ( + f"mariadb -h 127.0.0.1 -u{cls.db_user} -p{cls.db_password} " + f'-N -B -e "{statement}"' + ), + ], + check=False, + ) + return (p.stdout or "").strip() + + @classmethod + def read_payload(cls) -> str: + return cls.sql(f"SELECT v FROM {cls.db_name}.t") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/restore/test_cli_version_flag.py b/tests/unit/restore/test_cli_version_flag.py new file mode 100644 index 0000000..880800f --- /dev/null +++ b/tests/unit/restore/test_cli_version_flag.py @@ -0,0 +1,52 @@ +import unittest +from unittest.mock import patch + +from baudolo.restore import __main__ as cli + +ENGINES = { + "postgres": ("restore_postgres_sql", ["--db-name", "app"]), + "mariadb": ("restore_mariadb_sql", ["--db-name", "app"]), + "cluster": ("restore_cluster_sql", ["--instance", "central", "--db-user", "root"]), +} + + +class TestVersionFlagReachesEveryEngine(unittest.TestCase): + def call(self, engine: str, extra: list) -> dict: + target, required = ENGINES[engine] + argv = [ + engine, + "app_vol", + "hash", + "20260817000000", + "--container", + "db", + "--db-password", + "pw", + *required, + *extra, + ] + with patch.object(cli, target) as restore: + self.assertEqual(cli.main(argv), 0) + return restore.call_args.kwargs + + def test_the_gate_is_on_by_default(self) -> None: + for engine in ENGINES: + with self.subTest(engine=engine): + self.assertTrue(self.call(engine, [])["check_version"]) + + def test_the_flag_turns_it_off(self) -> None: + for engine in ENGINES: + with self.subTest(engine=engine): + kwargs = self.call(engine, ["--no-version-check"]) + self.assertFalse(kwargs["check_version"]) + + def test_empty_stays_independent_of_the_gate(self) -> None: + for engine in ENGINES: + with self.subTest(engine=engine): + kwargs = self.call(engine, ["--empty"]) + self.assertTrue(kwargs["empty"]) + self.assertTrue(kwargs["check_version"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/restore/test_cluster_replay.py b/tests/unit/restore/test_cluster_replay.py index 29e28bc..b91e66a 100644 --- a/tests/unit/restore/test_cluster_replay.py +++ b/tests/unit/restore/test_cluster_replay.py @@ -24,6 +24,7 @@ class TestClusterReplay(unittest.TestCase): password="pw", sql_path=sql.name, empty=empty, + check_version=False, ) return calls @@ -107,6 +108,7 @@ class TestClusterReplay(unittest.TestCase): password="pw", sql_path="/nonexistent/x.cluster.backup.sql", empty=False, + check_version=False, ) def test_the_path_helper_names_the_dumpall_file(self) -> None: diff --git a/tests/unit/restore/test_mariadb_empty_drop.py b/tests/unit/restore/test_mariadb_empty_drop.py index 9e26e2b..a82ade6 100644 --- a/tests/unit/restore/test_mariadb_empty_drop.py +++ b/tests/unit/restore/test_mariadb_empty_drop.py @@ -29,6 +29,7 @@ class TestMariadbEmptyDrop(unittest.TestCase): password="pw", sql_path=sql.name, empty=True, + check_version=False, ) drop_calls = [argv for argv in calls if any("DROP TABLE" in a for a in argv)] diff --git a/tests/unit/restore/test_postgres_single_transaction.py b/tests/unit/restore/test_postgres_single_transaction.py index 5353309..5aa87d5 100644 --- a/tests/unit/restore/test_postgres_single_transaction.py +++ b/tests/unit/restore/test_postgres_single_transaction.py @@ -24,6 +24,7 @@ class TestPostgresSingleTransaction(unittest.TestCase): password="pw", sql_path=sql.name, empty=True, + check_version=False, ) self.assertEqual(len(calls), 2, f"expected pre-clean + replay: {calls}") diff --git a/tests/unit/restore/test_version_gate.py b/tests/unit/restore/test_version_gate.py new file mode 100644 index 0000000..e90e465 --- /dev/null +++ b/tests/unit/restore/test_version_gate.py @@ -0,0 +1,244 @@ +import os +import tempfile +import unittest +from unittest.mock import MagicMock, patch + +from baudolo.restore.db import cluster as cluster_mod +from baudolo.restore.db import mariadb as mdb_mod +from baudolo.restore.db import postgres as pg_mod +from baudolo.restore.db import version as ver + +POSTGRES_HEADER = """-- +-- PostgreSQL database dump +-- + +\\restrict BbyzwODc1rWKL3rDyLhEjgCF0Kf2TU5ma7gcTs8eQI7copLtydXkc61zdULsPav + +-- Dumped from database version 17.11 +-- Dumped by pg_dump version 17.11 + +SET statement_timeout = 0; +""" + +MARIADB_HEADER = """/*M!999999\\- enable the sandbox mode */ +-- MariaDB dump 10.19-11.8.8-MariaDB, for debian-linux-gnu (x86_64) +-- +-- Host: 127.0.0.1 Database: mysql +-- ------------------------------------------------------ +-- Server version\t11.8.8-MariaDB-ubu2404 + +/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */; +""" + + +def cluster_header(roles: int) -> str: + """A pg_dumpall stream: banner, N roles, then the first database's dump.""" + head = [ + "--", + "-- PostgreSQL database cluster dump", + "--", + "", + "SET default_transaction_read_only = off;", + "", + "--", + "-- Roles", + "--", + ] + for i in range(roles): + head.append(f'CREATE ROLE "app{i}";') + head.append(f'ALTER ROLE "app{i}" WITH NOSUPERUSER INHERIT LOGIN;') + head.append("\\connect app") + head.append("") + return "\n".join(head) + "\n" + POSTGRES_HEADER + + +def dump_file(text: str) -> str: + path = os.path.join(tempfile.mkdtemp(), "app.backup.sql") + with open(path, "w", encoding="utf-8") as handle: + handle.write(text) + return path + + +class TestDumpVersion(unittest.TestCase): + """Headers captured from postgres:17-alpine and mariadb:11 themselves.""" + + def test_postgres_reads_the_source_server(self) -> None: + self.assertEqual( + ver.dump_version(dump_file(POSTGRES_HEADER), "postgres"), "17.11" + ) + + def test_mariadb_reads_the_server_not_the_dump_tool(self) -> None: + found = ver.dump_version(dump_file(MARIADB_HEADER), "mariadb") + self.assertEqual(found, "11.8.8-MariaDB-ubu2404") + self.assertNotEqual( + ver.major_of(found), + 10, + "10.19 is mariadb-dump's own version, not the server's", + ) + + def test_cluster_dump_states_its_version_far_below_the_header(self) -> None: + path = dump_file(cluster_header(roles=200)) + with open(path, encoding="utf-8") as handle: + offset = next(i for i, line in enumerate(handle) if "Dumped from" in line) + self.assertGreater(offset, 100, "fixture must exercise the deep scan") + self.assertEqual(ver.dump_version(path, "postgres"), "17.11") + + def test_a_version_beyond_the_scan_limit_is_refused_not_ignored(self) -> None: + path = dump_file(cluster_header(roles=ver.SCAN_LINES)) + with self.assertRaises(ver.VersionMismatch): + ver.dump_version(path, "postgres") + + def test_a_dump_without_a_version_header_is_refused(self) -> None: + path = dump_file("CREATE TABLE t (id int);\n") + with self.assertRaises(ver.VersionMismatch): + ver.dump_version(path, "postgres") + + +class TestMajorOf(unittest.TestCase): + def test_reads_the_leading_number(self) -> None: + self.assertEqual(ver.major_of("17.11"), 17) + self.assertEqual(ver.major_of("11.8.8-MariaDB-ubu2404"), 11) + self.assertEqual(ver.major_of("9.6.24"), 9) + self.assertEqual(ver.major_of("18beta1"), 18) + + def test_refuses_an_unreadable_version(self) -> None: + with self.assertRaises(ver.VersionMismatch): + ver.major_of("unknown") + + +class TestAssertReplayable(unittest.TestCase): + def test_newer_dump_into_older_engine_is_refused(self) -> None: + with self.assertRaises(ver.VersionMismatch) as caught: + ver.assert_replayable("/b/app.sql", "postgres", "17.11", "15.6") + self.assertIn("17.11", str(caught.exception)) + self.assertIn("15.6", str(caught.exception)) + + def test_same_major_passes(self) -> None: + ver.assert_replayable("/b/app.sql", "postgres", "17.4", "17.11") + + def test_older_dump_into_newer_engine_passes(self) -> None: + ver.assert_replayable("/b/app.sql", "postgres", "15.6", "17.11") + + +class TestServerVersion(unittest.TestCase): + def test_postgres_asks_over_pgpassword(self) -> None: + with patch.object(ver, "docker_exec") as exec_: + exec_.return_value = MagicMock(stdout=b" 17.11 \n") + found = ver.server_version("db", "postgres", "app", "pw") + self.assertEqual(found, "17.11") + argv = exec_.call_args.args[1] + self.assertIn("SHOW server_version", argv) + self.assertEqual(exec_.call_args.kwargs["docker_env"], {"PGPASSWORD": "pw"}) + + def test_mariadb_asks_through_the_detected_client(self) -> None: + with patch.object(ver, "docker_exec") as exec_: + exec_.return_value = MagicMock(stdout=b"11.8.8-MariaDB-ubu2404\n") + found = ver.server_version("db", "mariadb", "app", "pw", client="mysql") + self.assertEqual(found, "11.8.8-MariaDB-ubu2404") + self.assertEqual(exec_.call_args.args[1][0], "mysql") + + +class TestGateStopsBeforeDestroying(unittest.TestCase): + """--empty drops in one session and replays in the next, with no rollback + between them, so the refusal has to land before the first session.""" + + def setUp(self) -> None: + self.serving = patch.object(ver, "docker_exec").start() + self.addCleanup(patch.stopall) + + def serve(self, version: str) -> None: + self.serving.return_value = MagicMock(stdout=version.encode()) + + def test_postgres_refuses_without_running_the_preclean(self) -> None: + self.serve("15.6") + path = dump_file(POSTGRES_HEADER) + with ( + patch.object(pg_mod, "docker_exec") as replay, + self.assertRaises(ver.VersionMismatch), + ): + pg_mod.restore_postgres_sql( + container="db", + db_name="app", + user="app", + password="pw", + sql_path=path, + empty=True, + ) + replay.assert_not_called() + + def test_cluster_refuses_without_running_the_preclean(self) -> None: + self.serve("15.6") + path = dump_file(cluster_header(roles=3)) + with ( + patch.object(cluster_mod, "docker_exec") as replay, + self.assertRaises(ver.VersionMismatch), + ): + cluster_mod.restore_cluster_sql( + container="db", + user="postgres", + password="pw", + sql_path=path, + empty=True, + ) + replay.assert_not_called() + + def test_mariadb_refuses_without_dropping_tables(self) -> None: + self.serve("10.11.6-MariaDB") + path = dump_file(MARIADB_HEADER) + with ( + patch.object(mdb_mod, "_pick_client", return_value="mariadb"), + patch.object(mdb_mod, "docker_exec") as replay, + self.assertRaises(ver.VersionMismatch), + ): + mdb_mod.restore_mariadb_sql( + container="db", + db_name="app", + user="app", + password="pw", + sql_path=path, + empty=True, + ) + replay.assert_not_called() + + def test_matching_versions_let_the_replay_through(self) -> None: + self.serve("17.11") + path = dump_file(POSTGRES_HEADER) + with patch.object(pg_mod, "docker_exec") as replay: + pg_mod.restore_postgres_sql( + container="db", + db_name="app", + user="app", + password="pw", + sql_path=path, + empty=False, + ) + replay.assert_called_once() + + def test_the_escape_hatch_asks_the_engine_nothing(self) -> None: + path = dump_file("CREATE TABLE t (id int);\n") + with patch.object(pg_mod, "docker_exec"): + pg_mod.restore_postgres_sql( + container="db", + db_name="app", + user="app", + password="pw", + sql_path=path, + empty=False, + check_version=False, + ) + self.serving.assert_not_called() + + def test_a_missing_dump_is_reported_as_missing_not_as_a_mismatch(self) -> None: + with self.assertRaises(FileNotFoundError): + pg_mod.restore_postgres_sql( + container="db", + db_name="app", + user="app", + password="pw", + sql_path=os.path.join(tempfile.mkdtemp(), "absent.sql"), + empty=True, + ) + + +if __name__ == "__main__": + unittest.main()