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

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