rebuild(stage-2): establish runtime and persistence foundations
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from backend.database.migrations.registry import MIGRATIONS
|
||||
from backend.database.migrations.runner import Migration, MigrationError, MigrationRunner
|
||||
|
||||
__all__ = ["MIGRATIONS", "Migration", "MigrationError", "MigrationRunner"]
|
||||
@@ -0,0 +1,3 @@
|
||||
from backend.database.migrations.runner import Migration
|
||||
|
||||
MIGRATIONS: tuple[Migration, ...] = ()
|
||||
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import sqlite3
|
||||
from collections.abc import Callable, Iterable
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from backend.database.connection import Database
|
||||
|
||||
MigrationAction = Callable[[sqlite3.Connection], None]
|
||||
|
||||
|
||||
class MigrationError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Migration:
|
||||
version: int
|
||||
name: str
|
||||
signature: str
|
||||
upgrade: MigrationAction
|
||||
downgrade: MigrationAction
|
||||
|
||||
@property
|
||||
def checksum(self) -> str:
|
||||
source = f"{self.version}:{self.name}:{self.signature}"
|
||||
return hashlib.sha256(source.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
class MigrationRunner:
|
||||
def __init__(self, database: Database) -> None:
|
||||
self._database = database
|
||||
|
||||
def upgrade(self, migrations: Iterable[Migration]) -> tuple[int, ...]:
|
||||
ordered = self._validate(migrations)
|
||||
with self._database.transaction() as connection:
|
||||
self._ensure_ledger(connection)
|
||||
applied = self._applied(connection)
|
||||
self._verify_history(ordered, applied)
|
||||
completed: list[int] = []
|
||||
for migration in ordered:
|
||||
if migration.version in applied:
|
||||
continue
|
||||
self._run_action(connection, migration, migration.upgrade, "upgrade")
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO schema_migrations (version, name, checksum, applied_at)
|
||||
VALUES (?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
migration.version,
|
||||
migration.name,
|
||||
migration.checksum,
|
||||
datetime.now(UTC).isoformat(timespec="seconds"),
|
||||
),
|
||||
)
|
||||
completed.append(migration.version)
|
||||
return tuple(completed)
|
||||
|
||||
def downgrade(self, migrations: Iterable[Migration], target_version: int) -> tuple[int, ...]:
|
||||
if target_version < 0:
|
||||
raise MigrationError("Target version cannot be negative")
|
||||
ordered = self._validate(migrations)
|
||||
by_version = {migration.version: migration for migration in ordered}
|
||||
with self._database.transaction() as connection:
|
||||
self._ensure_ledger(connection)
|
||||
applied = self._applied(connection)
|
||||
self._verify_history(ordered, applied)
|
||||
pending = sorted(
|
||||
(version for version in applied if version > target_version), reverse=True
|
||||
)
|
||||
rolled_back: list[int] = []
|
||||
for version in pending:
|
||||
migration = by_version[version]
|
||||
self._run_action(connection, migration, migration.downgrade, "downgrade")
|
||||
connection.execute("DELETE FROM schema_migrations WHERE version = ?", (version,))
|
||||
rolled_back.append(version)
|
||||
return tuple(rolled_back)
|
||||
|
||||
@staticmethod
|
||||
def _validate(migrations: Iterable[Migration]) -> tuple[Migration, ...]:
|
||||
ordered = tuple(sorted(migrations, key=lambda migration: migration.version))
|
||||
versions = [migration.version for migration in ordered]
|
||||
if any(version <= 0 for version in versions):
|
||||
raise MigrationError("Migration versions must be positive integers")
|
||||
if len(versions) != len(set(versions)):
|
||||
raise MigrationError("Migration versions must be unique")
|
||||
return ordered
|
||||
|
||||
@staticmethod
|
||||
def _ensure_ledger(connection: sqlite3.Connection) -> None:
|
||||
connection.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
checksum TEXT NOT NULL,
|
||||
applied_at TEXT NOT NULL
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _applied(connection: sqlite3.Connection) -> dict[int, str]:
|
||||
return {
|
||||
int(row["version"]): str(row["checksum"])
|
||||
for row in connection.execute(
|
||||
"SELECT version, checksum FROM schema_migrations ORDER BY version"
|
||||
)
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _verify_history(migrations: tuple[Migration, ...], applied: dict[int, str]) -> None:
|
||||
by_version = {migration.version: migration for migration in migrations}
|
||||
known_versions = sorted(by_version)
|
||||
applied_versions = sorted(applied)
|
||||
unknown = sorted(set(applied) - set(known_versions))
|
||||
if unknown:
|
||||
values = ", ".join(f"{version:04d}" for version in unknown)
|
||||
raise MigrationError(f"Database contains unknown migrations: {values}")
|
||||
if applied_versions != known_versions[: len(applied_versions)]:
|
||||
raise MigrationError("Database migration history is not contiguous")
|
||||
for version in applied_versions:
|
||||
if applied[version] != by_version[version].checksum:
|
||||
raise MigrationError(
|
||||
f"Migration checksum changed: {version:04d} {by_version[version].name}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _run_action(
|
||||
connection: sqlite3.Connection,
|
||||
migration: Migration,
|
||||
action: MigrationAction,
|
||||
direction: str,
|
||||
) -> None:
|
||||
savepoint = f"migration_{migration.version:04d}_{direction}"
|
||||
connection.execute(f"SAVEPOINT {savepoint}")
|
||||
try:
|
||||
action(connection)
|
||||
connection.execute(f"RELEASE SAVEPOINT {savepoint}")
|
||||
except Exception as exc:
|
||||
connection.execute(f"ROLLBACK TO SAVEPOINT {savepoint}")
|
||||
connection.execute(f"RELEASE SAVEPOINT {savepoint}")
|
||||
raise MigrationError(
|
||||
f"Migration {direction} failed: {migration.version:04d} {migration.name}"
|
||||
) from exc
|
||||
Reference in New Issue
Block a user