"""Hybrid search orchestration.""" from __future__ import annotations import asyncio import logging import re from collections import defaultdict from dataclasses import dataclass, replace from typing import TYPE_CHECKING from pgvector.sqlalchemy import Vector from sqlalchemy import literal, select from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from python.ebook_search.bm25_corpus import ( BM25CorpusUnavailableError, load_bm25_corpus, score_bm25_corpus, ) from python.ebook_search.chunk_records import CHUNK_RECORD_COLUMNS from python.ebook_search.embeddings import MODEL_DIMENSIONS, embed_query, get_embedding_table from python.ebook_search.protected_phrases.matching import ( detect_protected_phrases_for_query, phrase_hits_for_chunks, ) from python.ebook_search.rerank import rerank_chunks from python.ebook_search.timing import RuntimeStep, async_timed_result, timed_result from python.orm.richie import ( EbookChapter, EbookChunk, EbookEmbeddingModel, EbookSource, ) if TYPE_CHECKING: from collections.abc import Mapping, Sequence import httpx from sqlalchemy.ext.asyncio import AsyncEngine from python.ebook_search.config import EbookSearchConfig from python.ebook_search.protected_phrases.models import PhraseMatch logger = logging.getLogger(__name__) @dataclass(frozen=True) class SearchResult: """One source chunk returned by search.""" chunk_id: int text: str source_title: str source_id: int | None = None score: float = 0.0 vector_score: float | None = None bm25_score: float | None = None fused_score: float | None = None rerank_score: float | None = None phrase_hit_count: int = 0 matched_phrases: tuple[str, ...] = () source_author: str | None = None chapter_title: str | None = None page_label: str | None = None rank_source: str = "Hybrid" @dataclass(frozen=True) class SearchResponse: """Search output for the UI.""" query: str results: list[SearchResult] rank_label: str timings: tuple[RuntimeStep, ...] = () phrase_matches: tuple[PhraseMatch, ...] = () @property def total_runtime_ms(self) -> float: """Return total measured runtime for the response.""" return sum(step.duration_ms for step in self.timings if step.counts_toward_total) @dataclass(frozen=True) class RetrievalResponse: """Parallel retrieval output for vector, BM25, and protected phrase candidates.""" vector_results: list[SearchResult] lexical_results: list[SearchResult] phrase_matches: list[PhraseMatch] timings: tuple[RuntimeStep, ...] async def search_ebooks( engine: AsyncEngine, client: httpx.AsyncClient, query: str, config: EbookSearchConfig, *, rerank: bool, phrase_matching: bool, ) -> SearchResponse: """Run hybrid vector/BM25 search and optional reranking. Phrase matching only runs when both the request asks for it and ``config.phrase_matching_enabled`` allows it. """ if not query.strip(): logger.info("ebook_search_empty_query") return SearchResponse(query=query, results=[], rank_label="Hybrid") phrase_matching = phrase_matching and config.phrase_matching_enabled logger.info(f"ebook_search_start query_length={len(query)} {rerank=} {phrase_matching=}") timings: list[RuntimeStep] = [] retrieval, timing = await async_timed_result( "Hybrid retrieval", parallel_retrieval(engine, client, query, config, phrase_matching=phrase_matching), ) phrase_matches = retrieval.phrase_matches timings.extend(retrieval.timings) timings.append(timing) fused, timing = timed_result( "Reciprocal rank fusion", reciprocal_rank_fusion, retrieval.vector_results, retrieval.lexical_results, rank_constant=config.rrf_rank_constant, ) timings.append(timing) phrase_boost_timing_name = "Phrase mention boost" if phrase_matching else "Phrase mention boost skipped" fused, timing = await async_timed_result( phrase_boost_timing_name, apply_phrase_mention_boosts( engine, fused, phrase_matches, config.phrase_hit_boost, phrase_matching=phrase_matching, ), ) timings.append(timing) rerank_enabled = config.rerank.enabled and rerank rerank_timing_name = "Rerank" if rerank_enabled else "Rerank skipped" response, timing = await async_timed_result( rerank_timing_name, apply_rerank(client, query, fused, config, rerank=rerank_enabled), ) timings.append(timing) response = replace(response, timings=tuple(timings), phrase_matches=tuple(phrase_matches)) logger.info( f"ebook_search_complete vector_candidates={len(retrieval.vector_results)} " f"lexical_candidates={len(retrieval.lexical_results)} fused_candidates={len(fused)} {phrase_matching=} " f"phrase_matches={len(phrase_matches)} returned={len(response.results)} {response.rank_label=} " f"{response.total_runtime_ms=:.1f}" ) return response async def query_phrase_matches( engine: AsyncEngine, query: str, config: EbookSearchConfig, *, phrase_matching: bool, ) -> list[PhraseMatch]: """Detect protected phrases in a query without making search fail when phrase tables are unavailable.""" if not phrase_matching: logger.info("ebook_protected_phrase_detection_skipped") return [] try: async with AsyncSession(engine) as session: return await detect_protected_phrases_for_query(session, query, config) except SQLAlchemyError as error: logger.warning(f"ebook_protected_phrase_detection_unavailable {error=}") return [] async def apply_phrase_mention_boosts( engine: AsyncEngine, candidates: list[SearchResult], phrase_matches: Sequence[PhraseMatch], phrase_hit_boost: float, *, phrase_matching: bool, ) -> list[SearchResult]: """Boost retrieved chunks that have indexed mentions for detected protected phrases when enabled.""" if not phrase_matching: logger.info(f"ebook_phrase_boost_skipped candidates={len(candidates)}") return candidates phrase_ids = sorted({match.phrase_id for match in phrase_matches}) if not candidates or not phrase_ids or phrase_hit_boost <= 0: return candidates chunk_ids = [candidate.chunk_id for candidate in candidates] try: async with AsyncSession(engine) as session: phrase_hits = await phrase_hits_for_chunks(session, chunk_ids=chunk_ids, phrase_ids=phrase_ids) except SQLAlchemyError as error: logger.warning(f"ebook_phrase_boost_unavailable {error=}") return candidates if not phrase_hits: return candidates hit_counts = { chunk_id: sum(hit.mention_count for hit in chunk_hits) for chunk_id, chunk_hits in phrase_hits.items() } boosted = [ replace( candidate, score=candidate.score + (hit_counts.get(candidate.chunk_id, 0) * phrase_hit_boost), fused_score=boosted_fused_score(candidate, hit_counts.get(candidate.chunk_id, 0), phrase_hit_boost), phrase_hit_count=hit_counts.get(candidate.chunk_id, 0), matched_phrases=tuple(hit.phrase_text for hit in phrase_hits.get(candidate.chunk_id, ())), rank_source=phrase_rank_source(candidate.rank_source, hit_counts.get(candidate.chunk_id, 0)), ) for candidate in candidates ] return sorted(boosted, key=lambda candidate: candidate.score, reverse=True) def boosted_fused_score(candidate: SearchResult, phrase_hit_count: int, phrase_hit_boost: float) -> float | None: """Return a fused score adjusted by phrase hits when a fused score exists.""" if candidate.fused_score is None: return None return candidate.fused_score + (phrase_hit_count * phrase_hit_boost) def phrase_rank_source(rank_source: str, phrase_hit_count: int) -> str: """Append phrase evidence to a rank-source label when a chunk was boosted.""" if phrase_hit_count <= 0 or "phrases" in rank_source: return rank_source return f"{rank_source} + phrases" async def parallel_retrieval( engine: AsyncEngine, client: httpx.AsyncClient, query: str, config: EbookSearchConfig, *, phrase_matching: bool, ) -> RetrievalResponse: """Run vector, BM25, and protected phrase retrieval concurrently with separate database sessions. BM25 scoring is pure CPU work over the cached corpus, so it runs in a worker thread instead of on the event loop. Protected phrase detection only depends on the query, so it joins the gather as a third task and returns immediately when phrase matching is disabled. """ phrase_timing_name = "Protected phrase detection" if phrase_matching else "Protected phrase detection skipped" ( (vector_results, vector_timing), (lexical_results, lexical_timing), (phrase_matches, phrase_timing), ) = await asyncio.gather( async_timed_result("Embedding + vector search", vector_candidates(engine, client, query, config)), async_timed_result("BM25 search", asyncio.to_thread(bm25_candidates, query, config)), async_timed_result( phrase_timing_name, query_phrase_matches(engine, query, config, phrase_matching=phrase_matching), ), ) logger.info( f"ebook_parallel_retrieval_complete vector_candidates={len(vector_results)} " f"lexical_candidates={len(lexical_results)} phrase_matches={len(phrase_matches)}" ) return RetrievalResponse( vector_results=vector_results, lexical_results=lexical_results, phrase_matches=phrase_matches, timings=( replace(vector_timing, counts_toward_total=False), replace(lexical_timing, counts_toward_total=False), replace(phrase_timing, counts_toward_total=False), ), ) async def apply_rerank( client: httpx.AsyncClient, query: str, candidates: list[SearchResult], config: EbookSearchConfig, *, rerank: bool, ) -> SearchResponse: """Rerank already-fused hybrid candidates when enabled for this request.""" if not rerank: logger.info(f"ebook_rerank_skipped candidates={len(candidates)}") return SearchResponse(query=query, results=candidates[: config.top_k], rank_label="Hybrid") reranked = await rerank_chunks(client, query, candidates[: config.rerank.candidates], config.rerank) logger.info( f"ebook_rerank_complete input_candidates={min(len(candidates), config.rerank.candidates)} " f"returned={len(reranked)}" ) return SearchResponse( query=query, results=[replace(result, rank_source="Hybrid + rerank") for result in reranked[: config.top_k]], rank_label="Hybrid + rerank", ) async def vector_candidates( engine: AsyncEngine, client: httpx.AsyncClient, query: str, config: EbookSearchConfig, ) -> list[SearchResult]: """Return pgvector cosine candidates for a natural-language query.""" async with AsyncSession(engine) as session: model = await session.scalar( select(EbookEmbeddingModel).where(EbookEmbeddingModel.name == config.embedding_model) ) if model is None: msg = f"Embedding model is not registered: {config.embedding_model}" raise ValueError(msg) expected_dimension = MODEL_DIMENSIONS[config.embedding_model] if model.dimension != expected_dimension: msg = f"Model row dimension {model.dimension} does not match configured dimension {expected_dimension}" raise ValueError(msg) embedding = await embed_query(client, query, config) limit = max(config.rerank.candidates, config.top_k) * config.vector_candidate_multiplier embedding_table = get_embedding_table(model.dimension) embedding_param = literal(embedding, type_=Vector(model.dimension)) distance = embedding_table.embedding.op("<=>")(embedding_param) score = (literal(1.0) - distance).label("score") statement = ( select( *CHUNK_RECORD_COLUMNS, score, ) .select_from(embedding_table) .join(EbookChunk, EbookChunk.id == embedding_table.chunk_id) .join(EbookSource, EbookSource.id == EbookChunk.source_id) .outerjoin(EbookChapter, EbookChapter.id == EbookChunk.chapter_id) .where(embedding_table.model_id == model.id) .order_by(distance) .limit(limit) ) rows = (await session.execute(statement)).mappings() results = [search_result_from_row(row) for row in rows] logger.info( f"ebook_vector_search_complete {config.embedding_model=} {model.dimension=} candidates={len(results)}" ) return results def bm25_candidates(query: str, config: EbookSearchConfig) -> list[SearchResult]: """Return BM25-ranked lexical candidates using the persisted corpus.""" try: corpus = load_bm25_corpus(config) except BM25CorpusUnavailableError as error: logger.warning(f"ebook_bm25_index_unavailable_skipping {error=}") return [] if not corpus.records: logger.info("ebook_bm25_search_complete corpus=0 candidates=0") return [] bm25_query = retrieval_query_from_text(query) scored_records = score_bm25_corpus(bm25_query, corpus, limit=config.bm25_candidate_limit) results = [ replace(search_result_from_row(record), score=score, vector_score=None, bm25_score=score) for record, score in scored_records ] max_score = results[0].bm25_score if results else 0.0 logger.info(f"ebook_bm25_search_complete corpus={len(corpus.records)} candidates={len(results)} {max_score=:.6f}") return results def reciprocal_rank_fusion( vector_results: list[SearchResult], lexical_results: list[SearchResult], rank_constant: int, ) -> list[SearchResult]: """Fuse vector and lexical rankings with Reciprocal Rank Fusion.""" by_chunk: dict[int, SearchResult] = {} scores: defaultdict[int, float] = defaultdict(float) vector_scores: dict[int, float] = {} bm25_scores: dict[int, float] = {} for rank, result in enumerate(vector_results, start=1): by_chunk.setdefault(result.chunk_id, result) vector_scores[result.chunk_id] = result.vector_score if result.vector_score is not None else result.score scores[result.chunk_id] += 1 / (rank_constant + rank) for rank, result in enumerate(lexical_results, start=1): by_chunk.setdefault(result.chunk_id, result) bm25_scores[result.chunk_id] = result.bm25_score if result.bm25_score is not None else result.score scores[result.chunk_id] += 1 / (rank_constant + rank) return sorted( ( replace( result, score=scores[result.chunk_id], vector_score=vector_scores.get(result.chunk_id), bm25_score=bm25_scores.get(result.chunk_id), fused_score=scores[result.chunk_id], rank_source="Hybrid", ) for result in by_chunk.values() ), key=lambda result: result.score, reverse=True, ) def search_result_from_row(row: Mapping[str, object]) -> SearchResult: """Convert a database row mapping into a search result.""" source_id = row.get("source_id") return SearchResult( chunk_id=int(row["chunk_id"]), text=str(row["text"]), source_id=int(source_id) if source_id is not None else None, source_title=str(row["source_title"]), source_author=optional_str(row["source_author"]), chapter_title=optional_str(row["chapter_title"]), page_label=optional_str(row["page_label"]), score=float(row["score"]) if "score" in row else 0.0, vector_score=float(row["score"]) if "score" in row else None, ) def optional_str(value: object) -> str | None: """Convert nullable database values to optional strings.""" if value is None: return None return str(value) TOKEN_RE = re.compile(r"[A-Za-z0-9_]+") def tokens(text_value: str) -> list[str]: """Extract tokens from a text value. This is a simple approximation of the tokenization used by PostgreSQL's full-text search, which is sufficient for BM25 candidate retrieval. It lowercases tokens and includes alphanumeric characters and underscores. """ return [match.group(0).lower() for match in TOKEN_RE.finditer(text_value)] QUERY_STOP_WORDS = { "a", "an", "and", "are", "as", "at", "does", "for", "in", "is", "of", "the", "to", "what", "when", "where", "which", "who", "why", } def retrieval_query_from_text(query: str) -> str: """Remove generic question words while preserving entity and series terms.""" keywords = [token for token in tokens(query) if token not in QUERY_STOP_WORDS] if not keywords: return query return " ".join(keywords)