228 lines
8.0 KiB
Python
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"]
|