fix(vector): harden packaged migrations
This commit is contained in:
@@ -6,6 +6,8 @@ import hashlib
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from importlib.resources import files
|
||||
from importlib.resources.abc import Traversable
|
||||
from pathlib import Path
|
||||
|
||||
import typer
|
||||
@@ -14,7 +16,7 @@ from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from tht.cli.vector_cmd import vector_app
|
||||
|
||||
MIGRATIONS_DIR = Path(__file__).parents[2] / "migrations" / "vector"
|
||||
MIGRATIONS_DIR = files("tht").joinpath("migrations", "vector")
|
||||
_MIGRATION_NAME = re.compile(r"^(?P<version>\d+)_(?P<name>[a-z0-9_]+)\.sql$")
|
||||
_LOCK_KEY = 7_304_708_654_221_909_028
|
||||
|
||||
@@ -27,7 +29,7 @@ class MigrationError(RuntimeError):
|
||||
class Migration:
|
||||
version: str
|
||||
name: str
|
||||
path: Path
|
||||
path: Traversable
|
||||
checksum: str
|
||||
|
||||
|
||||
@@ -38,32 +40,44 @@ class MigrationStatus:
|
||||
drifted: tuple[Migration, ...]
|
||||
|
||||
|
||||
def _discover(directory: Path) -> tuple[Migration, ...]:
|
||||
def _migration_source(directory: Traversable | Path | str) -> Traversable:
|
||||
return Path(directory) if isinstance(directory, (str, Path)) else directory
|
||||
|
||||
|
||||
def _discover(directory: Traversable | Path | str) -> tuple[Migration, ...]:
|
||||
source = _migration_source(directory)
|
||||
migrations = []
|
||||
seen_versions: set[str] = set()
|
||||
for path in sorted(directory.glob("*.sql")):
|
||||
seen_versions: set[int] = set()
|
||||
paths = [path for path in source.iterdir() if path.name.endswith(".sql")]
|
||||
parsed = []
|
||||
for path in paths:
|
||||
match = _MIGRATION_NAME.fullmatch(path.name)
|
||||
if match is None:
|
||||
raise MigrationError(f"Invalid migration filename: {path.name}")
|
||||
version = match.group("version")
|
||||
if version in seen_versions:
|
||||
raise MigrationError(f"Duplicate migration version: {version}")
|
||||
seen_versions.add(version)
|
||||
numeric_version = int(version)
|
||||
if numeric_version in seen_versions:
|
||||
raise MigrationError(f"Duplicate migration version: {numeric_version}")
|
||||
seen_versions.add(numeric_version)
|
||||
parsed.append((numeric_version, version, match.group("name"), path))
|
||||
for _, version, name, path in sorted(parsed, key=lambda item: item[0]):
|
||||
migrations.append(
|
||||
Migration(
|
||||
version=version,
|
||||
name=match.group("name"),
|
||||
name=name,
|
||||
path=path,
|
||||
checksum=hashlib.sha256(path.read_bytes()).hexdigest(),
|
||||
)
|
||||
)
|
||||
if not migrations:
|
||||
raise MigrationError(f"No migrations found in {directory}")
|
||||
raise MigrationError(f"No migrations found in {source}")
|
||||
return tuple(migrations)
|
||||
|
||||
|
||||
def _applied(connection) -> dict[str, str]:
|
||||
exists = connection.execute(text("SELECT to_regclass('public.tht_vector_migrations')")).scalar()
|
||||
exists = connection.execute(
|
||||
text("SELECT pg_catalog.to_regclass('public.tht_vector_migrations')")
|
||||
).scalar()
|
||||
if exists is None:
|
||||
return {}
|
||||
return dict(
|
||||
@@ -73,16 +87,32 @@ def _applied(connection) -> dict[str, str]:
|
||||
)
|
||||
|
||||
|
||||
def _reject_unknown_versions(
|
||||
migrations: tuple[Migration, ...], applied_checksums: dict[str, str]
|
||||
) -> None:
|
||||
local_versions = {migration.version for migration in migrations}
|
||||
unknown = sorted(
|
||||
set(applied_checksums) - local_versions,
|
||||
key=lambda version: (0, int(version)) if version.isdigit() else (1, version),
|
||||
)
|
||||
if unknown:
|
||||
raise MigrationError(
|
||||
"Database migration versions absent from local manifest: " + ", ".join(unknown)
|
||||
)
|
||||
|
||||
|
||||
def migration_status(
|
||||
database_url: str, migrations_dir: Path | str = MIGRATIONS_DIR
|
||||
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
|
||||
) -> MigrationStatus:
|
||||
migrations = _discover(Path(migrations_dir))
|
||||
migrations = _discover(migrations_dir)
|
||||
engine = create_engine(database_url)
|
||||
try:
|
||||
with engine.connect() as connection:
|
||||
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
|
||||
applied_checksums = _applied(connection)
|
||||
finally:
|
||||
engine.dispose()
|
||||
_reject_unknown_versions(migrations, applied_checksums)
|
||||
applied = tuple(
|
||||
migration
|
||||
for migration in migrations
|
||||
@@ -100,25 +130,31 @@ def migration_status(
|
||||
return MigrationStatus(applied=applied, pending=pending, drifted=drifted)
|
||||
|
||||
|
||||
def migrate(database_url: str, migrations_dir: Path | str = MIGRATIONS_DIR) -> MigrationStatus:
|
||||
migrations = _discover(Path(migrations_dir))
|
||||
def migrate(
|
||||
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
|
||||
) -> MigrationStatus:
|
||||
migrations = _discover(migrations_dir)
|
||||
engine = create_engine(database_url)
|
||||
current: Migration | None = None
|
||||
try:
|
||||
with engine.begin() as connection:
|
||||
connection.execute(text("SELECT pg_advisory_xact_lock(:key)"), {"key": _LOCK_KEY})
|
||||
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
|
||||
connection.execute(
|
||||
text("SELECT pg_catalog.pg_advisory_xact_lock(:key)"), {"key": _LOCK_KEY}
|
||||
)
|
||||
connection.exec_driver_sql(
|
||||
"""CREATE TABLE IF NOT EXISTS public.tht_vector_migrations (
|
||||
version text PRIMARY KEY,
|
||||
name text NOT NULL,
|
||||
checksum text NOT NULL,
|
||||
applied_at timestamptz NOT NULL DEFAULT now()
|
||||
applied_at timestamptz NOT NULL DEFAULT pg_catalog.now()
|
||||
)"""
|
||||
)
|
||||
connection.exec_driver_sql(
|
||||
"REVOKE ALL ON public.tht_vector_migrations FROM PUBLIC"
|
||||
)
|
||||
applied_checksums = _applied(connection)
|
||||
_reject_unknown_versions(migrations, applied_checksums)
|
||||
drifted = [
|
||||
item
|
||||
for item in migrations
|
||||
|
||||
Reference in New Issue
Block a user