"""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 _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) 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(self._schema), _vector_type(self._schema), table, where, _cosine_operator(self._schema), _vector_type(self._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") 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(self._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(self._schema)) raw = None try: raw = engine.raw_connection() 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 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"]