perf: replace window-function retrieval with UNION ALL + per-bank HNSW indexes (#541)
* 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_<uid16>). * fix: use correct embeddings.encode() in test
This commit is contained in:
parent
00ac3d8834
commit
43b3efc494
7 changed files with 547 additions and 134 deletions
|
|
@ -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")
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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="")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
193
hindsight-api/tests/test_hnsw_indexes.py
Normal file
193
hindsight-api/tests/test_hnsw_indexes.py
Normal file
|
|
@ -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)
|
||||
Loading…
Reference in a new issue