fix(vector): harden direct pgvector parity

This commit is contained in:
2026-07-12 01:09:24 +02:00
parent b09341f07e
commit e7e948c77c
3 changed files with 291 additions and 41 deletions
+110 -29
View File
@@ -26,6 +26,7 @@ COLLECTION_KINDS = {
"memory": {"memory", "solved_question"},
}
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
ALLOWED_KINDS = frozenset().union(*COLLECTION_KINDS.values())
_VECTOR_DIMENSION = re.compile(r"^vector\((\d+)\)$")
@@ -39,12 +40,18 @@ def _vector_literal(values: list[float]) -> str:
return "[" + ",".join(str(float(value)) for value in values) + "]"
def _validate_kinds(collection: str, kinds: list[str]) -> None:
def _validate_collection_kinds(collection: str, kinds: list[str]) -> None:
invalid = set(kinds) - COLLECTION_KINDS[collection]
if invalid:
raise VectorStoreError(f"Kind not allowed for {collection}: {', '.join(sorted(invalid))}")
def _validate_known_kinds(kinds: list[str]) -> None:
invalid = set(kinds) - ALLOWED_KINDS
if invalid:
raise VectorStoreError(f"Kind not allowed: {', '.join(sorted(invalid))}")
class PgVectorStore:
"""Direct store with independent reader and writer database credentials."""
@@ -72,7 +79,9 @@ class PgVectorStore:
upsert=writable,
)
def _probe(self, engine: Engine | None) -> tuple[bool | None, str | None, set[int]]:
def _probe(
self, engine: Engine | None, *, writable: bool
) -> tuple[bool | None, str | None, set[int]]:
if engine is None:
return None, None, set()
try:
@@ -81,28 +90,86 @@ class PgVectorStore:
with raw.cursor() as cursor:
cursor.execute("SELECT 1")
cursor.execute(
"""SELECT format_type(a.atttypid, a.atttypmod)
FROM pg_attribute a
JOIN pg_class c ON c.oid = a.attrelid
"""SELECT c.relname, format_type(a.atttypid, a.atttypmod),
has_table_privilege(current_user, c.oid, 'SELECT'),
has_table_privilege(current_user, c.oid, 'INSERT'),
has_table_privilege(current_user, c.oid, 'UPDATE'),
has_column_privilege(current_user, c.oid, 'record_key', 'SELECT')
AND has_column_privilege(
current_user, c.oid, 'content_hash', 'SELECT'
)
AND has_column_privilege(current_user, c.oid, 'kind', 'SELECT')
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
LEFT JOIN pg_attribute a ON a.attrelid = c.oid
AND a.attname = 'embedding' AND NOT a.attisdropped
WHERE n.nspname = %s AND c.relname = ANY(%s)
AND a.attname = 'embedding' AND NOT a.attisdropped""",
AND c.relkind IN ('r', 'p')""",
(self._schema, list(ALLOWED_COLLECTIONS)),
)
rows = cursor.fetchall()
present = {row[0] for row in rows}
missing_tables = sorted(ALLOWED_COLLECTIONS - present)
missing_embeddings = sorted(row[0] for row in rows if row[1] is None)
privilege_missing = sorted(
row[0]
for row in rows
if (writable and not (row[3] and row[4] and row[5]))
or (not writable and not row[2])
)
problems = []
if missing_tables:
problems.append("missing tables " + ", ".join(missing_tables))
if missing_embeddings:
problems.append(
"missing embedding columns " + ", ".join(missing_embeddings)
)
if privilege_missing:
authority = "write" if writable else "read"
problems.append(
f"missing {authority} privileges " + ", ".join(privilege_missing)
)
if problems:
return False, "vector schema incomplete: " + "; ".join(problems), set()
dimensions = {
int(match.group(1))
for (type_name,) in cursor.fetchall()
for _, type_name, *_ in rows
if (match := _VECTOR_DIMENSION.match(type_name))
}
invalid_types = sorted(
row[0]
for row in rows
if row[1] is not None and not _VECTOR_DIMENSION.match(row[1])
)
if invalid_types:
return (
False,
"vector schema incomplete: invalid embedding types "
+ ", ".join(invalid_types),
set(),
)
if self._expected_dimension is not None:
mismatches = sorted(
f"{name}={int(match.group(1))}"
for name, type_name, *_ in rows
if (match := _VECTOR_DIMENSION.match(type_name))
and int(match.group(1)) != self._expected_dimension
)
if mismatches:
return (
False,
"embedding dimension mismatch: " + ", ".join(mismatches),
dimensions,
)
return True, None, dimensions
finally:
raw.close()
except Exception as exc:
return False, str(exc), set()
return False, f"vector database probe failed: {type(exc).__name__}", set()
def health(self) -> VectorHealth:
read_ok, read_detail, read_dimensions = self._probe(self._reader)
write_ok, write_detail, write_dimensions = self._probe(self._writer)
read_ok, read_detail, read_dimensions = self._probe(self._reader, writable=False)
write_ok, write_detail, write_dimensions = self._probe(self._writer, writable=True)
dimensions = tuple(sorted(read_dimensions | write_dimensions))
compatible = (
None
@@ -138,22 +205,27 @@ class PgVectorStore:
raise VectorReadUnavailable("Vector reader credential is not configured")
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
raise VectorStoreError("Query embedding dimension does not match configured dimension")
if kinds:
_validate_known_kinds(kinds)
hits: list[VectorHit] = []
raw = self._reader.raw_connection()
try:
with raw.cursor() as cursor:
for collection in collections:
table = _collection(self._schema, collection)
if kinds:
_validate_kinds(collection, kinds)
where = sql.SQL(" WHERE kind = ANY(%s)") if kinds else sql.SQL("")
collection_kinds = (
sorted(set(kinds) & COLLECTION_KINDS[collection]) if kinds else None
)
if kinds and not collection_kinds:
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)
params = [_vector_literal(embedding)]
if kinds:
params.append(kinds)
if collection_kinds:
params.append(collection_kinds)
params.extend((_vector_literal(embedding), limit))
cursor.execute(query, params)
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
@@ -169,7 +241,7 @@ class PgVectorStore:
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
engine = self._require_writer()
table = _collection(self._schema, collection)
_validate_kinds(collection, kinds)
_validate_collection_kinds(collection, kinds)
raw = engine.raw_connection()
try:
with raw.cursor() as cursor:
@@ -187,18 +259,20 @@ class PgVectorStore:
engine = self._require_writer()
table = _collection(self._schema, collection)
for write_record in records:
_validate_kinds(collection, [write_record.record.kind])
_validate_collection_kinds(collection, [write_record.record.kind])
if (
self._expected_dimension is not None
and len(write_record.embedding) != self._expected_dimension
):
raise VectorStoreError("Embedding dimension does not match configured dimension")
query = sql.SQL(
insert = sql.SQL(
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
"ON CONFLICT (record_key) DO UPDATE SET kind = EXCLUDED.kind, "
"content_hash = EXCLUDED.content_hash, metadata = EXCLUDED.metadata, "
"embedding = EXCLUDED.embedding, indexed_at = now()"
"ON CONFLICT (record_key) DO NOTHING"
).format(table)
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)
raw = engine.raw_connection()
try:
@@ -213,16 +287,23 @@ class PgVectorStore:
"content": record.content,
**record.metadata,
}
metadata_json = json.dumps(metadata)
vector = _vector_literal(write_record.embedding)
cursor.execute(
query,
(
record.id,
record.kind,
write_record.content_hash,
json.dumps(metadata),
_vector_literal(write_record.embedding),
),
insert,
(record.id, record.kind, write_record.content_hash, metadata_json, vector),
)
if cursor.rowcount == 0:
cursor.execute(
update,
(
record.kind,
write_record.content_hash,
metadata_json,
vector,
record.id,
),
)
raw.commit()
except Exception:
raw.rollback()