fix(vector): harden packaged migrations

This commit is contained in:
2026-07-12 01:32:33 +02:00
parent c0d50e9b08
commit 3588a7749b
14 changed files with 337 additions and 54 deletions
+53 -17
View File
@@ -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