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

192 lines
6.6 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 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<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: 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"]