feat(restore): refuse a dump the target engine cannot 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. The version each side is on decides that up front, so the refusal lands before the first session opens.

Both engines state their origin in the dump's own header and spell it differently. Postgres names the source server; MariaDB opens with mariadb-dump's own version and names the server further down, so matching the first number would read the tool on one engine and the server on the other. A pg_dumpall stream carries no version line of its own at all - the first belongs to the first database's embedded pg_dump output, arbitrarily far down - hence the deep scan.

Restoring forward across a major version stays allowed; that is the upgrade path. Only backward is refused, with --no-version-check as the way out.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-08-17 05:25:29 +02:00
parent f437787e64
commit bb647c66ec
12 changed files with 791 additions and 9 deletions

View File

@@ -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

View File

@@ -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:

View File

@@ -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(

View File

@@ -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:

View File

@@ -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.")

View File

@@ -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],

View File

@@ -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()

View File

@@ -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()

View File

@@ -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:

View File

@@ -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)]

View File

@@ -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}")

View File

@@ -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()