"""Versioned, transactional migrations for the direct pgvector schema.""" from __future__ import annotations 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 from sqlalchemy import create_engine, text from sqlalchemy.exc import SQLAlchemyError from tht.cli.vector_cmd import vector_app MIGRATIONS_DIR = files("tht").joinpath("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: Traversable checksum: str @dataclass(frozen=True) class MigrationStatus: applied: tuple[Migration, ...] pending: tuple[Migration, ...] drifted: 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[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") 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=name, path=path, checksum=hashlib.sha256(path.read_bytes()).hexdigest(), ) ) if not migrations: raise MigrationError(f"No migrations found in {source}") return tuple(migrations) def _applied(connection) -> dict[str, str]: exists = connection.execute( text("SELECT pg_catalog.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 _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: Traversable | Path | str = MIGRATIONS_DIR ) -> MigrationStatus: 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 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: 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.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 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 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"]