"""Versioned, transactional migrations for the direct pgvector schema.""" from __future__ import annotations import hashlib import json import re from dataclasses import dataclass from pathlib import Path import typer from sqlalchemy import create_engine, text from sqlalchemy.exc import SQLAlchemyError from tht.cli.vector_cmd import vector_app MIGRATIONS_DIR = Path(__file__).parents[2] / "migrations" / "vector" _MIGRATION_NAME = re.compile(r"^(?P\d+)_(?P[a-z0-9_]+)\.sql$") _LOCK_KEY = 7_304_708_654_221_909_028 class MigrationError(RuntimeError): """Raised when migration discovery or application is unsafe.""" @dataclass(frozen=True) class Migration: version: str name: str path: Path checksum: str @dataclass(frozen=True) class MigrationStatus: applied: tuple[Migration, ...] pending: tuple[Migration, ...] drifted: tuple[Migration, ...] def _discover(directory: Path) -> tuple[Migration, ...]: migrations = [] seen_versions: set[str] = set() for path in sorted(directory.glob("*.sql")): 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) migrations.append( Migration( version=version, name=match.group("name"), path=path, checksum=hashlib.sha256(path.read_bytes()).hexdigest(), ) ) if not migrations: raise MigrationError(f"No migrations found in {directory}") return tuple(migrations) def _applied(connection) -> dict[str, str]: exists = connection.execute(text("SELECT to_regclass('public.tht_vector_migrations')")).scalar() if exists is None: return {} return dict( connection.execute( text("SELECT version, checksum FROM public.tht_vector_migrations") ).all() ) def migration_status( database_url: str, migrations_dir: Path | str = MIGRATIONS_DIR ) -> MigrationStatus: migrations = _discover(Path(migrations_dir)) engine = create_engine(database_url) try: with engine.connect() as connection: applied_checksums = _applied(connection) finally: engine.dispose() applied = tuple( migration for migration in migrations if applied_checksums.get(migration.version) == migration.checksum ) drifted = tuple( migration for migration in migrations if migration.version in applied_checksums and applied_checksums[migration.version] != migration.checksum ) pending = tuple( migration for migration in migrations if migration.version not in applied_checksums ) 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)) 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( """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() )""" ) connection.exec_driver_sql( "REVOKE ALL ON public.tht_vector_migrations FROM PUBLIC" ) applied_checksums = _applied(connection) drifted = [ item for item in migrations if item.version in applied_checksums and applied_checksums[item.version] != item.checksum ] if drifted: versions = ", ".join(item.version for item in drifted) raise MigrationError(f"Migration checksum drift: {versions}") for current in migrations: if current.version in applied_checksums: continue connection.exec_driver_sql(current.path.read_text()) connection.execute( text( "INSERT INTO public.tht_vector_migrations (version, name, checksum) " "VALUES (:version, :name, :checksum)" ), { "version": current.version, "name": current.name, "checksum": current.checksum, }, ) except MigrationError: raise except SQLAlchemyError as exc: filename = current.path.name if current is not None else "migration setup" raise MigrationError(f"Failed to apply {filename}: {type(exc).__name__}") from exc finally: engine.dispose() return migration_status(database_url, migrations_dir) def _payload(status: MigrationStatus) -> dict[str, list[str]]: return { "applied": [item.version for item in status.applied], "drifted": [item.version for item in status.drifted], "pending": [item.version for item in status.pending], } @vector_app.command("migrate") def migrate_cmd( database_url: str = typer.Option( ..., "--database-url", envvar="THT_VECTOR_ADMIN_URL", help="Admin PostgreSQL URL." ), status_only: bool = typer.Option(False, "--status", help="Inspect without applying."), json_output: bool = typer.Option(False, "--json", help="Emit pristine JSON."), ) -> None: """Apply or inspect the local pgvector schema migrations.""" try: status = migration_status(database_url) if status_only else migrate(database_url) except (MigrationError, SQLAlchemyError) as exc: if json_output: typer.echo(json.dumps({"error": str(exc)}, sort_keys=True)) else: typer.echo(f"ERROR: {exc}", err=True) raise typer.Exit(code=1) from None payload = _payload(status) if json_output: typer.echo(json.dumps(payload, sort_keys=True)) else: typer.echo( f"Applied: {len(status.applied)}; pending: {len(status.pending)}; " f"drifted: {len(status.drifted)}" ) __all__ = ["MigrationError", "MigrationStatus", "migrate", "migration_status"]