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

500 lines
22 KiB
Python

"""Direct PostgreSQL/pgvector implementation of the vector port."""
import json
import re
from psycopg2 import sql
from sqlalchemy import Engine
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 Exception 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 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"]