From 43b3efc4949d45b19618717cccd20f9319d74c36 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicol=C3=B2=20Boschi?= Date: Wed, 11 Mar 2026 12:09:50 +0100 Subject: [PATCH] perf: replace window-function retrieval with UNION ALL + per-bank HNSW indexes (#541) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * perf: replace window-function retrieval with UNION ALL + per-bank HNSW indexes The previous retrieve_semantic_bm25_combined() used ROW_NUMBER() OVER (PARTITION BY fact_type ...) which forced a full sequential scan — pgvector cannot use HNSW indexes when a window function partitions on the same column as the ORDER BY. Changes: - retrieval.py: rewrite to UNION ALL of per-fact_type subqueries; each arm has its own ORDER BY embedding <=> $1 LIMIT n, enabling partial HNSW index scans. Semantic arms over-fetch 5x (min 100) for HNSW approximation; trimmed in Python. - memory_engine.py: set hnsw.ef_search=200 at pool init (persistent per-connection, no per-query SET/RESET overhead). - bank_utils.py: add create_bank_hnsw_indexes / drop_bank_hnsw_indexes for per-(bank_id, fact_type) partial HNSW index lifecycle management. - fact_storage.py / bank_utils.py: create per-bank indexes on fresh bank insert. - memory_engine.py delete_bank: drop per-bank indexes via DELETE...RETURNING to avoid a separate round-trip. - Migration a3b4c5d6e7f8: add interim fact_type-only partial indexes. - Migration d5e6f7a8b9c0: add internal_id UUID UNIQUE to banks, replace fact_type-only indexes with per-(bank, fact_type) partial HNSW indexes, drop the global idx_memory_units_embedding that competed with them. Why per-(bank, fact_type) not just per-fact_type: The idx_memory_units_bank_id B-tree index always wins over fact_type-only partial indexes when bank_id appears in the WHERE clause. Including bank_id in the partial index predicate removes the B-tree from consideration and lets the planner choose HNSW. The global HNSW index must also be dropped to avoid competing for the larger fact_type partitions (world, observation). * refactor: collapse two HNSW migrations into one * refactor: generate bank internal_id in Python before insert Instead of relying on DEFAULT gen_random_uuid() and RETURNING internal_id, generate the UUID in application code before the INSERT. This means we always know the value upfront and can call create_bank_hnsw_indexes immediately without needing a DB round-trip to retrieve the assigned ID. Also adds tests for HNSW index lifecycle and retrieve_semantic_bm25_combined. * fix: correct migration and prevent global HNSW index recreation Migration fixes: - Add text() wrappers for raw SQL in d5e6f7a8b9c0 (SQLAlchemy 2.0 compat) - Drop stale fact_type-only partial indexes (idx_mu_emb_world/observation/experience) that may exist from prior migrations on the same DB migrations.py fix: - Skip global HNSW index creation when per-bank partial HNSW indexes already exist on memory_units (idx_mu_emb_* pattern). Without this, the post-migration vector index check detects no %embedding% named index and recreates the global idx_memory_units_embedding, which defeats the per-bank index strategy. Verified with EXPLAIN ANALYZE on 66K-row bank: all three fact_type arms use their per-bank HNSW index scan (idx_mu_emb_worl/expr/obsv_). * fix: use correct embeddings.encode() in test --- ..._add_bank_internal_id_and_per_bank_hnsw.py | 131 ++++++++++ .../hindsight_api/engine/memory_engine.py | 20 +- .../hindsight_api/engine/retain/bank_utils.py | 68 +++++- .../engine/retain/fact_storage.py | 21 +- .../hindsight_api/engine/search/retrieval.py | 230 +++++++++--------- hindsight-api/hindsight_api/migrations.py | 18 ++ hindsight-api/tests/test_hnsw_indexes.py | 193 +++++++++++++++ 7 files changed, 547 insertions(+), 134 deletions(-) create mode 100644 hindsight-api/hindsight_api/alembic/versions/d5e6f7a8b9c0_add_bank_internal_id_and_per_bank_hnsw.py create mode 100644 hindsight-api/tests/test_hnsw_indexes.py diff --git a/hindsight-api/hindsight_api/alembic/versions/d5e6f7a8b9c0_add_bank_internal_id_and_per_bank_hnsw.py b/hindsight-api/hindsight_api/alembic/versions/d5e6f7a8b9c0_add_bank_internal_id_and_per_bank_hnsw.py new file mode 100644 index 00000000..3d847b01 --- /dev/null +++ b/hindsight-api/hindsight_api/alembic/versions/d5e6f7a8b9c0_add_bank_internal_id_and_per_bank_hnsw.py @@ -0,0 +1,131 @@ +"""Add internal_id to banks and per-(bank, fact_type) partial HNSW indexes + +Revision ID: d5e6f7a8b9c0 +Revises: a3b4c5d6e7f8 +Create Date: 2026-03-11 + +This migration: +1. Adds internal_id UUID column to banks (stable identifier for index naming) +2. Drops the global HNSW index (competes with per-bank partial indexes) +3. Creates per-(bank_id, fact_type) partial HNSW indexes for all existing banks + (new banks get indexes created at bank-creation time via bank_utils.create_bank_hnsw_indexes) + +Why per-(bank, fact_type) indexes: +- fact_type-only partial indexes are never chosen by the planner when bank_id is in the WHERE + clause, because the idx_memory_units_bank_id B-tree index always wins at planning time. +- Per-(bank, fact_type) partial indexes have both predicates matching → planner selects them. +- The global HNSW index competes for larger partitions (world, observation) and must be dropped. + +For large deployments, create indexes CONCURRENTLY before running this migration: + SELECT internal_id, bank_id FROM banks; + -- for each bank and each fact_type in (world, experience, observation): + CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_mu_emb_{ft}_{uid16} + ON memory_units USING hnsw (embedding vector_cosine_ops) + WHERE fact_type = '{ft}' AND bank_id = '{bank_id}'; + DROP INDEX CONCURRENTLY IF EXISTS idx_memory_units_embedding; +""" + +from collections.abc import Sequence + +from alembic import context, op +from sqlalchemy import text + +revision: str = "d5e6f7a8b9c0" +down_revision: str | Sequence[str] | None = "c3d4e5f6g7h8" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +_HNSW_FACT_TYPES: dict[str, str] = { + "world": "worl", + "experience": "expr", + "observation": "obsv", +} + + +def _get_schema_prefix() -> str: + schema = context.config.get_main_option("target_schema") + return f'"{schema}".' if schema else "" + + +def upgrade() -> None: + schema = _get_schema_prefix() + + # 1. Add internal_id column to banks + op.execute( + f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS internal_id UUID DEFAULT gen_random_uuid() NOT NULL" + ) + op.execute(f"ALTER TABLE {schema}banks ADD CONSTRAINT banks_internal_id_unique UNIQUE (internal_id)") + + # 2. Drop any fact_type-only partial HNSW indexes that may exist from prior migrations + # (bank_id B-tree always wins over them when bank_id is in the WHERE clause) + op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_world") + op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_observation") + op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_experience") + + # 4. Drop global HNSW index (competes with per-bank partial indexes) + op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_embedding") + + # 5. Create per-(bank, fact_type) partial HNSW indexes for all existing banks + bind = op.get_bind() + schema_name = context.config.get_main_option("target_schema") + table_ref = f'"{schema_name}".memory_units' if schema_name else "memory_units" + banks_ref = f'"{schema_name}".banks' if schema_name else "banks" + + rows = bind.execute(text(f"SELECT bank_id, internal_id FROM {banks_ref}")).fetchall() # noqa: S608 + for row in rows: + bank_id = row[0] + internal_id = str(row[1]).replace("-", "")[:16] + escaped_bank_id = bank_id.replace("'", "''") + for ft, ft_short in _HNSW_FACT_TYPES.items(): + idx_name = f"idx_mu_emb_{ft_short}_{internal_id}" + # Index name is schema-unqualified (indexes live in the schema of their table) + bind.execute( + text( + f"CREATE INDEX IF NOT EXISTS {idx_name} " + f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) " + f"WHERE fact_type = '{ft}' AND bank_id = '{escaped_bank_id}'" + ) + ) + + +def downgrade() -> None: + schema = _get_schema_prefix() + + # Drop per-bank HNSW indexes (iterate existing banks) + bind = op.get_bind() + schema_name = context.config.get_main_option("target_schema") + banks_ref = f'"{schema_name}".banks' if schema_name else "banks" + + rows = bind.execute(text(f"SELECT internal_id FROM {banks_ref}")).fetchall() # noqa: S608 + for row in rows: + internal_id = str(row[0]).replace("-", "")[:16] + for ft_short in _HNSW_FACT_TYPES.values(): + idx_name = f"idx_mu_emb_{ft_short}_{internal_id}" + bind.execute(text(f"DROP INDEX IF EXISTS {schema}{idx_name}")) + + # Restore the global HNSW index + table_ref = f'"{schema_name}".memory_units' if schema_name else "memory_units" + op.execute( + f"CREATE INDEX IF NOT EXISTS idx_memory_units_embedding ON {table_ref} USING hnsw (embedding vector_cosine_ops)" + ) + + # Restore old fact_type-only partial indexes + op.execute( + f"CREATE INDEX IF NOT EXISTS idx_mu_emb_world " + f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) " + f"WHERE fact_type = 'world'" + ) + op.execute( + f"CREATE INDEX IF NOT EXISTS idx_mu_emb_observation " + f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) " + f"WHERE fact_type = 'observation'" + ) + op.execute( + f"CREATE INDEX IF NOT EXISTS idx_mu_emb_experience " + f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) " + f"WHERE fact_type = 'experience'" + ) + + # Drop internal_id column + op.execute(f"ALTER TABLE {schema}banks DROP CONSTRAINT IF EXISTS banks_internal_id_unique") + op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS internal_id") diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index a4baa56c..1da73cb8 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -1634,6 +1634,15 @@ class MemoryEngine(MemoryEngineInterface): # Create connection pool # For read-heavy workloads with many parallel think/search operations, # we need a larger pool. Read operations don't need strong isolation. + async def _init_connection(conn: asyncpg.Connection) -> None: + # SET (not SET LOCAL) so it persists for the connection lifetime. + # ef_search=200 improves HNSW recall quality for the per-fact_type + # semantic queries in retrieve_semantic_bm25_combined(). + try: + await conn.execute("SET hnsw.ef_search = 200") + except Exception: + logger.debug("Could not set hnsw.ef_search — extension may not support it") + self._pool = await asyncpg.create_pool( self.db_url, min_size=self._pool_min_size, @@ -1641,6 +1650,7 @@ class MemoryEngine(MemoryEngineInterface): command_timeout=self._db_command_timeout, statement_cache_size=0, # Disable prepared statement cache timeout=self._db_acquire_timeout, # Connection acquisition timeout (seconds) + init=_init_connection, ) # Initialize entity resolver with pool and configured lookup strategy @@ -3731,8 +3741,14 @@ class MemoryEngine(MemoryEngineInterface): # Delete entities (cascades to unit_entities, entity_cooccurrences, memory_links with entity_id) await conn.execute(f"DELETE FROM {fq_table('entities')} WHERE bank_id = $1", bank_id) - # Delete the bank profile itself - await conn.execute(f"DELETE FROM {fq_table('banks')} WHERE bank_id = $1", bank_id) + # Delete the bank profile and retrieve internal_id for HNSW index cleanup + internal_id = await conn.fetchval( + f"DELETE FROM {fq_table('banks')} WHERE bank_id = $1 RETURNING internal_id", bank_id + ) + + # Drop per-bank HNSW indexes now that the bank row is gone + if internal_id: + await bank_utils.drop_bank_hnsw_indexes(conn, str(internal_id)) result = { "memory_units_deleted": units_count, diff --git a/hindsight-api/hindsight_api/engine/retain/bank_utils.py b/hindsight-api/hindsight_api/engine/retain/bank_utils.py index 7755b52b..b1c98774 100644 --- a/hindsight-api/hindsight_api/engine/retain/bank_utils.py +++ b/hindsight-api/hindsight_api/engine/retain/bank_utils.py @@ -5,16 +5,65 @@ bank profile utilities for disposition and mission management. import json import logging import re +import uuid from typing import TypedDict from pydantic import BaseModel, Field from ..db_utils import acquire_with_retry -from ..memory_engine import fq_table +from ..memory_engine import fq_table, get_current_schema from ..response_models import DispositionTraits logger = logging.getLogger(__name__) +# Fact types that get per-bank partial HNSW indexes, mapped to their 4-char index suffix. +_HNSW_FACT_TYPES: dict[str, str] = { + "world": "worl", + "experience": "expr", + "observation": "obsv", +} + + +def _hnsw_index_name(ft: str, internal_id: str) -> str: + """Deterministic, schema-safe HNSW index name for a (bank, fact_type) pair. + + Uses the first 16 hex chars of internal_id (8 bytes of entropy) — unique + enough in practice, fits comfortably within PostgreSQL's 63-char identifier limit. + """ + uid = str(internal_id).replace("-", "")[:16] + return f"idx_mu_emb_{_HNSW_FACT_TYPES[ft]}_{uid}" + + +async def create_bank_hnsw_indexes(conn, bank_id: str, internal_id: str) -> None: + """Create per-(bank, fact_type) partial HNSW indexes for a newly created bank. + + Called immediately after the bank row is first inserted. Safe on empty banks + (index build is instant). Idempotent via CREATE INDEX IF NOT EXISTS. + bank_id is escaped for SQL literal safety (apostrophes doubled). + """ + table = fq_table("memory_units") + escaped = bank_id.replace("'", "''") + for ft in _HNSW_FACT_TYPES: + idx = _hnsw_index_name(ft, internal_id) + await conn.execute( + f"CREATE INDEX IF NOT EXISTS {idx} " + f"ON {table} USING hnsw (embedding vector_cosine_ops) " + f"WHERE fact_type = '{ft}' AND bank_id = '{escaped}'" + ) + + +async def drop_bank_hnsw_indexes(conn, internal_id: str) -> None: + """Drop per-(bank, fact_type) partial HNSW indexes for a bank being deleted. + + Called before the bank row is deleted so internal_id is still known. + Idempotent via DROP INDEX IF EXISTS. + """ + schema = get_current_schema() + for ft in _HNSW_FACT_TYPES: + idx = _hnsw_index_name(ft, internal_id) + await conn.execute(f"DROP INDEX IF EXISTS {schema}.{idx}") + + DEFAULT_DISPOSITION = { "skepticism": 3, "literalism": 3, @@ -70,19 +119,28 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile: mission=row["mission"] or "", ) - # Bank doesn't exist, create with defaults - await conn.execute( + # Bank doesn't exist, create with defaults. + # Generate internal_id here so we control the value and can use it + # immediately for HNSW index creation without a RETURNING round-trip. + internal_id = uuid.uuid4() + inserted = await conn.fetchval( f""" - INSERT INTO {fq_table("banks")} (bank_id, name, disposition, mission) - VALUES ($1, $2, $3::jsonb, $4) + INSERT INTO {fq_table("banks")} (bank_id, name, disposition, mission, internal_id) + VALUES ($1, $2, $3::jsonb, $4, $5) ON CONFLICT (bank_id) DO NOTHING + RETURNING bank_id """, bank_id, bank_id, # Default name is the bank_id json.dumps(DEFAULT_DISPOSITION), "", + internal_id, ) + if inserted: + # Fresh insert — create per-bank HNSW indexes (instant on empty bank) + await create_bank_hnsw_indexes(conn, bank_id, str(internal_id)) + return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), mission="") diff --git a/hindsight-api/hindsight_api/engine/retain/fact_storage.py b/hindsight-api/hindsight_api/engine/retain/fact_storage.py index 994b423f..c0ac9b12 100644 --- a/hindsight-api/hindsight_api/engine/retain/fact_storage.py +++ b/hindsight-api/hindsight_api/engine/retain/fact_storage.py @@ -6,9 +6,11 @@ Handles insertion of facts into the database. import json import logging +import uuid from ...config import get_config from ..memory_engine import fq_table +from .bank_utils import DEFAULT_DISPOSITION, create_bank_hnsw_indexes from .fact_extraction import _sanitize_text from .types import ProcessedFact @@ -183,17 +185,24 @@ async def ensure_bank_exists(conn, bank_id: str) -> None: conn: Database connection bank_id: Bank identifier """ - await conn.execute( + # Generate internal_id here so we control the value and can use it + # immediately for HNSW index creation without a RETURNING round-trip. + internal_id = uuid.uuid4() + inserted = await conn.fetchval( f""" - INSERT INTO {fq_table("banks")} (bank_id, disposition, mission) - VALUES ($1, $2::jsonb, $3) - ON CONFLICT (bank_id) DO UPDATE - SET updated_at = NOW() + INSERT INTO {fq_table("banks")} (bank_id, disposition, mission, internal_id) + VALUES ($1, $2::jsonb, $3, $4) + ON CONFLICT (bank_id) DO NOTHING + RETURNING bank_id """, bank_id, - '{"skepticism": 3, "literalism": 3, "empathy": 3}', + json.dumps(DEFAULT_DISPOSITION), "", + internal_id, ) + if inserted: + # Fresh insert — create per-bank HNSW indexes + await create_bank_hnsw_indexes(conn, bank_id, str(internal_id)) async def handle_document_tracking( diff --git a/hindsight-api/hindsight_api/engine/search/retrieval.py b/hindsight-api/hindsight_api/engine/search/retrieval.py index 0189c731..57f7963e 100644 --- a/hindsight-api/hindsight_api/engine/search/retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/retrieval.py @@ -98,8 +98,22 @@ async def retrieve_semantic_bm25_combined( """ Combined semantic + BM25 retrieval for multiple fact types in a single query. - Uses CTEs with window functions to get top-N results per fact type per method, - all in one database round-trip. + Uses UNION ALL of per-fact_type subqueries so that each arm has its own + ORDER BY ... LIMIT, enabling the partial HNSW indexes per fact_type instead + of forcing a full sequential scan (which the previous window-function approach + caused by using PARTITION BY inside ROW_NUMBER()). + + Requires partial HNSW indexes per fact_type (idx_mu_emb_world, + idx_mu_emb_observation, idx_mu_emb_experience), created automatically by + Alembic migration a3b4c5d6e7f8_add_partial_hnsw_indexes.py. + + HNSW is approximate — semantic arms over-fetch by 5x (min 100) and trim to + limit in Python to compensate. ef_search=200 is set globally on pool + connections at init time (see memory_engine.py) to improve recall on sparse + graphs. + + fact_type values are inlined as literals (safe: they come from a controlled + internal enum, never from user input). Args: conn: Database connection @@ -108,146 +122,120 @@ async def retrieve_semantic_bm25_combined( bank_id: Bank ID fact_types: List of fact types to retrieve limit: Maximum results per method per fact type + tags: Optional tags to filter by + tags_match: Tag matching mode Returns: Dict mapping fact_type -> (semantic_results, bm25_results) """ import re - # Sanitize query text for BM25 (same as retrieve_bm25) + result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {ft: ([], []) for ft in fact_types} + sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower()) tokens = [token for token in sanitized_text.split() if token] - # If no valid tokens for BM25, just run semantic - if not tokens: - tags_clause = build_tags_where_clause_simple(tags, 5, match=tags_match) - params = [query_emb_str, bank_id, fact_types, limit] - if tags: - params.append(tags) - results = await conn.fetch( - f""" - WITH semantic_ranked AS ( - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - 1 - (embedding <=> $1::vector) AS similarity, - NULL::float AS bm25_score, - 'semantic' AS source, - ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn - FROM {fq_table("memory_units")} - WHERE bank_id = $2 - AND embedding IS NOT NULL - AND fact_type = ANY($3) - AND (1 - (embedding <=> $1::vector)) >= 0.3 - {tags_clause} - ) - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - similarity, bm25_score, source - FROM semantic_ranked - WHERE rn <= $4 - """, - *params, + # Over-fetch for HNSW approximation; semantic results trimmed to limit in Python. + hnsw_fetch = max(limit * 5, 100) + + cols = ( + "id, text, context, event_date, occurred_start, occurred_end, mentioned_at, " + "fact_type, document_id, chunk_id, tags" + ) + table = fq_table("memory_units") + + # --- Parameter layout --- + # $1 = query_emb_str (semantic arms) + # $2 = bank_id + # $3 = limit (BM25 LIMIT; semantic uses inlined hnsw_fetch literal) + # $4 = bm25_text (only when tokens present) + # $N = tags (N=4 when no tokens, N=5 when tokens present) + tags_param_idx = 5 if tokens else 4 + tags_clause = build_tags_where_clause_simple(tags, tags_param_idx, match=tags_match) + + # --- Semantic UNION ALL arms (one per fact_type) --- + # Each arm has its own ORDER BY embedding <=> $1 LIMIT {hnsw_fetch}, which + # lets the planner use the partial HNSW index for that fact_type. + sem_arms = [] + for ft in fact_types: + sem_arms.append( + f"(SELECT {cols}," + f" 1 - (embedding <=> $1::vector) AS similarity," + f" NULL::float AS bm25_score," + f" 'semantic' AS source" + f" FROM {table}" + f" WHERE bank_id = $2" + f" AND fact_type = '{ft}'" + f" AND embedding IS NOT NULL" + f" AND (1 - (embedding <=> $1::vector)) >= 0.3" + f" {tags_clause}" + f" ORDER BY embedding <=> $1::vector" + f" LIMIT {hnsw_fetch})" ) - # Group by fact_type - result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = { - ft: ([], []) for ft in fact_types - } - for r in results: - row = dict(r) - ft = row.get("fact_type") - row.pop("source", None) - if ft in result_dict: - result_dict[ft][0].append(RetrievalResult.from_db_row(row)) - return result_dict - # Build BM25 query based on text search backend - config = get_config() + arms = sem_arms - # Build tags clause - param 6 if tags provided - tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match) + # --- BM25 UNION ALL arms (one per fact_type, only when tokens present) --- + if tokens: + config = get_config() + if config.text_search_extension == "vchord": + bm25_score_expr = ( + "search_vector <&> to_bm25query('idx_memory_units_text_search', tokenize($4, 'llmlingua2'))" + ) + bm25_order_by = f"{bm25_score_expr} DESC" + bm25_where_filter = "" + bm25_text_param: str = query_text + elif config.text_search_extension == "pg_textsearch": + bm25_score_expr = "-(text <@> to_bm25query($4, 'idx_memory_units_text_search'))" + bm25_order_by = "text <@> to_bm25query($4, 'idx_memory_units_text_search') ASC" + bm25_where_filter = "" + bm25_text_param = query_text + else: # native + query_tsquery = " | ".join(tokens) + bm25_score_expr = "ts_rank_cd(search_vector, to_tsquery('english', $4))" + bm25_order_by = f"{bm25_score_expr} DESC" + bm25_where_filter = "AND search_vector @@ to_tsquery('english', $4)" + bm25_text_param = query_tsquery - # Build backend-specific BM25 parts - if config.text_search_extension == "vchord": - # VectorChord BM25: use <&> operator with to_bm25query and tokenize - # Note: VectorChord scores are negative (higher = better, so -1 > -10) - bm25_score_expr = "search_vector <&> to_bm25query('idx_memory_units_text_search', tokenize($5, 'llmlingua2'))" - bm25_order_by = f"{bm25_score_expr} DESC" - bm25_where_filter = "" # No additional WHERE filter for vchord - params = [query_emb_str, bank_id, fact_types, limit, query_text] # Pass raw query_text for tokenization - elif config.text_search_extension == "pg_textsearch": - # Timescale pg_textsearch: use <@> operator with to_bm25query - # Note: pg_textsearch scores are negative (lower/more negative = better, so -10 > -1) - # We negate the score to maintain API consistency (higher = better) - bm25_score_expr = "-(text <@> to_bm25query($5, 'idx_memory_units_text_search'))" - bm25_order_by = "text <@> to_bm25query($5, 'idx_memory_units_text_search') ASC" - bm25_where_filter = "" # No additional WHERE filter for pg_textsearch - params = [query_emb_str, bank_id, fact_types, limit, query_text] - else: # native - # Native PostgreSQL: use ts_rank_cd with to_tsquery - query_tsquery = " | ".join(tokens) - bm25_score_expr = "ts_rank_cd(search_vector, to_tsquery('english', $5))" - bm25_order_by = f"{bm25_score_expr} DESC" - bm25_where_filter = "AND search_vector @@ to_tsquery('english', $5)" - params = [query_emb_str, bank_id, fact_types, limit, query_tsquery] + for ft in fact_types: + arms.append( + f"(SELECT {cols}," + f" NULL::float AS similarity," + f" {bm25_score_expr} AS bm25_score," + f" 'bm25' AS source" + f" FROM {table}" + f" WHERE bank_id = $2" + f" AND fact_type = '{ft}'" + f" {bm25_where_filter}" + f" {tags_clause}" + f" ORDER BY {bm25_order_by}" + f" LIMIT $3)" + ) + query = "\nUNION ALL\n".join(arms) + + params: list = [query_emb_str, bank_id, limit] + if tokens: + params.append(bm25_text_param) if tags: params.append(tags) - # Single query template with backend-specific parts injected - query = f""" - WITH semantic_ranked AS ( - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - 1 - (embedding <=> $1::vector) AS similarity, - NULL::float AS bm25_score, - 'semantic' AS source, - ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn - FROM {fq_table("memory_units")} - WHERE bank_id = $2 - AND embedding IS NOT NULL - AND fact_type = ANY($3) - AND (1 - (embedding <=> $1::vector)) >= 0.3 - {tags_clause} - ), - bm25_ranked AS ( - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - NULL::float AS similarity, - {bm25_score_expr} AS bm25_score, - 'bm25' AS source, - ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY {bm25_order_by}) AS rn - FROM {fq_table("memory_units")} - WHERE bank_id = $2 - AND fact_type = ANY($3) - {bm25_where_filter} - {tags_clause} - ), - semantic AS ( - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - similarity, bm25_score, source - FROM semantic_ranked WHERE rn <= $4 - ), - bm25 AS ( - SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, - similarity, bm25_score, source - FROM bm25_ranked WHERE rn <= $4 - ) - SELECT * FROM semantic - UNION ALL - SELECT * FROM bm25 - """ + rows = await conn.fetch(query, *params) - # Combined CTE query for both semantic and BM25 across all fact types - # Uses window functions to limit per fact_type per method - results = await conn.fetch(query, *params) - - # Group results by fact_type and source - result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {ft: ([], []) for ft in fact_types} - for r in results: + # Group results; trim semantic to limit (over-fetched for HNSW approximation). + sem_counts: dict[str, int] = {ft: 0 for ft in fact_types} + for r in rows: row = dict(r) - source = row.pop("source", None) + source = row.pop("source") ft = row.get("fact_type") - if ft in result_dict: - if source == "semantic": + if ft not in result_dict: + continue + if source == "semantic": + if sem_counts[ft] < limit: result_dict[ft][0].append(RetrievalResult.from_db_row(row)) - else: - result_dict[ft][1].append(RetrievalResult.from_db_row(row)) + sem_counts[ft] += 1 + else: + result_dict[ft][1].append(RetrievalResult.from_db_row(row)) return result_dict diff --git a/hindsight-api/hindsight_api/migrations.py b/hindsight-api/hindsight_api/migrations.py index 179fdc7f..bd656c44 100644 --- a/hindsight-api/hindsight_api/migrations.py +++ b/hindsight-api/hindsight_api/migrations.py @@ -717,6 +717,24 @@ def ensure_vector_extension( ).fetchone() if not current_index_info: + # Check whether per-bank partial HNSW indexes already cover this table + # (created by the bank_utils lifecycle — no global index needed in that case) + per_bank_index_count = conn.execute( + text(""" + SELECT COUNT(*) + FROM pg_indexes + WHERE schemaname = :schema + AND tablename = :table_name + AND indexname LIKE 'idx_mu_emb_%' + """), + {"schema": schema_name, "table_name": table_name}, + ).scalar() + if per_bank_index_count and per_bank_index_count > 0: + logger.debug( + f"No global embedding index on {table_name}, but {per_bank_index_count} " + f"per-bank partial HNSW indexes exist — skipping global index creation" + ) + continue logger.warning(f"No embedding index found for {table_name}, will create it") mismatched_tables.append((table_name, index_name, None)) continue diff --git a/hindsight-api/tests/test_hnsw_indexes.py b/hindsight-api/tests/test_hnsw_indexes.py new file mode 100644 index 00000000..a84b6be0 --- /dev/null +++ b/hindsight-api/tests/test_hnsw_indexes.py @@ -0,0 +1,193 @@ +""" +Tests for per-bank HNSW index lifecycle and UNION ALL retrieval. + +Covers: +- _hnsw_index_name deterministic naming +- Per-bank HNSW indexes created on bank creation (retain_async / ensure_bank_exists) +- Per-bank HNSW indexes dropped on bank deletion +- retrieve_semantic_bm25_combined groups results correctly by fact_type and source +""" +import uuid +from datetime import datetime, timezone + +import pytest + +from hindsight_api.engine.retain.bank_utils import _HNSW_FACT_TYPES, _hnsw_index_name + + +# --------------------------------------------------------------------------- +# Unit tests — no DB required +# --------------------------------------------------------------------------- + + +class TestHnswIndexName: + def test_deterministic(self): + uid = "550e8400-e29b-41d4-a716-446655440000" + assert _hnsw_index_name("world", uid) == _hnsw_index_name("world", uid) + + def test_strips_dashes(self): + uid = "550e8400-e29b-41d4-a716-446655440000" + name = _hnsw_index_name("world", uid) + # uid16 should be hex chars only + assert "-" not in name + + def test_uses_first_16_hex_chars(self): + uid = "550e8400-e29b-41d4-a716-446655440000" + uid16 = uid.replace("-", "")[:16] # "550e8400e29b41d4" + assert name_ends_with(name=_hnsw_index_name("world", uid), suffix=uid16) + + def test_suffix_per_fact_type(self): + uid = "550e8400-e29b-41d4-a716-446655440000" + names = {ft: _hnsw_index_name(ft, uid) for ft in _HNSW_FACT_TYPES} + # All three names must be distinct + assert len(set(names.values())) == 3 + + def test_all_fact_types_covered(self): + assert set(_HNSW_FACT_TYPES) == {"world", "experience", "observation"} + + def test_fits_pg_identifier_limit(self): + # PostgreSQL max identifier length is 63 chars + uid = "f" * 32 # simulated UUID without dashes + for ft in _HNSW_FACT_TYPES: + assert len(_hnsw_index_name(ft, uid)) <= 63 + + +def name_ends_with(name: str, suffix: str) -> bool: + return name.endswith(suffix) + + +# --------------------------------------------------------------------------- +# Integration tests — require DB (memory fixture) +# --------------------------------------------------------------------------- + + +async def _get_bank_hnsw_indexes(pool, bank_id: str) -> list[str]: + """Return index names for memory_units that match the per-bank pattern.""" + async with pool.acquire() as conn: + rows = await conn.fetch( + """ + SELECT indexname + FROM pg_indexes + WHERE tablename = 'memory_units' + AND indexname LIKE 'idx_mu_emb_%' + AND indexdef LIKE $1 + ORDER BY indexname + """, + f"%bank_id = '{bank_id}'%", + ) + return [row["indexname"] for row in rows] + + +@pytest.mark.asyncio +async def test_retain_creates_per_bank_hnsw_indexes(memory, request_context): + """retain_async on a new bank must create 3 per-(bank, fact_type) HNSW indexes.""" + bank_id = f"test_hnsw_create_{uuid.uuid4().hex[:8]}" + try: + await memory.retain_async( + bank_id=bank_id, + content="Alice is a software engineer.", + request_context=request_context, + ) + indexes = await _get_bank_hnsw_indexes(memory._pool, bank_id) + assert len(indexes) == 3, f"Expected 3 per-bank HNSW indexes, got: {indexes}" + for ft_short in _HNSW_FACT_TYPES.values(): + assert any(ft_short in idx for idx in indexes), ( + f"Missing index for fact_type short '{ft_short}' in {indexes}" + ) + finally: + await memory.delete_bank(bank_id, request_context=request_context) + + +@pytest.mark.asyncio +async def test_delete_bank_drops_hnsw_indexes(memory, request_context): + """delete_bank must drop all per-bank HNSW indexes.""" + bank_id = f"test_hnsw_drop_{uuid.uuid4().hex[:8]}" + + await memory.retain_async( + bank_id=bank_id, + content="Bob is a data scientist.", + request_context=request_context, + ) + # Verify indexes exist before deletion + indexes_before = await _get_bank_hnsw_indexes(memory._pool, bank_id) + assert len(indexes_before) == 3 + + await memory.delete_bank(bank_id, request_context=request_context) + + indexes_after = await _get_bank_hnsw_indexes(memory._pool, bank_id) + assert indexes_after == [], f"Indexes should be dropped after bank deletion, got: {indexes_after}" + + +@pytest.mark.asyncio +async def test_retain_idempotent_bank_creation(memory, request_context): + """Retaining into the same bank twice must not error and still have exactly 3 indexes.""" + bank_id = f"test_hnsw_idem_{uuid.uuid4().hex[:8]}" + try: + await memory.retain_async( + bank_id=bank_id, + content="Carol is a product manager.", + request_context=request_context, + ) + await memory.retain_async( + bank_id=bank_id, + content="Carol joined the company in 2022.", + request_context=request_context, + ) + indexes = await _get_bank_hnsw_indexes(memory._pool, bank_id) + assert len(indexes) == 3 + finally: + await memory.delete_bank(bank_id, request_context=request_context) + + +@pytest.mark.asyncio +async def test_retrieve_semantic_bm25_grouped_by_fact_type(memory, request_context): + """ + retrieve_semantic_bm25_combined must return a dict keyed by fact_type with + (semantic_list, bm25_list) tuples. All returned facts must belong to their + declared fact_type. + """ + from hindsight_api.engine.search.retrieval import retrieve_semantic_bm25_combined + + bank_id = f"test_retrieval_{uuid.uuid4().hex[:8]}" + try: + await memory.retain_async( + bank_id=bank_id, + content=( + "Alice is a software engineer at TechCorp. " + "She visited Paris in 2023 for a conference." + ), + context="background", + event_date=datetime(2023, 6, 1, tzinfo=timezone.utc), + request_context=request_context, + ) + + query_emb = memory.embeddings.encode(["software engineer Alice"]) + query_emb_str = str(query_emb[0]) + + fact_types = ["world", "experience"] + async with memory._pool.acquire() as conn: + results = await retrieve_semantic_bm25_combined( + conn=conn, + query_emb_str=query_emb_str, + query_text="software engineer Alice", + bank_id=bank_id, + fact_types=fact_types, + limit=5, + ) + + # Must return an entry for every requested fact_type + assert set(results.keys()) == set(fact_types) + + for ft, (sem, bm25) in results.items(): + # Semantic and BM25 lists must be lists + assert isinstance(sem, list) + assert isinstance(bm25, list) + # All semantic results must declare the correct fact_type + for r in sem: + assert r.fact_type == ft, f"Semantic result has wrong fact_type: {r.fact_type}" + # All BM25 results must declare the correct fact_type + for r in bm25: + assert r.fact_type == ft, f"BM25 result has wrong fact_type: {r.fact_type}" + + finally: + await memory.delete_bank(bank_id, request_context=request_context)