fix(vector): harden packaged migrations

This commit is contained in:
2026-07-12 01:32:33 +02:00
parent c0d50e9b08
commit 3588a7749b
14 changed files with 337 additions and 54 deletions
+23 -8
View File
@@ -27,7 +27,7 @@ COLLECTION_KINDS = {
}
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
ALLOWED_KINDS = frozenset().union(*COLLECTION_KINDS.values())
_VECTOR_DIMENSION = re.compile(r"^vector\((\d+)\)$")
_VECTOR_DIMENSION = re.compile(r"^(?:[a-z_][a-z0-9_]*\.)?vector\((\d+)\)$")
def _collection(schema: str, name: str) -> sql.Identifier:
@@ -40,6 +40,14 @@ def _vector_literal(values: list[float]) -> str:
return "[" + ",".join(str(float(value)) for value in values) + "]"
def _vector_type(schema: str) -> sql.Identifier:
return sql.Identifier(schema, "vector")
def _cosine_operator(schema: str) -> sql.Composed:
return sql.SQL("OPERATOR({}.<=>)").format(sql.Identifier(schema))
def _validate_collection_kinds(collection: str, kinds: list[str]) -> None:
invalid = set(kinds) - COLLECTION_KINDS[collection]
if invalid:
@@ -248,9 +256,16 @@ class PgVectorStore:
continue
where = sql.SQL(" WHERE kind = ANY(%s)") if collection_kinds else sql.SQL("")
query = sql.SQL(
"SELECT metadata, 1 - (embedding <=> %s::vector) AS similarity "
"FROM {}{} ORDER BY embedding <=> %s::vector LIMIT %s"
).format(table, where)
"SELECT metadata, 1 - (embedding {} %s::{}) AS similarity "
"FROM {}{} ORDER BY embedding {} %s::{} LIMIT %s"
).format(
_cosine_operator(self._schema),
_vector_type(self._schema),
table,
where,
_cosine_operator(self._schema),
_vector_type(self._schema),
)
params = [_vector_literal(embedding)]
if collection_kinds:
params.append(collection_kinds)
@@ -295,13 +310,13 @@ class PgVectorStore:
raise VectorStoreError("Embedding dimension does not match configured dimension")
insert = sql.SQL(
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
"VALUES (%s, %s, %s, %s::jsonb, %s::{}) "
"ON CONFLICT (record_key) DO NOTHING"
).format(table)
).format(table, _vector_type(self._schema))
update = sql.SQL(
"UPDATE {} SET kind = %s, content_hash = %s, metadata = %s::jsonb, "
"embedding = %s::vector, indexed_at = now() WHERE record_key = %s"
).format(table)
"embedding = %s::{}, indexed_at = pg_catalog.now() WHERE record_key = %s"
).format(table, _vector_type(self._schema))
raw = engine.raw_connection()
try:
with raw.cursor() as cursor:
+53 -17
View File
@@ -6,6 +6,8 @@ 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
@@ -14,7 +16,7 @@ from sqlalchemy.exc import SQLAlchemyError
from tht.cli.vector_cmd import vector_app
MIGRATIONS_DIR = Path(__file__).parents[2] / "migrations" / "vector"
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
@@ -27,7 +29,7 @@ class MigrationError(RuntimeError):
class Migration:
version: str
name: str
path: Path
path: Traversable
checksum: str
@@ -38,32 +40,44 @@ class MigrationStatus:
drifted: tuple[Migration, ...]
def _discover(directory: Path) -> 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[str] = set()
for path in sorted(directory.glob("*.sql")):
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")
if version in seen_versions:
raise MigrationError(f"Duplicate migration version: {version}")
seen_versions.add(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=match.group("name"),
name=name,
path=path,
checksum=hashlib.sha256(path.read_bytes()).hexdigest(),
)
)
if not migrations:
raise MigrationError(f"No migrations found in {directory}")
raise MigrationError(f"No migrations found in {source}")
return tuple(migrations)
def _applied(connection) -> dict[str, str]:
exists = connection.execute(text("SELECT to_regclass('public.tht_vector_migrations')")).scalar()
exists = connection.execute(
text("SELECT pg_catalog.to_regclass('public.tht_vector_migrations')")
).scalar()
if exists is None:
return {}
return dict(
@@ -73,16 +87,32 @@ def _applied(connection) -> dict[str, str]:
)
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: Path | str = MIGRATIONS_DIR
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
) -> MigrationStatus:
migrations = _discover(Path(migrations_dir))
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
@@ -100,25 +130,31 @@ def migration_status(
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))
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.execute(text("SELECT pg_advisory_xact_lock(:key)"), {"key": _LOCK_KEY})
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 now()
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
@@ -0,0 +1,3 @@
CREATE SCHEMA IF NOT EXISTS vectors;
REVOKE ALL ON SCHEMA vectors FROM PUBLIC;
CREATE EXTENSION IF NOT EXISTS vector WITH SCHEMA vectors;
@@ -0,0 +1,32 @@
CREATE TABLE IF NOT EXISTS vectors.schema_records (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
CREATE TABLE IF NOT EXISTS vectors.evidence (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
CREATE TABLE IF NOT EXISTS vectors.memory (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
REVOKE ALL ON ALL TABLES IN SCHEMA vectors FROM PUBLIC;
REVOKE ALL ON ALL SEQUENCES IN SCHEMA vectors FROM PUBLIC;
@@ -0,0 +1,23 @@
DO $roles$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_catalog.pg_roles WHERE rolname = 'vector_reader') THEN
CREATE ROLE vector_reader NOLOGIN;
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_catalog.pg_roles WHERE rolname = 'vector_writer') THEN
CREATE ROLE vector_writer NOLOGIN;
END IF;
END
$roles$;
REVOKE ALL ON SCHEMA vectors FROM vector_reader, vector_writer;
REVOKE ALL ON ALL TABLES IN SCHEMA vectors FROM vector_reader, vector_writer;
REVOKE ALL ON ALL SEQUENCES IN SCHEMA vectors FROM vector_reader, vector_writer;
GRANT USAGE ON SCHEMA vectors TO vector_reader, vector_writer;
GRANT SELECT ON ALL TABLES IN SCHEMA vectors TO vector_reader;
GRANT INSERT, UPDATE
ON vectors.schema_records, vectors.evidence, vectors.memory TO vector_writer;
GRANT SELECT (record_key, kind, content_hash)
ON vectors.schema_records, vectors.evidence, vectors.memory TO vector_writer;
GRANT USAGE ON ALL SEQUENCES IN SCHEMA vectors TO vector_writer;