fix: (search) add logic to use external vectore store for message search (#479)

* fix: (search) add logic to use external vectore store for message search

* fix: (search) oversample to reduce duplicate errors
This commit is contained in:
Vineeth Voruganti 2026-04-02 13:00:54 -04:00 committed by GitHub
parent 0533c6dd26
commit d803c546f0
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 153 additions and 41 deletions

View File

@ -1,3 +1,4 @@
from collections.abc import Sequence
from datetime import datetime
from logging import getLogger
from typing import Any
@ -18,6 +19,21 @@ from .session import get_or_create_session
logger = getLogger(__name__)
def _deduplicate_messages(
messages: Sequence[models.Message], limit: int
) -> list[models.Message]:
"""Deduplicate messages by public_id, preserving input order."""
seen: set[str] = set()
result: list[models.Message] = []
for msg in messages:
if msg.public_id not in seen:
seen.add(msg.public_id)
result.append(msg)
if len(result) >= limit:
break
return result
def _apply_token_limit(
base_conditions: list[ColumnElement[Any]], token_limit: int
) -> Select[tuple[models.Message]]:
@ -578,6 +594,78 @@ async def update_message(
return honcho_message
async def _search_messages_external(
db: AsyncSession,
workspace_name: str,
query_embedding: list[float],
limit: int,
*,
session_name: str | None = None,
after_date: datetime | None = None,
before_date: datetime | None = None,
) -> list[models.Message]:
"""Query the external vector store for messages and fetch them from the DB.
Multiple vector records can map to the same message (chunked embeddings),
so we oversample from the vector store and deduplicate by message_id.
Date filters are applied at the DB level since external vector stores
don't support temporal filtering.
"""
external_vector_store = get_external_vector_store()
if external_vector_store is None:
return []
namespace = external_vector_store.get_vector_namespace("message", workspace_name)
vector_filters: dict[str, Any] = {}
if session_name:
vector_filters["session_name"] = session_name
# Oversample: chunks can map to the same message, and date filters are
# applied post-fetch (vector stores don't support temporal filtering),
# so fetch extra to compensate for both deduplication and filtering.
has_date_filters = after_date is not None or before_date is not None
oversample = 6 if has_date_filters else 3
vector_results = await external_vector_store.query(
namespace,
query_embedding,
top_k=limit * oversample,
filters=vector_filters if vector_filters else None,
)
if not vector_results:
return []
# Deduplicate by message_id preserving similarity order
seen: dict[str, None] = {}
for vr in vector_results:
mid = vr.metadata.get("message_id")
if mid and mid not in seen:
seen[mid] = None
message_ids = list(seen.keys())
if not message_ids:
return []
# Fetch from DB with optional date filtering
fetch_stmt = (
select(models.Message)
.where(models.Message.public_id.in_(message_ids))
.where(models.Message.workspace_name == workspace_name)
)
if after_date:
fetch_stmt = fetch_stmt.where(models.Message.created_at >= after_date)
if before_date:
fetch_stmt = fetch_stmt.where(models.Message.created_at <= before_date)
result = await db.execute(fetch_stmt)
messages_by_id = {msg.public_id: msg for msg in result.scalars().all()}
# Preserve vector store similarity order, apply limit
return [messages_by_id[mid] for mid in message_ids if mid in messages_by_id][:limit]
async def search_messages(
db: AsyncSession,
workspace_name: str,
@ -612,25 +700,36 @@ async def search_messages(
embedding if embedding is not None else await embedding_client.embed(query)
)
# First, find the top matching messages
match_stmt = (
select(models.Message)
.join(
models.MessageEmbedding,
models.Message.public_id == models.MessageEmbedding.message_id,
)
.where(models.MessageEmbedding.workspace_name == workspace_name)
.order_by(models.MessageEmbedding.embedding.cosine_distance(query_embedding))
.limit(limit)
)
if session_name:
match_stmt = match_stmt.where(
models.MessageEmbedding.session_name == session_name
if settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED:
# pgvector path: cosine distance in SQL
# Oversample because a message with multiple embedding chunks can
# produce duplicate rows; we deduplicate in Python to preserve HNSW
# index usage (a DISTINCT ON subquery would prevent the index scan).
match_stmt = (
select(models.Message)
.join(
models.MessageEmbedding,
models.Message.public_id == models.MessageEmbedding.message_id,
)
.where(models.MessageEmbedding.workspace_name == workspace_name)
.order_by(
models.MessageEmbedding.embedding.cosine_distance(query_embedding)
)
.limit(limit * 2)
)
result = await db.execute(match_stmt)
matched_messages = list(result.scalars().all())
if session_name:
match_stmt = match_stmt.where(
models.MessageEmbedding.session_name == session_name
)
result = await db.execute(match_stmt)
matched_messages = _deduplicate_messages(result.scalars().all(), limit)
else:
# External vector store path
matched_messages = await _search_messages_external(
db, workspace_name, query_embedding, limit, session_name=session_name
)
return await _build_merged_snippets(
db, workspace_name, matched_messages, context_window
@ -767,34 +866,47 @@ async def search_messages_temporal(
embedding if embedding is not None else await embedding_client.embed(query)
)
# Build query with date filters
match_stmt = (
select(models.Message)
.join(
models.MessageEmbedding,
models.Message.public_id == models.MessageEmbedding.message_id,
)
.where(models.MessageEmbedding.workspace_name == workspace_name)
)
if session_name:
match_stmt = match_stmt.where(
models.MessageEmbedding.session_name == session_name
if settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED:
# pgvector path: cosine distance in SQL with date filters
# Oversample to handle chunk duplicates (see search_messages comment)
match_stmt = (
select(models.Message)
.join(
models.MessageEmbedding,
models.Message.public_id == models.MessageEmbedding.message_id,
)
.where(models.MessageEmbedding.workspace_name == workspace_name)
)
# Apply date filters on the Message table
if after_date:
match_stmt = match_stmt.where(models.Message.created_at >= after_date)
if before_date:
match_stmt = match_stmt.where(models.Message.created_at <= before_date)
if session_name:
match_stmt = match_stmt.where(
models.MessageEmbedding.session_name == session_name
)
# Order by similarity and limit
match_stmt = match_stmt.order_by(
models.MessageEmbedding.embedding.cosine_distance(query_embedding)
).limit(limit)
# Apply date filters on the Message table
if after_date:
match_stmt = match_stmt.where(models.Message.created_at >= after_date)
if before_date:
match_stmt = match_stmt.where(models.Message.created_at <= before_date)
result = await db.execute(match_stmt)
matched_messages = list(result.scalars().all())
# Order by similarity and limit
match_stmt = match_stmt.order_by(
models.MessageEmbedding.embedding.cosine_distance(query_embedding)
).limit(limit * 2)
result = await db.execute(match_stmt)
matched_messages = _deduplicate_messages(result.scalars().all(), limit)
else:
# External vector store path with post-fetch date filtering
matched_messages = await _search_messages_external(
db,
workspace_name,
query_embedding,
limit,
session_name=session_name,
after_date=after_date,
before_date=before_date,
)
return await _build_merged_snippets(
db, workspace_name, matched_messages, context_window