feat(vector): version pgvector schema
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
"""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"]
|
||||
Reference in New Issue
Block a user