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

This commit is contained in:
Vineeth Voruganti 2026-04-02 12:35:59 -04:00
parent 0533c6dd26
commit 7fbac692ee
1 changed files with 129 additions and 41 deletions

View File

@ -578,6 +578,74 @@ 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, so fetch extra
vector_results = await external_vector_store.query(
namespace,
query_embedding,
top_k=limit * 3,
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 +680,33 @@ 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
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)
)
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 = list(result.scalars().all())
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 +843,46 @@ 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
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)
result = await db.execute(match_stmt)
matched_messages = list(result.scalars().all())
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