""" Reciprocal Rank Fusion (RRF) utilities for combining search results. RRF is a method to combine multiple ranked lists by computing the reciprocal of each item's rank in each list, then summing these reciprocal ranks. """ import re from typing import Any, TypeVar from sqlalchemy import Select, and_, func, or_, select from sqlalchemy.ext.asyncio import AsyncSession from src import models from src.config import settings from src.embedding_client import embedding_client from src.exceptions import ValidationException from src.models import session_peers_table from src.utils.filter import apply_filter from src.utils.formatting import ILIKE_ESCAPE_CHAR, escape_ilike_pattern from src.vector_store import get_external_vector_store T = TypeVar("T") def reciprocal_rank_fusion(*ranked_lists: list[T], k: int = 60, limit: int) -> list[T]: """ Combine multiple ranked lists using Reciprocal Rank Fusion (RRF). RRF assigns a score to each item based on the formula: RRF_score = sum(1 / (k + rank_i)) for all lists where the item appears Where: - k is a constant (typically 60) that controls the impact of high-ranked items - rank_i is the rank of the item in list i (1-indexed) Args: *ranked_lists: Variable number of ranked lists to combine k: RRF constant parameter (default: 60) limit: Maximum number of results to return Returns: list of items ranked by RRF score (highest score first) """ if not ranked_lists: return [] # dictionary to store RRF scores for each item rrf_scores: dict[T, float] = {} # Process each ranked list for ranked_list in ranked_lists: for rank, item in enumerate(ranked_list, 1): # 1-indexed ranking if item not in rrf_scores: rrf_scores[item] = 0.0 # Add reciprocal rank contribution from this list rrf_scores[item] += 1.0 / (k + rank) # Sort items by RRF score (descending order) sorted_items = sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True) # Extract just the items (not the scores) result = [item for item, _ in sorted_items] return result[:limit] async def _semantic_search( db: AsyncSession, query: str, workspace_name: str, limit: int, filters: dict[str, Any] | None = None, ) -> list[models.Message]: """ Perform semantic search using external vector store for message embeddings. Args: db: Database session query: Search query workspace_name: Name of the workspace to search in limit: Maximum number of results to return filters: Optional filters to apply at vector store level (supports: session_id, peer_id) Returns: list of messages ordered by semantic similarity """ try: embedding_query = await embedding_client.embed(query) except ValueError as e: raise ValidationException( f"Query exceeds maximum token limit of {settings.MAX_EMBEDDING_TOKENS}." ) from e # Query Postgres / pgvector directly if settings.EMBED_MESSAGES and ( settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED ): # Join message_embeddings with messages to get full message objects distance_expr = models.MessageEmbedding.embedding.cosine_distance( embedding_query ) stmt = ( select(models.Message) .join( models.MessageEmbedding, models.Message.public_id == models.MessageEmbedding.message_id, ) .where(models.MessageEmbedding.embedding.isnot(None)) .where(models.MessageEmbedding.workspace_name == workspace_name) ) # Apply all additional filters using the standard filter utility # filters dict uses external names (session_id, peer_id) which apply_filter will map # to internal column names (session_name, peer_name) if filters: # Create a copy with workspace added internal_filters = filters.copy() internal_filters["workspace_id"] = workspace_name stmt = apply_filter(stmt, models.Message, internal_filters) # Order by cosine distance and limit stmt = stmt.order_by(distance_expr).limit(limit) result = await db.execute(stmt) return list(result.scalars().all()) # FALLBACK: Use external vector store (Turbopuffer, LanceDB) external_vector_store = get_external_vector_store() if external_vector_store is None: return [] namespace = external_vector_store.get_vector_namespace("message", workspace_name) # Build vector store filters from the provided filters vector_filters: dict[str, Any] = {} if filters: # Map external filter keys to vector store metadata keys if "session_id" in filters: vector_filters["session_name"] = filters["session_id"] if "peer_id" in filters: vector_filters["peer_name"] = filters["peer_id"] # Query external vector store for similar message embeddings # Since all filters are applied at the vector store level, we don't need to oversample vector_results = await external_vector_store.query( namespace, embedding_query, top_k=limit, filters=vector_filters if vector_filters else None, ) if not vector_results: return [] # Extract message IDs from vector metadata # Use dict to deduplicate while preserving order (dict keys maintain insertion order in Python 3.7+) seen_message_ids: dict[str, None] = {} for result in vector_results: message_id = result.metadata.get("message_id") if message_id and message_id not in seen_message_ids: seen_message_ids[message_id] = None message_ids = list(seen_message_ids.keys()) # Fetch messages from database by the IDs from vector search and reapply filters semantic_query = select(models.Message).where( models.Message.public_id.in_(message_ids) ) semantic_query = apply_filter(semantic_query, models.Message, filters) result = await db.execute(semantic_query) messages = {msg.public_id: msg for msg in result.scalars().all()} # Return messages in order of similarity (preserving vector store order) ordered_messages: list[models.Message] = [] for msg_id in message_ids: if msg_id in messages: ordered_messages.append(messages[msg_id]) return ordered_messages async def _filter_by_peer_perspective( db: AsyncSession, messages: list[models.Message], workspace_name: str, peer_name: str, ) -> list[models.Message]: """ Filter messages by peer perspective (temporal session membership). Only keeps messages from sessions where the peer was a member at the time the message was created (between joined_at and left_at). Args: db: Database session messages: List of messages to filter workspace_name: Name of the workspace peer_name: Name of the peer whose perspective to use Returns: Filtered list of messages """ if not messages: return [] # Get all session memberships for this peer in this workspace session_memberships_query = ( select(session_peers_table) .where(session_peers_table.c.workspace_name == workspace_name) .where(session_peers_table.c.peer_name == peer_name) ) result = await db.execute(session_memberships_query) memberships = result.all() # Build a lookup of session -> time windows session_windows: dict[str, list[tuple[Any, Any]]] = {} for membership in memberships: session_name = membership.session_name if session_name not in session_windows: session_windows[session_name] = [] session_windows[session_name].append((membership.joined_at, membership.left_at)) # Filter messages filtered_messages: list[models.Message] = [] for msg in messages: if msg.session_name not in session_windows: continue # Check if message was created during any of the peer's active windows in this session for joined_at, left_at in session_windows[msg.session_name]: if msg.created_at >= joined_at and ( left_at is None or msg.created_at <= left_at ): filtered_messages.append(msg) break # Don't add the same message twice return filtered_messages async def _fulltext_search( db: AsyncSession, query: str, stmt: Select[tuple[models.Message]], limit: int, ) -> list[models.Message]: """ Perform full-text search using PostgreSQL FTS and ILIKE fallback. Args: db: Database session query: Search query stmt: Base SQL query conditions limit: Maximum number of results to return Returns: list of messages ordered by text search relevance """ # Check if query contains special characters that FTS might not handle well has_special_chars = bool( re.search(r'[~`!@#$%^&*()_+=\[\]{};\':"\\|,.<>/?-]', query) ) # Escape ILIKE pattern characters to treat user input literally escaped_query = escape_ilike_pattern(query) if has_special_chars: # For queries with special characters, use exact string matching (ILIKE) search_condition = models.Message.content.ilike( f"%{escaped_query}%", escape=ILIKE_ESCAPE_CHAR ) fulltext_query = stmt.where(search_condition).order_by( models.Message.created_at.desc() ) else: # For natural language queries, use full text search with ranking fts_condition = func.to_tsvector("english", models.Message.content).op("@@")( func.plainto_tsquery("english", query) ) # Combine FTS with ILIKE as fallback for better coverage combined_condition = or_( fts_condition, models.Message.content.ilike( f"%{escaped_query}%", escape=ILIKE_ESCAPE_CHAR ), ) fulltext_query = stmt.where(combined_condition).order_by( # Order by FTS relevance first, then by creation time func.coalesce( func.ts_rank( func.to_tsvector("english", models.Message.content), func.plainto_tsquery("english", query), ), 0, ).desc(), models.Message.created_at.desc(), ) fulltext_query = fulltext_query.limit(limit) result = await db.execute(fulltext_query) return list(result.scalars().all()) async def search( db: AsyncSession, query: str, *, filters: dict[str, Any] | None = None, limit: int = 10, ) -> list[models.Message]: """ Search across message content using a hybrid approach with Reciprocal Rank Fusion (RRF). This function combines semantic search and full-text search results using RRF when both are available, providing better search results than either method alone. Args: db: Database session query: Search query to match against message content filters: Optional filters to scope search (must include workspace_id for semantic search). Special filter 'peer_perspective' will search across all messages from sessions that the peer is/was a member of, filtered by the time window when they were actually in the session. limit: Maximum number of results to return Returns: list of messages that match the search query, ordered by RRF relevance or individual search relevance Raises: ValidationException: If query exceeds maximum token limit for embeddings """ # Base query conditions stmt = select(models.Message) # Handle special peer_perspective filter peer_perspective_name: str | None = None if filters and "peer_perspective" in filters: peer_perspective_name = filters["peer_perspective"] # Remove from filters dict so apply_filter doesn't try to handle it filters = {k: v for k, v in filters.items() if k != "peer_perspective"} # Safety: peer_perspective must be scoped to a workspace if not filters or ( "workspace_id" not in filters and "workspace_name" not in filters ): raise ValidationException( "peer_perspective requires a workspace scope (workspace_id or workspace_name)." ) # Join with session_peers_table to get messages from sessions the peer was in # Only include messages created during the time window the peer was active stmt = stmt.join( session_peers_table, and_( models.Message.session_name == session_peers_table.c.session_name, models.Message.workspace_name == session_peers_table.c.workspace_name, models.Message.created_at >= session_peers_table.c.joined_at, or_( session_peers_table.c.left_at.is_(None), models.Message.created_at <= session_peers_table.c.left_at, ), ), ).where(session_peers_table.c.peer_name == peer_perspective_name) stmt = apply_filter(stmt, models.Message, filters) search_results: list[list[models.Message]] = [] # Perform semantic search if enabled and we have workspace context # workspace_id is required for semantic search to determine the vector namespace workspace_name: str | None = filters.get("workspace_id") if filters else None if settings.EMBED_MESSAGES and isinstance(workspace_name, str): # Type narrowing: workspace_name is guaranteed to be str in this block # Get more results for fusion (increase if peer_perspective filtering is applied post-search) semantic_limit = limit * 4 if peer_perspective_name else limit * 2 semantic_results = await _semantic_search( db=db, query=query, workspace_name=workspace_name, limit=semantic_limit, filters=filters, ) # Apply peer_perspective filtering to semantic results if needed # Vector store can't handle temporal filtering (joined_at/left_at), so filter post-search if peer_perspective_name: semantic_results = await _filter_by_peer_perspective( db, semantic_results, workspace_name, peer_perspective_name ) search_results.append(semantic_results) # Perform full-text search # Get more results for fusion fulltext_limit = limit * 2 fulltext_results = await _fulltext_search( db=db, query=query, stmt=stmt, limit=fulltext_limit ) search_results.append(fulltext_results) # Combine results using RRF if we have multiple search methods if len(search_results) > 1: # Use RRF to combine semantic and full-text results combined_results = reciprocal_rank_fusion(*search_results, limit=limit) elif len(search_results) == 1: # Single search method - apply limit directly combined_results = search_results[0] combined_results = combined_results[:limit] else: # No search results combined_results = [] return combined_results