"""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"]