fix: (search) add logic to use external vectore store for message search
This commit is contained in:
parent
0533c6dd26
commit
7fbac692ee
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue