From a26cfb2baf726adf42a4a00b89029c57e7679aba Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Thu, 2 Apr 2026 12:55:32 -0400 Subject: [PATCH] fix: (search) oversample to reduce duplicate errors --- src/crud/message.py | 36 ++++++++++++++++++++++++++++++------ 1 file changed, 30 insertions(+), 6 deletions(-) diff --git a/src/crud/message.py b/src/crud/message.py index 4bd7754f..41e0053b 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -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]]: @@ -606,11 +622,15 @@ async def _search_messages_external( if session_name: vector_filters["session_name"] = session_name - # Oversample: chunks can map to the same message, so fetch extra + # 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 * 3, + top_k=limit * oversample, filters=vector_filters if vector_filters else None, ) @@ -682,6 +702,9 @@ async def search_messages( 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( @@ -692,7 +715,7 @@ async def search_messages( .order_by( models.MessageEmbedding.embedding.cosine_distance(query_embedding) ) - .limit(limit) + .limit(limit * 2) ) if session_name: @@ -701,7 +724,7 @@ async def search_messages( ) result = await db.execute(match_stmt) - matched_messages = list(result.scalars().all()) + matched_messages = _deduplicate_messages(result.scalars().all(), limit) else: # External vector store path matched_messages = await _search_messages_external( @@ -845,6 +868,7 @@ async def search_messages_temporal( 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( @@ -868,10 +892,10 @@ async def search_messages_temporal( # Order by similarity and limit match_stmt = match_stmt.order_by( models.MessageEmbedding.embedding.cosine_distance(query_embedding) - ).limit(limit) + ).limit(limit * 2) result = await db.execute(match_stmt) - matched_messages = list(result.scalars().all()) + matched_messages = _deduplicate_messages(result.scalars().all(), limit) else: # External vector store path with post-fetch date filtering matched_messages = await _search_messages_external(