345 lines
15 KiB
Python
345 lines
15 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"^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 _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,
|
|
)
|
|
|
|
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 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,
|
|
) -> 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 = self._reader.raw_connection()
|
|
try:
|
|
with raw.cursor() as cursor:
|
|
for collection in collections:
|
|
table = _collection(self._schema, collection)
|
|
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 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())
|
|
finally:
|
|
raw.close()
|
|
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[: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 = engine.raw_connection()
|
|
try:
|
|
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())
|
|
finally:
|
|
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")
|
|
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 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:
|
|
with raw.cursor() as cursor:
|
|
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 Exception:
|
|
raw.rollback()
|
|
raise
|
|
finally:
|
|
raw.close()
|
|
return len(records)
|
|
|
|
|
|
__all__ = ["ALLOWED_COLLECTIONS", "PgVectorStore"]
|