fix(vector): harden packaged migrations
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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;
|
||||
Reference in New Issue
Block a user