Files
ThothII/harness/tht/cli/vector_migrate_cmd.py
T

228 lines
8.0 KiB
Python

"""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<version>\d+)_(?P<name>[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"]