Files
ThothII/harness/tht/adapters/vector/pgvector.py
T

525 lines
24 KiB
Python

"""Direct PostgreSQL/pgvector implementation of the vector port."""
import json
import re
from psycopg2 import Error as PsycopgError
from psycopg2 import sql
from sqlalchemy import Engine
from sqlalchemy.exc import SQLAlchemyError
from tht.config import DatabaseConfig
from tht.db.connection import make_engine
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorReadUnavailable,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
require_positive_limit,
)
from tht.vectorstore.store import VectorHit, hit_from_metadata
COLLECTION_KINDS = {
"schema_records": {"schema_table", "schema_column"},
"evidence": {"evidence"},
"memory": {"memory", "solved_question"},
}
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
ALLOWED_KINDS = frozenset().union(*COLLECTION_KINDS.values())
_VECTOR_DIMENSION = re.compile(r"^(?:[a-z_][a-z0-9_]*\.)?vector\((\d+)\)$")
def _collection(schema: str, name: str) -> sql.Identifier:
if name not in ALLOWED_COLLECTIONS:
raise VectorStoreError(f"Collection not allowed: {name}")
return sql.Identifier(schema, name)
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 _vector_sql_names(cursor, table_schema: str, collection: str) -> tuple[str, str]:
"""Discover pgvector type and operator namespaces from the embedding column."""
cursor.execute(
"""SELECT type_ns.nspname, operator_ns.nspname
FROM pg_catalog.pg_attribute attribute
JOIN pg_catalog.pg_class table_class
ON table_class.oid = attribute.attrelid
JOIN pg_catalog.pg_namespace table_ns
ON table_ns.oid = table_class.relnamespace
JOIN pg_catalog.pg_type vector_type
ON vector_type.oid = attribute.atttypid
JOIN pg_catalog.pg_namespace type_ns
ON type_ns.oid = vector_type.typnamespace
JOIN pg_catalog.pg_operator cosine
ON cosine.oprname = %s
AND cosine.oprleft = vector_type.oid
AND cosine.oprright = vector_type.oid
JOIN pg_catalog.pg_namespace operator_ns
ON operator_ns.oid = cosine.oprnamespace
WHERE table_ns.nspname = %s
AND table_class.relname = %s
AND attribute.attname = %s
AND NOT attribute.attisdropped
ORDER BY cosine.oid
LIMIT 1""",
("<=>", table_schema, collection, "embedding"),
)
row = cursor.fetchone()
if row is None:
raise VectorStoreError(f"Collection {collection} has no usable pgvector embedding")
return row[0], row[1]
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."""
def __init__(
self,
read_config: DatabaseConfig | None,
write_config: DatabaseConfig | None = None,
*,
expected_dimension: int | None = None,
):
self._reader = make_engine(read_config) if read_config is not None else None
self._writer = make_engine(write_config) if write_config is not None else None
config = read_config or write_config
self._schema = config.db_schema if config is not None else "vectors"
if read_config and write_config and read_config.db_schema != write_config.db_schema:
raise VectorStoreError("Reader and writer vector schemas must match")
self._expected_dimension = expected_dimension
@property
def capabilities(self) -> VectorCapabilities:
writable = self._writer is not None
return VectorCapabilities(
search=self._reader is not None,
existing_hashes=writable,
upsert=writable,
metadata_filter=self._reader is not None,
delete_generation=writable,
list_evidence_generations=writable,
)
def _probe(
self, engine: Engine | None, *, writable: bool
) -> tuple[bool | None, str | None, set[int]]:
if engine is None:
return None, None, set()
try:
raw = engine.raw_connection()
try:
with raw.cursor() as cursor:
cursor.execute("SELECT 1")
cursor.execute(
"SELECT has_schema_privilege(current_user, %s, 'USAGE')",
(self._schema,),
)
schema_usage = bool(cursor.fetchone()[0])
if not schema_usage:
return False, "vector schema incomplete: missing schema usage", set()
cursor.execute(
"""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'),
CASE WHEN id_attr.attname IS NOT NULL THEN
pg_get_serial_sequence(
format('%%I.%%I', n.nspname, c.relname), 'id'
)
END AS id_sequence,
CASE WHEN id_attr.attname IS NOT NULL THEN
has_sequence_privilege(
current_user,
pg_get_serial_sequence(
format('%%I.%%I', n.nspname, c.relname), 'id'
),
'USAGE'
)
END AS sequence_usage
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
LEFT JOIN pg_attribute id_attr ON id_attr.attrelid = c.oid
AND id_attr.attname = 'id' AND NOT id_attr.attisdropped
WHERE n.nspname = %s AND c.relname = ANY(%s)
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])
)
missing_sequences = sorted(
row[0] for row in rows if writable and row[6] is None
)
sequence_privilege_missing = sorted(
row[0] for row in rows if writable and row[6] is not None and not row[7]
)
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 missing_sequences:
problems.append("missing id sequences " + ", ".join(missing_sequences))
if sequence_privilege_missing:
problems.append(
"missing sequence privileges " + ", ".join(sequence_privilege_missing)
)
if problems:
return False, "vector schema incomplete: " + "; ".join(problems), set()
dimensions = {
int(match.group(1))
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 (AttributeError, TypeError, ValueError, PsycopgError, SQLAlchemyError) as exc:
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, writable=False)
write_ok, write_detail, write_dimensions = self._probe(self._writer, writable=True)
dimensions = tuple(sorted(read_dimensions | write_dimensions))
compatible = (
None
if self._expected_dimension is None or not dimensions
else dimensions == (self._expected_dimension,)
)
reachable = [value for value in (read_ok, write_ok) if value is not None]
details = [value for value in (read_detail, write_detail) if value]
return VectorHealth(
ok=bool(reachable) and all(reachable) and compatible is not False,
detail="; ".join(details) or None,
read_configured=self._reader is not None,
read_reachable=read_ok,
read_detail=read_detail,
write_configured=self._writer is not None,
write_reachable=write_ok,
write_detail=write_detail,
expected_dimension=self._expected_dimension,
observed_dimensions=dimensions,
dimension_compatible=compatible,
)
def search(
self,
collections: list[str],
embedding: list[float],
*,
limit: int,
kinds: list[str] | None = None,
metadata_filter: dict[str, object] | None = None,
) -> list[VectorHit]:
require_positive_limit(limit)
if self._reader is None:
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 = None
try:
raw = self._reader.raw_connection()
with raw.cursor() as cursor:
for collection in collections:
table = _collection(self._schema, collection)
type_schema, operator_schema = _vector_sql_names(
cursor, self._schema, collection
)
collection_kinds = (
sorted(set(kinds) & COLLECTION_KINDS[collection]) if kinds else None
)
if kinds and not collection_kinds:
continue
clauses = []
filter_params = []
if collection_kinds:
clauses.append(sql.SQL("kind = ANY(%s)"))
filter_params.append(collection_kinds)
if metadata_filter is not None:
if collection != "evidence" or set(metadata_filter) != {
"vector_generation", "document_ids", "workspace_id"
}:
raise VectorStoreError("Unsupported vector metadata filter")
generation = metadata_filter["vector_generation"]
document_ids = metadata_filter["document_ids"]
workspace_id = metadata_filter["workspace_id"]
if not isinstance(generation, str) or not isinstance(document_ids, list) or not isinstance(workspace_id, str):
raise VectorStoreError("Invalid vector metadata filter")
clauses.append(sql.SQL("metadata->>'vector_generation' = %s"))
clauses.append(sql.SQL("metadata->>'document_id' = ANY(%s)"))
clauses.append(sql.SQL("metadata->>'workspace_id' = %s"))
filter_params.extend((generation, document_ids, workspace_id))
where = (
sql.SQL(" WHERE ") + sql.SQL(" AND ").join(clauses)
if clauses else sql.SQL("")
)
query = sql.SQL(
"SELECT metadata, 1 - (embedding {} %s::{}) AS similarity "
"FROM {}{} ORDER BY embedding {} %s::{}, record_key LIMIT %s"
).format(
_cosine_operator(operator_schema),
_vector_type(type_schema),
table,
where,
_cosine_operator(operator_schema),
_vector_type(type_schema),
)
params = [_vector_literal(embedding)]
params.extend(filter_params)
params.extend((_vector_literal(embedding), limit))
cursor.execute(query, params)
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
except VectorStoreError:
raise
except Exception as exc:
raise VectorReadUnavailable("Vector read operation unavailable") from exc
finally:
if raw is not None:
raw.close()
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:limit]
def _require_writer(self) -> Engine:
if self._writer is None:
raise VectorWriteUnavailable("Vector writer credential is not configured")
return self._writer
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
engine = self._require_writer()
table = _collection(self._schema, collection)
_validate_collection_kinds(collection, kinds)
raw = None
try:
raw = engine.raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL("SELECT record_key, content_hash FROM {} WHERE kind = ANY(%s)").format(
table
),
(kinds,),
)
return dict(cursor.fetchall())
except VectorStoreError:
raise
except Exception as exc:
raise VectorWriteUnavailable("Vector write operation unavailable") from exc
finally:
if raw is not None:
raw.close()
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
engine = self._require_writer()
table = _collection(self._schema, collection)
for write_record in records:
_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")
raw = None
try:
raw = engine.raw_connection()
with raw.cursor() as cursor:
type_schema, _ = _vector_sql_names(cursor, self._schema, collection)
insert = sql.SQL(
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
"VALUES (%s, %s, %s, %s::jsonb, %s::{}) "
"ON CONFLICT (record_key) DO NOTHING"
).format(table, _vector_type(type_schema))
update = sql.SQL(
"UPDATE {} SET kind = %s, content_hash = %s, metadata = %s::jsonb, "
"embedding = %s::{}, indexed_at = pg_catalog.now() WHERE record_key = %s"
).format(table, _vector_type(type_schema))
for write_record in records:
record = write_record.record
metadata = {
"kind": record.kind,
"ref": record.ref,
"record_key": record.id,
"title": record.title,
"content": record.content,
**record.metadata,
}
metadata_json = json.dumps(metadata)
vector = _vector_literal(write_record.embedding)
cursor.execute(
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 VectorStoreError:
if raw is not None:
raw.rollback()
raise
except Exception as exc:
if raw is not None:
raw.rollback()
raise VectorWriteUnavailable("Vector write operation unavailable") from exc
finally:
if raw is not None:
raw.close()
return len(records)
def delete_generation(self, collection: str, generation: str, workspace_id: str) -> int:
if collection != "evidence" or re.fullmatch(r"gen:[0-9a-f]{32}", generation) is None:
raise VectorStoreError("Only exact Evidence generations may be deleted")
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise VectorStoreError("Invalid Evidence workspace namespace")
raw = None
try:
raw = self._require_writer().raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL(
"DELETE FROM {} WHERE kind = 'evidence' "
"AND metadata->>'vector_generation' = %s "
"AND metadata->>'workspace_id' = %s"
).format(_collection(self._schema, collection)),
(generation, workspace_id),
)
count = cursor.rowcount
raw.commit()
return count
except Exception as exc:
if raw is not None:
raw.rollback()
raise VectorWriteUnavailable("Vector generation cleanup unavailable") from exc
finally:
if raw is not None:
raw.close()
def delete_kinds(self, collection: str, kinds: list[str]) -> int:
_collection(self._schema, collection)
_validate_collection_kinds(collection, kinds)
raw = None
try:
raw = self._require_writer().raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL("DELETE FROM {} WHERE kind = ANY(%s)").format(
_collection(self._schema, collection)
),
(kinds,),
)
count = cursor.rowcount
raw.commit()
return count
except Exception as exc:
if raw is not None:
raw.rollback()
raise VectorWriteUnavailable("Vector kind cleanup unavailable") from exc
finally:
if raw is not None:
raw.close()
def list_evidence_generations(self, collection: str, workspace_id: str) -> list[str]:
if collection != "evidence":
raise VectorStoreError("Only exact Evidence generations may be listed")
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise VectorStoreError("Invalid Evidence workspace namespace")
raw = None
try:
raw = self._require_writer().raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL(
"SELECT DISTINCT metadata->>'vector_generation' FROM {} "
"WHERE kind = 'evidence' AND metadata->>'vector_generation' "
"~ '^gen:[0-9a-f]{{32}}$' AND metadata->>'workspace_id' = %s ORDER BY 1"
).format(_collection(self._schema, collection)),
(workspace_id,),
)
return [row[0] for row in cursor.fetchall()]
except Exception as exc:
raise VectorWriteUnavailable("Vector generation inventory unavailable") from exc
finally:
if raw is not None:
raw.close()
__all__ = ["ALLOWED_COLLECTIONS", "PgVectorStore"]