refactor(ebook-search): simplify search and phrase matching
treefmt / nix fmt (pull_request) Successful in 5s
pytest / pytest (pull_request) Successful in 28s
test ebook search / test-ebook-search (pull_request) Failing after 35s
build_systems / build-bob (pull_request) Successful in 51s
build_systems / build-brain (pull_request) Successful in 50s
build_systems / build-rhapsody-in-green (pull_request) Successful in 1m3s
build_systems / build-jeeves (pull_request) Successful in 2m20s

This commit is contained in:
2026-07-15 15:25:11 -04:00
parent c9f859cac5
commit ceb2dbb2b3
16 changed files with 490 additions and 557 deletions
+6 -11
View File
@@ -13,7 +13,7 @@ from python.ebook_search.api.dependencies import ( # noqa: TC001 FastAPI resolv
AppEngine, AppEngine,
AppHttpClient, AppHttpClient,
) )
from python.ebook_search.api.web import templates from python.ebook_search.api.web import error_response, templates
from python.ebook_search.embeddings import embed_missing_chunks, embedding_model_stats from python.ebook_search.embeddings import embed_missing_chunks, embedding_model_stats
from python.ebook_search.ingest import ingest_configured_paths from python.ebook_search.ingest import ingest_configured_paths
from python.ebook_search.protected_phrases.generate_ngrams import generate_candidate_phrases_for_books from python.ebook_search.protected_phrases.generate_ngrams import generate_candidate_phrases_for_books
@@ -50,7 +50,7 @@ async def scan_library(request: Request, config: AppConfig, session: AsyncDbSess
await session.commit() await session.commit()
except Exception as error: except Exception as error:
logger.exception("ebook_admin_scan_failed") logger.exception("ebook_admin_scan_failed")
return templates.TemplateResponse(request, "partials/error.html", {"message": str(error)}, status_code=500) return error_response(request, error)
logger.info(f"ebook_admin_scan_complete {count=}") logger.info(f"ebook_admin_scan_complete {count=}")
if count > 0: if count > 0:
@@ -65,7 +65,7 @@ async def generate_all_phrases(request: Request, config: AppConfig, engine: AppE
result = await generate_candidate_phrases_for_books(engine, config) result = await generate_candidate_phrases_for_books(engine, config)
except Exception as error: except Exception as error:
logger.exception("ebook_admin_generate_phrases_failed") logger.exception("ebook_admin_generate_phrases_failed")
return templates.TemplateResponse(request, "partials/error.html", {"message": str(error)}, status_code=500) return error_response(request, error)
logger.info( logger.info(
f"ebook_admin_generate_phrases_complete {result.books_seen=} {result.books_built=} {result.candidate_phrases=}" f"ebook_admin_generate_phrases_complete {result.books_seen=} {result.books_built=} {result.candidate_phrases=}"
@@ -128,7 +128,7 @@ async def run_phrase_judgment(
result = await judge_candidate_phrases_for_books(engine, config, source_ids=source_ids) result = await judge_candidate_phrases_for_books(engine, config, source_ids=source_ids)
except Exception as error: except Exception as error:
logger.exception("ebook_admin_judge_phrases_failed") logger.exception("ebook_admin_judge_phrases_failed")
return templates.TemplateResponse(request, "partials/error.html", {"message": str(error)}, status_code=500) return error_response(request, error)
logger.info( logger.info(
f"ebook_admin_judge_phrases_complete {result.books_seen=} {result.books_judged=} {result.books_failed=} " f"ebook_admin_judge_phrases_complete {result.books_seen=} {result.books_judged=} {result.books_failed=} "
@@ -161,7 +161,7 @@ async def embed_missing(
await session.commit() await session.commit()
except Exception as error: except Exception as error:
logger.exception("ebook_admin_embed_missing_failed") logger.exception("ebook_admin_embed_missing_failed")
return templates.TemplateResponse(request, "partials/error.html", {"message": str(error)}, status_code=500) return error_response(request, error)
logger.info(f"ebook_admin_embed_missing_complete {count=}") logger.info(f"ebook_admin_embed_missing_complete {count=}")
return templates.TemplateResponse( return templates.TemplateResponse(
@@ -192,12 +192,7 @@ async def embed_all(
logger.info(f"ebook_admin_embed_all_batch_complete {batches=} {count=} {total=}") logger.info(f"ebook_admin_embed_all_batch_complete {batches=} {count=} {total=}")
except Exception as error: except Exception as error:
logger.exception(f"ebook_admin_embed_all_failed {batches=} {total=}") logger.exception(f"ebook_admin_embed_all_failed {batches=} {total=}")
return templates.TemplateResponse( return error_response(request, f"Embed all failed after {total} chunks in {batches} batches: {error}")
request,
"partials/error.html",
{"message": f"Embed all failed after {total} chunks in {batches} batches: {error}"},
status_code=500,
)
logger.info(f"ebook_admin_embed_all_complete {batches=} {total=}") logger.info(f"ebook_admin_embed_all_complete {batches=} {total=}")
return templates.TemplateResponse( return templates.TemplateResponse(
+2 -9
View File
@@ -15,6 +15,7 @@ from python.ebook_search.api.dependencies import (
from python.ebook_search.api.judge_tasks import is_judging_book, pop_book_judgment_outcome, start_book_phrase_judgment from python.ebook_search.api.judge_tasks import is_judging_book, pop_book_judgment_outcome, start_book_phrase_judgment
from python.ebook_search.api.web import templates from python.ebook_search.api.web import templates
from python.ebook_search.protected_phrases.generate_ngrams import recalculate_candidate_phrases_for_book from python.ebook_search.protected_phrases.generate_ngrams import recalculate_candidate_phrases_for_book
from python.ebook_search.protected_phrases.store import count_protected_phrases
from python.fastapi_tools import AsyncDbSession # noqa: TC001 FastAPI resolves this annotated dependency at runtime from python.fastapi_tools import AsyncDbSession # noqa: TC001 FastAPI resolves this annotated dependency at runtime
from python.orm.richie import EbookCandidatePhrase, EbookChapter, EbookChunk, EbookProtectedPhrase, EbookSource from python.orm.richie import EbookCandidatePhrase, EbookChapter, EbookChunk, EbookProtectedPhrase, EbookSource
@@ -71,14 +72,6 @@ async def get_judged_candidate_count(session: AsyncSession, book_id: int) -> int
) )
async def get_protected_count(session: AsyncSession, book_id: int) -> int:
"""Return the number of protected phrases for one book."""
return (
await session.scalar(select(func.count(EbookProtectedPhrase.id)).where(EbookProtectedPhrase.book_id == book_id))
or 0
)
async def get_candidates(session: AsyncSession, book_id: int) -> list[EbookCandidatePhrase]: async def get_candidates(session: AsyncSession, book_id: int) -> list[EbookCandidatePhrase]:
"""Return the indexed candidates for one book.""" """Return the indexed candidates for one book."""
return list( return list(
@@ -122,7 +115,7 @@ async def book_detail(source_id: int, request: Request, session: AsyncDbSession)
chunk_count = await get_chunk_count(session, source.id) chunk_count = await get_chunk_count(session, source.id)
candidate_count = await get_candidate_count(session, source.id) candidate_count = await get_candidate_count(session, source.id)
judged_candidate_count = await get_judged_candidate_count(session, source.id) judged_candidate_count = await get_judged_candidate_count(session, source.id)
protected_count = await get_protected_count(session, source.id) protected_count = await count_protected_phrases(session, source.id)
candidates = await get_candidates(session, source.id) candidates = await get_candidates(session, source.id)
protected_phrases = await get_protected_phrases(session, source.id) protected_phrases = await get_protected_phrases(session, source.id)
else: else:
+2 -2
View File
@@ -16,7 +16,7 @@ from python.ebook_search.api.dependencies import ( # noqa: TC001 FastAPI resolv
AppEngine, AppEngine,
AppHttpClient, AppHttpClient,
) )
from python.ebook_search.api.web import templates from python.ebook_search.api.web import error_response, templates
from python.ebook_search.guardrails import ( from python.ebook_search.guardrails import (
CitationReport, CitationReport,
is_confident, is_confident,
@@ -95,7 +95,7 @@ async def search(
) )
except Exception as error: except Exception as error:
logger.exception("ebook_search_request_failed") logger.exception("ebook_search_request_failed")
return templates.TemplateResponse(request, "partials/error.html", {"message": str(error)}, status_code=500) return error_response(request, error)
answer_start = perf_counter() answer_start = perf_counter()
answer, low_confidence, citation_report = await build_answer(client, query, response, config) answer, low_confidence, citation_report = await build_answer(client, query, response, config)
+10
View File
@@ -3,9 +3,14 @@
from __future__ import annotations from __future__ import annotations
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING
from fastapi.templating import Jinja2Templates from fastapi.templating import Jinja2Templates
if TYPE_CHECKING:
from fastapi import Request
from fastapi.responses import HTMLResponse
PACKAGE_DIR = Path(__file__).resolve().parent PACKAGE_DIR = Path(__file__).resolve().parent
TEMPLATE_DIR = PACKAGE_DIR / "templates" TEMPLATE_DIR = PACKAGE_DIR / "templates"
STATIC_DIR = PACKAGE_DIR / "static" STATIC_DIR = PACKAGE_DIR / "static"
@@ -21,3 +26,8 @@ def static_version(filename: str) -> int:
templates = Jinja2Templates(directory=TEMPLATE_DIR) templates = Jinja2Templates(directory=TEMPLATE_DIR)
templates.env.globals["static_version"] = static_version templates.env.globals["static_version"] = static_version
def error_response(request: Request, message: object) -> HTMLResponse:
"""Render the shared error partial for a failed UI request."""
return templates.TemplateResponse(request, "partials/error.html", {"message": str(message)}, status_code=500)
+2 -7
View File
@@ -15,6 +15,7 @@ from typing import TYPE_CHECKING
import bm25s import bm25s
from sqlalchemy import func, select, union_all from sqlalchemy import func, select, union_all
from python.ebook_search.chunk_records import CHUNK_RECORD_COLUMNS
from python.orm.richie import EbookChapter, EbookChunk, EbookSource from python.orm.richie import EbookChapter, EbookChunk, EbookSource
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -169,13 +170,7 @@ async def fetch_bm25_corpus_records(session: AsyncSession) -> tuple[list[dict[st
""" """
statement = ( statement = (
select( select(
EbookChunk.id.label("chunk_id"), *CHUNK_RECORD_COLUMNS,
EbookChunk.text.label("text"),
EbookSource.id.label("source_id"),
EbookSource.title.label("source_title"),
EbookSource.author.label("source_author"),
EbookChapter.title.label("chapter_title"),
EbookChunk.page_label.label("page_label"),
EbookChunk.search_text.label("bm25_text"), EbookChunk.search_text.label("bm25_text"),
) )
.select_from(EbookChunk) .select_from(EbookChunk)
+13
View File
@@ -0,0 +1,13 @@
"""Shared database columns used to build search-result records."""
from python.orm.richie import EbookChapter, EbookChunk, EbookSource
CHUNK_RECORD_COLUMNS = (
EbookChunk.id.label("chunk_id"),
EbookChunk.text.label("text"),
EbookSource.id.label("source_id"),
EbookSource.title.label("source_title"),
EbookSource.author.label("source_author"),
EbookChapter.title.label("chapter_title"),
EbookChunk.page_label.label("page_label"),
)
-6
View File
@@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
from os import getenv
from typing import Annotated, Self from typing import Annotated, Self
from pydantic import AliasChoices, Field, field_validator, model_validator from pydantic import AliasChoices, Field, field_validator, model_validator
@@ -32,11 +31,6 @@ def normalize_embedding_alias(model: str) -> str:
return standard_model return standard_model
def normalize_embedding_model(default: str = "qwen3-embedding-0.6b") -> str:
"""Normalize the configured embedding alias to its provider model name."""
return normalize_embedding_alias(getenv("EBOOK_SEARCH_EMBEDDING_MODEL", default))
class RerankConfig(BaseSettings): class RerankConfig(BaseSettings):
"""vLLM reranker settings.""" """vLLM reranker settings."""
+28 -14
View File
@@ -64,17 +64,13 @@ async def check_embedding_endpoint(
timeout_seconds: float = 5.0, timeout_seconds: float = 5.0,
) -> bool: ) -> bool:
"""Return whether the configured embedding endpoint answers a model listing.""" """Return whether the configured embedding endpoint answers a model listing."""
try: return await _check_endpoint(
response = await client.get( client,
f"{config.embedding_base_url.rstrip('/')}/models", base_url=config.embedding_base_url,
headers=auth_headers(config.embedding_api_key), api_key=config.embedding_api_key,
timeout=timeout_seconds, timeout_seconds=timeout_seconds,
) unavailable_log=f"ebook_embedding_endpoint_unreachable {config.embedding_base_url=}",
response.raise_for_status() )
except httpx.HTTPError as error:
logger.warning(f"ebook_embedding_endpoint_unreachable {config.embedding_base_url=} {error=}")
return False
return True
async def check_chat_endpoint( async def check_chat_endpoint(
@@ -84,15 +80,33 @@ async def check_chat_endpoint(
timeout_seconds: float = 5.0, timeout_seconds: float = 5.0,
) -> bool: ) -> bool:
"""Return whether the configured chat (answering) endpoint answers a model listing.""" """Return whether the configured chat (answering) endpoint answers a model listing."""
return await _check_endpoint(
client,
base_url=config.vllm_base_url,
api_key=config.vllm_api_key,
timeout_seconds=timeout_seconds,
unavailable_log=f"ebook_chat_endpoint_unreachable {config.vllm_base_url=}",
)
async def _check_endpoint(
client: httpx.AsyncClient,
*,
base_url: str,
api_key: str,
timeout_seconds: float,
unavailable_log: str,
) -> bool:
"""Return whether an OpenAI-compatible endpoint answers a model listing."""
try: try:
response = await client.get( response = await client.get(
f"{config.vllm_base_url.rstrip('/')}/models", f"{base_url.rstrip('/')}/models",
headers=auth_headers(config.vllm_api_key), headers=auth_headers(api_key),
timeout=timeout_seconds, timeout=timeout_seconds,
) )
response.raise_for_status() response.raise_for_status()
except httpx.HTTPError as error: except httpx.HTTPError as error:
logger.warning(f"ebook_chat_endpoint_unreachable {config.vllm_base_url=} {error=}") logger.warning(f"{unavailable_log} {error=}")
return False return False
return True return True
@@ -64,40 +64,30 @@ def normalize_candidate_phrase(
phrase_text: str, phrase_text: str,
config: EbookSearchConfig, config: EbookSearchConfig,
*, *,
min_tokens: int | None = None,
max_tokens: int | None = None, max_tokens: int | None = None,
strip_leading_article: bool = False,
) -> tuple[str, str, int] | None: ) -> tuple[str, str, int] | None:
"""Normalize a candidate phrase and validate token bounds. """Normalize a candidate phrase and validate token bounds.
Args: Args:
phrase_text (str): Raw phrase text to normalize. phrase_text (str): Raw phrase text to normalize.
config (EbookSearchConfig): Runtime phrase-tuning settings. config (EbookSearchConfig): Runtime phrase-tuning settings.
min_tokens (int | None): Minimum token count override; defaults to ``config.phrase_min_tokens``.
max_tokens (int | None): Maximum token count override; defaults to ``config.phrase_max_tokens``. max_tokens (int | None): Maximum token count override; defaults to ``config.phrase_max_tokens``.
strip_leading_article (bool): Whether to drop a single leading English article.
Returns: Returns:
tuple[str, str, int] | None: Display text, normalized phrase, and token count, or ``None`` tuple[str, str, int] | None: Display text, normalized phrase, and token count, or ``None``
when the phrase falls outside the token bounds or is ignored. when the phrase falls outside the token bounds or is ignored.
""" """
normalized_tokens = tokenize_with_offsets(phrase_text) normalized_tokens = tokenize_with_offsets(phrase_text)
start = 0
if strip_leading_article and normalized_tokens and normalized_tokens[0].text in {"the", "a", "an"}:
start = 1
selected_tokens = normalized_tokens[start:]
min_count = config.phrase_min_tokens if min_tokens is None else min_tokens
max_count = config.phrase_max_tokens if max_tokens is None else max_tokens max_count = config.phrase_max_tokens if max_tokens is None else max_tokens
if len(selected_tokens) < min_count or len(selected_tokens) > max_count: if len(normalized_tokens) < config.phrase_min_tokens or len(normalized_tokens) > max_count:
return None return None
phrase_norm = " ".join(token.text for token in selected_tokens) phrase_norm = " ".join(token.text for token in normalized_tokens)
if phrase_norm in get_ignored_phrases(): if phrase_norm in get_ignored_phrases():
return None return None
display_text = phrase_text[selected_tokens[0].start_char : selected_tokens[-1].end_char].strip() display_text = phrase_text[normalized_tokens[0].start_char : normalized_tokens[-1].end_char].strip()
return display_text or phrase_norm, phrase_norm, len(selected_tokens) return display_text or phrase_norm, phrase_norm, len(normalized_tokens)
def count_raw_ngrams(tokens: Sequence[str], config: EbookSearchConfig) -> Counter[str]: def count_raw_ngrams(tokens: Sequence[str], config: EbookSearchConfig) -> Counter[str]:
+252 -386
View File
@@ -6,12 +6,10 @@ import logging
from collections import defaultdict from collections import defaultdict
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from sqlalchemy import and_, delete, func, or_, select from sqlalchemy import and_, delete, or_, select, union
from python.ebook_search.protected_phrases.config import get_ignored_phrases from python.ebook_search.protected_phrases.config import get_ignored_phrases
from python.ebook_search.protected_phrases.models import ( from python.ebook_search.protected_phrases.models import (
ChunkPhraseHit,
HydratedPhraseMatch,
PhraseLookup, PhraseLookup,
PhraseMatch, PhraseMatch,
) )
@@ -29,11 +27,107 @@ if TYPE_CHECKING:
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from python.ebook_search.config import EbookSearchConfig from python.ebook_search.config import EbookSearchConfig
from python.ebook_search.protected_phrases.text_normalization import NormalizedToken
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
async def detect_protected_phrases_for_query(
session: AsyncSession,
query_text: str,
config: EbookSearchConfig,
) -> list[PhraseMatch]:
"""Find query phrases with indexed exact matches on canonical and alias norms.
Args:
session (AsyncSession): Active database session.
query_text (str): User query text to detect phrases in.
config (EbookSearchConfig): Runtime phrase-tuning settings.
Returns:
list[PhraseMatch]: Metadata-backed, overlap-resolved phrase matches for the query.
"""
tokens_ = tokenize_with_offsets(query_text)
windows_by_norm: defaultdict[str, list[tuple[int, int]]] = defaultdict(list)
token_texts = [token.text for token in tokens_]
max_tokens = max(config.phrase_max_tokens, config.phrase_max_entity_tokens)
for phrase_norm, start, end in generate_query_ngrams(
token_texts,
min_n=config.phrase_min_tokens,
max_n=max_tokens,
):
windows_by_norm[phrase_norm].append((start, end))
if not windows_by_norm:
return []
query_norms = tuple(windows_by_norm)
matched_norms = union(
select(
EbookProtectedPhrase.id.label("phrase_id"),
EbookProtectedPhrase.phrase_norm.label("matched_norm"),
).where(EbookProtectedPhrase.phrase_norm.in_(query_norms)),
select(
EbookPhraseAlias.phrase_id.label("phrase_id"),
EbookPhraseAlias.alias_norm.label("matched_norm"),
).where(EbookPhraseAlias.alias_norm.in_(query_norms)),
).subquery()
statement = select(EbookProtectedPhrase, matched_norms.c.matched_norm).join(
matched_norms,
matched_norms.c.phrase_id == EbookProtectedPhrase.id,
)
matches: list[PhraseMatch] = []
for phrase, matched_norm in await session.execute(statement):
for start, end in windows_by_norm[matched_norm]:
matches.append(
PhraseMatch(
phrase_id=phrase.id,
matched_norm=matched_norm,
phrase_text=phrase.phrase_text,
phrase_norm=phrase.phrase_norm,
canonical_id=phrase.canonical_id,
phrase_type=phrase.phrase_type,
token_count=end - start,
confidence=phrase.confidence,
importance=phrase.importance,
allow_nested=phrase.allow_nested,
suppress_children=phrase.suppress_children,
start_token=start,
end_token=end,
start_char=tokens_[start].start_char,
end_char=tokens_[end - 1].end_char,
book_id=phrase.book_id,
series_id=phrase.series_id,
)
)
return resolve_overlaps(matches)
async def index_chunk_phrase_mentions_for_book(
session: AsyncSession,
book_id: int,
config: EbookSearchConfig,
) -> int:
"""Rebuild chunk phrase mentions for all chunks in one book.
Args:
session (AsyncSession): Active database session.
book_id (int): Book whose chunk mentions are rebuilt.
config (EbookSearchConfig): Runtime phrase-tuning settings.
Returns:
int: Total number of chunk phrase mentions indexed for the book.
"""
lookup = await load_phrase_lookup(session, config, book_id=book_id)
await session.execute(delete(EbookChunkPhraseMention).where(EbookChunkPhraseMention.book_id == book_id))
chunks = await session.scalars(select(EbookChunk).where(EbookChunk.source_id == book_id).order_by(EbookChunk.id))
count = 0
for chunk in chunks:
count += index_chunk_phrase_mentions(session, chunk, lookup=lookup)
await session.flush()
logger.info(f"ebook_chunk_phrase_mentions_indexed {book_id=} {count=}")
return count
async def load_phrase_lookup( async def load_phrase_lookup(
session: AsyncSession, session: AsyncSession,
config: EbookSearchConfig, config: EbookSearchConfig,
@@ -52,40 +146,29 @@ async def load_phrase_lookup(
Returns: Returns:
PhraseLookup: Normalized phrase and alias maps with the token-window bounds to test. PhraseLookup: Normalized phrase and alias maps with the token-window bounds to test.
""" """
norm_to_ids: defaultdict[str, list[int]] = defaultdict(list) phrase_ids_by_norm: defaultdict[str, set[int]] = defaultdict(set)
alias_to_ids: defaultdict[str, list[int]] = defaultdict(list) phrases_by_id: dict[int, EbookProtectedPhrase] = {}
max_tokens = config.phrase_max_tokens max_tokens = config.phrase_max_tokens
phrase_statement = select( statement = select(
EbookProtectedPhrase.id, EbookProtectedPhrase,
EbookProtectedPhrase.phrase_norm, EbookPhraseAlias.alias_norm,
EbookProtectedPhrase.token_count, ).outerjoin(EbookPhraseAlias, EbookPhraseAlias.phrase_id == EbookProtectedPhrase.id)
)
scope_filter = protected_phrase_scope_filter(book_id=book_id, series_id=series_id) scope_filter = protected_phrase_scope_filter(book_id=book_id, series_id=series_id)
if scope_filter is not None: if scope_filter is not None:
phrase_statement = phrase_statement.where(scope_filter) statement = statement.where(scope_filter)
for row in await session.execute(phrase_statement): for phrase, alias_norm in await session.execute(statement):
phrase_id = int(row.id) phrases_by_id[phrase.id] = phrase
phrase_norm = str(row.phrase_norm) phrase_ids_by_norm[phrase.phrase_norm].add(phrase.id)
norm_to_ids[phrase_norm].append(phrase_id) max_tokens = max(max_tokens, phrase.token_count)
max_tokens = max(max_tokens, int(row.token_count)) if alias_norm is not None:
phrase_ids_by_norm[alias_norm].add(phrase.id)
alias_statement = select( max_tokens = max(max_tokens, len(alias_norm.split()))
EbookPhraseAlias.alias_norm,
EbookPhraseAlias.phrase_id,
).join(EbookProtectedPhrase, EbookProtectedPhrase.id == EbookPhraseAlias.phrase_id)
if scope_filter is not None:
alias_statement = alias_statement.where(scope_filter)
for row in await session.execute(alias_statement):
alias_norm = str(row.alias_norm)
alias_to_ids[alias_norm].append(int(row.phrase_id))
max_tokens = max(max_tokens, len(alias_norm.split()))
return PhraseLookup( return PhraseLookup(
norm_to_phrase_ids={key: tuple(values) for key, values in norm_to_ids.items()}, phrase_ids_by_norm={key: tuple(sorted(values)) for key, values in phrase_ids_by_norm.items()},
alias_to_phrase_ids={key: tuple(values) for key, values in alias_to_ids.items()}, phrases_by_id=phrases_by_id,
min_tokens=config.phrase_min_tokens, min_tokens=config.phrase_min_tokens,
max_tokens=max_tokens, max_tokens=max_tokens,
) )
@@ -111,6 +194,144 @@ def protected_phrase_scope_filter(*, book_id: int | None, series_id: int | None)
return and_(*conditions) return and_(*conditions)
def is_inside(child: PhraseMatch, parent: PhraseMatch) -> bool:
"""Return whether one token span is strictly inside another.
Args:
child (PhraseMatch): Candidate nested match.
parent (PhraseMatch): Candidate enclosing match.
Returns:
bool: True when ``child`` lies within ``parent`` and is not the same span.
"""
return (
child.start_token >= parent.start_token
and child.end_token <= parent.end_token
and (child.start_token, child.end_token, child.phrase_id)
!= (parent.start_token, parent.end_token, parent.phrase_id)
)
def index_chunk_phrase_mentions(session: AsyncSession, chunk: EbookChunk, *, lookup: PhraseLookup) -> int:
"""Store protected phrase mentions for one chunk.
Args:
session (AsyncSession): Active database session.
chunk (EbookChunk): Chunk whose text is scanned for phrase mentions.
lookup (PhraseLookup): In-memory phrase and alias lookup maps.
Returns:
int: Number of phrase mentions stored for the chunk.
"""
tokens_ = tokenize_with_offsets(chunk.text)
token_texts = [token.text for token in tokens_]
raw_matches: list[PhraseMatch] = []
phrase_windows = generate_query_ngrams(token_texts, min_n=lookup.min_tokens, max_n=lookup.max_tokens)
for matched_norm, start, end in phrase_windows:
for phrase_id in lookup.phrase_ids_by_norm.get(matched_norm, ()):
phrase = lookup.phrases_by_id[phrase_id]
raw_matches.append(
PhraseMatch(
phrase_id=phrase_id,
matched_norm=matched_norm,
phrase_text=phrase.phrase_text,
phrase_norm=phrase.phrase_norm,
canonical_id=phrase.canonical_id,
phrase_type=phrase.phrase_type,
confidence=phrase.confidence,
importance=phrase.importance,
allow_nested=phrase.allow_nested,
suppress_children=phrase.suppress_children,
start_token=start,
end_token=end,
token_count=end - start,
start_char=tokens_[start].start_char,
end_char=tokens_[end - 1].end_char,
book_id=phrase.book_id,
series_id=phrase.series_id,
)
)
matches = resolve_overlaps(raw_matches)
for match in matches:
session.add(
EbookChunkPhraseMention(
chunk_id=chunk.id,
phrase_id=match.phrase_id,
book_id=match.book_id if match.book_id is not None else chunk.source_id,
series_id=match.series_id,
start_char=match.start_char if match.start_char is not None else 0,
end_char=match.end_char,
)
)
return len(matches)
def resolve_overlaps(matches: Sequence[PhraseMatch]) -> list[PhraseMatch]:
"""Resolve overlapping phrase matches without relying only on longest match.
Args:
matches (Sequence[PhraseMatch]): Metadata-backed matches that may overlap.
Returns:
list[PhraseMatch]: The kept, non-suppressed matches.
"""
sorted_matches = sorted(
matches,
key=lambda match: (match.start_token, -match.token_count, -match.importance, -match.confidence),
)
kept: list[PhraseMatch] = []
for candidate in sorted_matches:
if any(should_suppress(candidate, existing) for existing in kept):
continue
kept.append(candidate)
return kept
def should_suppress(candidate: PhraseMatch, kept: PhraseMatch) -> bool:
"""Return whether an already-kept match should suppress a candidate.
Args:
candidate (PhraseMatch): Match being considered for keeping.
kept (PhraseMatch): Match already kept that may suppress the candidate.
Returns:
bool: True when the candidate should be dropped in favor of the kept match.
"""
if not overlaps(candidate, kept):
return False
if candidate.canonical_id == kept.canonical_id:
return rank_match(kept) >= rank_match(candidate)
if is_inside(candidate, kept) and kept.suppress_children and not candidate.allow_nested:
return True
return not candidate.allow_nested and rank_match(kept) > rank_match(candidate)
def overlaps(first: PhraseMatch, second: PhraseMatch) -> bool:
"""Return whether two token spans overlap.
Args:
first (PhraseMatch): First match to compare.
second (PhraseMatch): Second match to compare.
Returns:
bool: True when the two token spans share at least one token position.
"""
return not (first.end_token <= second.start_token or first.start_token >= second.end_token)
def rank_match(match: PhraseMatch) -> tuple[float, float, int]:
"""Rank phrase matches by importance, confidence, then token count.
Args:
match (PhraseMatch): Match to build a sort key for.
Returns:
tuple[float, float, int]: A comparable key of importance, confidence, and token count.
"""
return (match.importance, match.confidence, match.token_count)
def generate_query_ngrams( def generate_query_ngrams(
tokens_: Sequence[str], tokens_: Sequence[str],
min_n: int, min_n: int,
@@ -134,358 +355,3 @@ def generate_query_ngrams(
if phrase_norm in get_ignored_phrases(): if phrase_norm in get_ignored_phrases():
continue continue
yield phrase_norm, start, end yield phrase_norm, start, end
def detect_phrase_candidates(query_text: str, lookup: PhraseLookup) -> list[PhraseMatch]:
"""Detect protected phrase windows in a user query using RAM hash lookups.
Args:
query_text (str): User query text to scan.
lookup (PhraseLookup): In-memory phrase and alias lookup maps.
Returns:
list[PhraseMatch]: Unhydrated phrase matches found in the query.
"""
return detect_phrase_candidates_from_tokens(tokenize_with_offsets(query_text), lookup)
def detect_phrase_candidates_in_text(text: str, lookup: PhraseLookup) -> list[PhraseMatch]:
"""Detect protected phrase windows in arbitrary text with character offsets.
Args:
text (str): Arbitrary text, such as a chunk, to scan.
lookup (PhraseLookup): In-memory phrase and alias lookup maps.
Returns:
list[PhraseMatch]: Unhydrated phrase matches found in the text.
"""
return detect_phrase_candidates_from_tokens(tokenize_with_offsets(text), lookup)
def detect_phrase_candidates_from_tokens(tokens_: Sequence[NormalizedToken], lookup: PhraseLookup) -> list[PhraseMatch]:
"""Detect protected phrase windows from already-normalized tokens.
Args:
tokens_ (Sequence[NormalizedToken]): Normalized tokens with character offsets.
lookup (PhraseLookup): In-memory phrase and alias lookup maps.
Returns:
list[PhraseMatch]: Deduplicated unhydrated phrase matches with token and character spans.
"""
matches: list[PhraseMatch] = []
seen: set[tuple[int | None, str, int, int]] = set()
token_texts = [token.text for token in tokens_]
for phrase_norm, start, end in generate_query_ngrams(token_texts, min_n=lookup.min_tokens, max_n=lookup.max_tokens):
phrase_ids = lookup.norm_to_phrase_ids.get(phrase_norm, ())
alias_ids = lookup.alias_to_phrase_ids.get(phrase_norm, ())
for phrase_id in (*phrase_ids, *alias_ids):
key = (phrase_id, phrase_norm, start, end)
if key in seen:
continue
seen.add(key)
matches.append(
PhraseMatch(
phrase_norm=phrase_norm,
phrase_id=phrase_id,
start_token=start,
end_token=end,
token_count=end - start,
start_char=tokens_[start].start_char,
end_char=tokens_[end - 1].end_char,
)
)
return matches
async def hydrate_matches(session: AsyncSession, matches: Sequence[PhraseMatch]) -> list[HydratedPhraseMatch]:
"""Fetch protected phrase metadata for raw phrase matches.
Args:
session (AsyncSession): Active database session.
matches (Sequence[PhraseMatch]): Unhydrated matches to enrich.
Returns:
list[HydratedPhraseMatch]: Matches with protected-phrase metadata attached.
"""
if not matches:
return []
phrase_ids = sorted({match.phrase_id for match in matches if match.phrase_id is not None})
if not phrase_ids:
return []
rows = {
row.id: row
for row in await session.scalars(select(EbookProtectedPhrase).where(EbookProtectedPhrase.id.in_(phrase_ids)))
}
hydrated: list[HydratedPhraseMatch] = []
for match in matches:
if match.phrase_id is None:
continue
phrase = rows.get(match.phrase_id)
if phrase is None:
continue
hydrated.append(
HydratedPhraseMatch(
phrase_id=phrase.id,
matched_norm=match.phrase_norm,
phrase_text=phrase.phrase_text,
phrase_norm=phrase.phrase_norm,
canonical_id=phrase.canonical_id,
phrase_type=phrase.phrase_type,
token_count=match.token_count,
confidence=phrase.confidence,
importance=phrase.importance,
allow_nested=phrase.allow_nested,
suppress_children=phrase.suppress_children,
start_token=match.start_token,
end_token=match.end_token,
start_char=match.start_char,
end_char=match.end_char,
book_id=phrase.book_id,
series_id=phrase.series_id,
)
)
return hydrated
def overlaps(first: HydratedPhraseMatch, second: HydratedPhraseMatch) -> bool:
"""Return whether two token spans overlap.
Args:
first (HydratedPhraseMatch): First match to compare.
second (HydratedPhraseMatch): Second match to compare.
Returns:
bool: True when the two token spans share at least one token position.
"""
return not (first.end_token <= second.start_token or first.start_token >= second.end_token)
def is_inside(child: HydratedPhraseMatch, parent: HydratedPhraseMatch) -> bool:
"""Return whether one token span is strictly inside another.
Args:
child (HydratedPhraseMatch): Candidate nested match.
parent (HydratedPhraseMatch): Candidate enclosing match.
Returns:
bool: True when ``child`` lies within ``parent`` and is not the same span.
"""
return (
child.start_token >= parent.start_token
and child.end_token <= parent.end_token
and (child.start_token, child.end_token, child.phrase_id)
!= (parent.start_token, parent.end_token, parent.phrase_id)
)
def rank_match(match: HydratedPhraseMatch) -> tuple[float, float, int]:
"""Rank phrase matches by importance, confidence, then token count.
Args:
match (HydratedPhraseMatch): Match to build a sort key for.
Returns:
tuple[float, float, int]: A comparable key of importance, confidence, and token count.
"""
return (match.importance, match.confidence, match.token_count)
def should_suppress(candidate: HydratedPhraseMatch, kept: HydratedPhraseMatch) -> bool:
"""Return whether an already-kept match should suppress a candidate.
Args:
candidate (HydratedPhraseMatch): Match being considered for keeping.
kept (HydratedPhraseMatch): Match already kept that may suppress the candidate.
Returns:
bool: True when the candidate should be dropped in favor of the kept match.
"""
if not overlaps(candidate, kept):
return False
if candidate.canonical_id == kept.canonical_id:
return rank_match(kept) >= rank_match(candidate)
if is_inside(candidate, kept) and kept.suppress_children and not candidate.allow_nested:
return True
return not candidate.allow_nested and rank_match(kept) > rank_match(candidate)
def resolve_overlaps(matches: Sequence[HydratedPhraseMatch]) -> list[HydratedPhraseMatch]:
"""Resolve overlapping phrase matches without relying only on longest match.
Args:
matches (Sequence[HydratedPhraseMatch]): Hydrated matches that may overlap.
Returns:
list[HydratedPhraseMatch]: The kept, non-suppressed matches.
"""
sorted_matches = sorted(
matches,
key=lambda match: (match.start_token, -match.token_count, -match.importance, -match.confidence),
)
kept: list[HydratedPhraseMatch] = []
for candidate in sorted_matches:
if any(should_suppress(candidate, existing) for existing in kept):
continue
kept.append(candidate)
return kept
async def detect_protected_phrases_for_query(
session: AsyncSession,
query_text: str,
config: EbookSearchConfig,
*,
lookup: PhraseLookup | None = None,
book_id: int | None = None,
series_id: int | None = None,
) -> list[HydratedPhraseMatch]:
"""Run the full online protected-phrase query-detection pipeline.
Args:
session (AsyncSession): Active database session.
query_text (str): User query text to detect phrases in.
config (EbookSearchConfig): Runtime phrase-tuning settings.
lookup (PhraseLookup | None): Optional preloaded lookup; loaded on demand when ``None``.
book_id (int | None): Optional book scope for lookup loading.
series_id (int | None): Optional series scope for lookup loading.
Returns:
list[HydratedPhraseMatch]: Hydrated, overlap-resolved phrase matches for the query.
"""
active_lookup = (
lookup
if lookup is not None
else await load_phrase_lookup(session, config, book_id=book_id, series_id=series_id)
)
return resolve_overlaps(await hydrate_matches(session, detect_phrase_candidates(query_text, active_lookup)))
async def index_chunk_phrase_mentions_for_book(
session: AsyncSession,
book_id: int,
config: EbookSearchConfig,
*,
series_id: int | None = None,
lookup: PhraseLookup | None = None,
) -> int:
"""Rebuild chunk phrase mentions for all chunks in one book.
Args:
session (AsyncSession): Active database session.
book_id (int): Book whose chunk mentions are rebuilt.
config (EbookSearchConfig): Runtime phrase-tuning settings.
series_id (int | None): Optional series scope for lookup loading.
lookup (PhraseLookup | None): Optional preloaded lookup; loaded on demand when ``None``.
Returns:
int: Total number of chunk phrase mentions indexed for the book.
"""
active_lookup = (
lookup
if lookup is not None
else await load_phrase_lookup(session, config, book_id=book_id, series_id=series_id)
)
await session.execute(delete(EbookChunkPhraseMention).where(EbookChunkPhraseMention.book_id == book_id))
chunks = await session.scalars(select(EbookChunk).where(EbookChunk.source_id == book_id).order_by(EbookChunk.id))
count = 0
for chunk in chunks:
count += await index_chunk_phrase_mentions(session, chunk, lookup=active_lookup)
await session.flush()
logger.info(f"ebook_chunk_phrase_mentions_indexed {book_id=} {count=}")
return count
async def index_chunk_phrase_mentions(session: AsyncSession, chunk: EbookChunk, *, lookup: PhraseLookup) -> int:
"""Store protected phrase mentions for one chunk.
Args:
session (AsyncSession): Active database session.
chunk (EbookChunk): Chunk whose text is scanned for phrase mentions.
lookup (PhraseLookup): In-memory phrase and alias lookup maps.
Returns:
int: Number of phrase mentions stored for the chunk.
"""
raw_matches = detect_phrase_candidates_in_text(chunk.text, lookup)
hydrated = resolve_overlaps(await hydrate_matches(session, raw_matches))
for match in hydrated:
session.add(
EbookChunkPhraseMention(
chunk_id=chunk.id,
phrase_id=match.phrase_id,
book_id=match.book_id if match.book_id is not None else chunk.source_id,
series_id=match.series_id,
start_char=match.start_char if match.start_char is not None else 0,
end_char=match.end_char,
)
)
return len(hydrated)
async def phrase_hits_for_chunks(
session: AsyncSession,
*,
chunk_ids: Sequence[int],
phrase_ids: Sequence[int],
) -> dict[int, tuple[ChunkPhraseHit, ...]]:
"""Return matched protected phrases with mention counts by chunk id using indexed chunk mentions.
Args:
session (AsyncSession): Active database session.
chunk_ids (Sequence[int]): Chunk ids to look up mentions for.
phrase_ids (Sequence[int]): Protected phrase ids to restrict the results to.
Returns:
dict[int, tuple[ChunkPhraseHit, ...]]: Phrase hits per chunk id, ordered by mention count.
"""
if not chunk_ids or not phrase_ids:
return {}
mention_count = func.count(EbookChunkPhraseMention.phrase_id).label("mention_count")
statement = (
select(
EbookChunkPhraseMention.chunk_id,
EbookProtectedPhrase.id.label("phrase_id"),
EbookProtectedPhrase.phrase_text,
mention_count,
)
.join(EbookProtectedPhrase, EbookProtectedPhrase.id == EbookChunkPhraseMention.phrase_id)
.where(
EbookChunkPhraseMention.chunk_id.in_(chunk_ids),
EbookChunkPhraseMention.phrase_id.in_(phrase_ids),
)
.group_by(EbookChunkPhraseMention.chunk_id, EbookProtectedPhrase.id, EbookProtectedPhrase.phrase_text)
.order_by(EbookChunkPhraseMention.chunk_id, mention_count.desc(), EbookProtectedPhrase.phrase_text)
)
hits: defaultdict[int, list[ChunkPhraseHit]] = defaultdict(list)
for row in await session.execute(statement):
hits[int(row.chunk_id)].append(
ChunkPhraseHit(
phrase_id=int(row.phrase_id),
phrase_text=str(row.phrase_text),
mention_count=int(row.mention_count),
)
)
return {chunk_id: tuple(chunk_hits) for chunk_id, chunk_hits in hits.items()}
async def phrase_hit_counts_for_chunks(
session: AsyncSession,
*,
chunk_ids: Sequence[int],
phrase_ids: Sequence[int],
) -> dict[int, int]:
"""Return phrase-hit counts by chunk id using indexed chunk mentions.
Args:
session (AsyncSession): Active database session.
chunk_ids (Sequence[int]): Chunk ids to count mentions for.
phrase_ids (Sequence[int]): Protected phrase ids to restrict the counts to.
Returns:
dict[int, int]: Total mention count per chunk id.
"""
hits = await phrase_hits_for_chunks(session, chunk_ids=chunk_ids, phrase_ids=phrase_ids)
return {chunk_id: sum(hit.mention_count for hit in chunk_hits) for chunk_id, chunk_hits in hits.items()}
@@ -8,6 +8,8 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Mapping from collections.abc import Mapping
from python.orm.richie import EbookProtectedPhrase
@dataclass(slots=True) @dataclass(slots=True)
class PhraseCandidate: class PhraseCandidate:
@@ -71,47 +73,24 @@ class LLMJudgment:
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class PhraseLookup: class PhraseLookup:
"""In-memory lookup maps used for constant-time phrase-window checks. """In-memory phrase metadata used for constant-time text-window checks.
Attributes: Attributes:
norm_to_phrase_ids (Mapping[str, tuple[int, ...]]): Normalized phrase to protected phrase ids. phrase_ids_by_norm (Mapping[str, tuple[int, ...]]): Canonical and alias norms to phrase ids.
alias_to_phrase_ids (Mapping[str, tuple[int, ...]]): Normalized alias to protected phrase ids. phrases_by_id (Mapping[int, EbookProtectedPhrase]): Protected phrase metadata by id.
min_tokens (int): Smallest token-window size to test. min_tokens (int): Smallest token-window size to test.
max_tokens (int): Largest token-window size to test. max_tokens (int): Largest token-window size to test.
""" """
norm_to_phrase_ids: Mapping[str, tuple[int, ...]] phrase_ids_by_norm: Mapping[str, tuple[int, ...]]
alias_to_phrase_ids: Mapping[str, tuple[int, ...]] phrases_by_id: Mapping[int, EbookProtectedPhrase]
min_tokens: int min_tokens: int
max_tokens: int max_tokens: int
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class PhraseMatch: class PhraseMatch:
"""An unhydrated query or chunk phrase match. """A detected phrase match with protected-phrase metadata attached.
Attributes:
phrase_norm (str): Normalized text of the matched window.
start_token (int): Index of the first matched token.
end_token (int): Index one past the last matched token.
token_count (int): Number of tokens in the match.
phrase_id (int | None): Matched protected phrase id when known.
start_char (int | None): Start character offset in the source text.
end_char (int | None): End character offset in the source text.
"""
phrase_norm: str
start_token: int
end_token: int
token_count: int
phrase_id: int | None = None
start_char: int | None = None
end_char: int | None = None
@dataclass(frozen=True, slots=True)
class HydratedPhraseMatch:
"""A phrase match with protected-phrase metadata attached.
Attributes: Attributes:
phrase_id (int): Protected phrase id. phrase_id (int): Protected phrase id.
@@ -152,21 +131,6 @@ class HydratedPhraseMatch:
series_id: int | None = None series_id: int | None = None
@dataclass(frozen=True, slots=True)
class ChunkPhraseHit:
"""One protected phrase with its mention count inside one retrieved chunk.
Attributes:
phrase_id (int): Protected phrase id.
phrase_text (str): Display text of the protected phrase.
mention_count (int): Indexed mentions of the phrase in the chunk.
"""
phrase_id: int
phrase_text: str
mention_count: int
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class PhraseCandidateGenerationResult: class PhraseCandidateGenerationResult:
"""Summary of candidate phrase extraction for indexed books. """Summary of candidate phrase extraction for indexed books.
+23 -31
View File
@@ -19,6 +19,7 @@ from python.ebook_search.bm25_corpus import (
load_bm25_corpus, load_bm25_corpus,
score_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.embeddings import MODEL_DIMENSIONS, embed_query, get_embedding_table
from python.ebook_search.protected_phrases.matching import ( from python.ebook_search.protected_phrases.matching import (
detect_protected_phrases_for_query, detect_protected_phrases_for_query,
@@ -40,7 +41,7 @@ if TYPE_CHECKING:
from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine
from python.ebook_search.config import EbookSearchConfig from python.ebook_search.config import EbookSearchConfig
from python.ebook_search.protected_phrases.models import HydratedPhraseMatch from python.ebook_search.protected_phrases.models import PhraseMatch
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -74,7 +75,7 @@ class SearchResponse:
results: list[SearchResult] results: list[SearchResult]
rank_label: str rank_label: str
timings: tuple[RuntimeStep, ...] = () timings: tuple[RuntimeStep, ...] = ()
phrase_matches: tuple[HydratedPhraseMatch, ...] = () phrase_matches: tuple[PhraseMatch, ...] = ()
@property @property
def total_runtime_ms(self) -> float: def total_runtime_ms(self) -> float:
@@ -88,7 +89,7 @@ class RetrievalResponse:
vector_results: list[SearchResult] vector_results: list[SearchResult]
lexical_results: list[SearchResult] lexical_results: list[SearchResult]
phrase_matches: list[HydratedPhraseMatch] phrase_matches: list[PhraseMatch]
timings: tuple[RuntimeStep, ...] timings: tuple[RuntimeStep, ...]
@@ -151,18 +152,17 @@ async def search_ebooks(
return response return response
def skip_phrase_matches() -> list[HydratedPhraseMatch]:
"""Return no protected phrase matches when phrase matching is disabled."""
logger.info("ebook_protected_phrase_detection_skipped")
return []
async def query_phrase_matches( async def query_phrase_matches(
engine: AsyncEngine, engine: AsyncEngine,
query: str, query: str,
config: EbookSearchConfig, config: EbookSearchConfig,
) -> list[HydratedPhraseMatch]: *,
phrase_matching: bool,
) -> list[PhraseMatch]:
"""Detect protected phrases in a query without making search fail when phrase tables are unavailable.""" """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: try:
async with AsyncSession(engine) as session: async with AsyncSession(engine) as session:
return await detect_protected_phrases_for_query(session, query, config) return await detect_protected_phrases_for_query(session, query, config)
@@ -180,7 +180,7 @@ def skip_phrase_mention_boosts(candidates: list[SearchResult]) -> list[SearchRes
async def apply_phrase_mention_boosts( async def apply_phrase_mention_boosts(
engine: AsyncEngine, engine: AsyncEngine,
candidates: list[SearchResult], candidates: list[SearchResult],
phrase_matches: Sequence[HydratedPhraseMatch], phrase_matches: Sequence[PhraseMatch],
phrase_hit_boost: float, phrase_hit_boost: float,
) -> list[SearchResult]: ) -> list[SearchResult]:
"""Boost retrieved chunks that have indexed mentions for detected protected phrases.""" """Boost retrieved chunks that have indexed mentions for detected protected phrases."""
@@ -242,23 +242,21 @@ async def parallel_retrieval(
BM25 scoring is pure CPU work over the cached corpus, so it runs in a worker thread 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 instead of on the event loop. Protected phrase detection only depends on the query, so
it joins the gather as a third task when phrase matching is enabled. it joins the gather as a third task and returns immediately when phrase matching is disabled.
""" """
phrase_task = ( phrase_timing_name = "Protected phrase detection" if phrase_matching else "Protected phrase detection skipped"
asyncio.create_task( (
async_timed_result("Protected phrase detection", query_phrase_matches(engine, query, config)) (vector_results, vector_timing),
) (lexical_results, lexical_timing),
if phrase_matching (phrase_matches, phrase_timing),
else None ) = await asyncio.gather(
)
(vector_results, vector_timing), (lexical_results, lexical_timing) = await asyncio.gather(
async_timed_result("Embedding + vector search", vector_candidates(engine, client, query, config)), 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("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),
),
) )
if phrase_task is not None:
phrase_matches, phrase_timing = await phrase_task
else:
phrase_matches, phrase_timing = timed_result("Protected phrase detection skipped", skip_phrase_matches)
logger.info( logger.info(
f"ebook_parallel_retrieval_complete vector_candidates={len(vector_results)} " f"ebook_parallel_retrieval_complete vector_candidates={len(vector_results)} "
@@ -334,13 +332,7 @@ async def vector_candidates(
score = (literal(1.0) - distance).label("score") score = (literal(1.0) - distance).label("score")
statement = ( statement = (
select( select(
EbookChunk.id.label("chunk_id"), *CHUNK_RECORD_COLUMNS,
EbookChunk.text.label("text"),
EbookSource.id.label("source_id"),
EbookSource.title.label("source_title"),
EbookSource.author.label("source_author"),
EbookChapter.title.label("chapter_title"),
EbookChunk.page_label.label("page_label"),
score, score,
) )
.select_from(embedding_table) .select_from(embedding_table)
+8 -7
View File
@@ -25,7 +25,7 @@ from python.ebook_search.bm25_corpus import (
score_bm25_corpus, score_bm25_corpus,
write_bm25_corpus, write_bm25_corpus,
) )
from python.ebook_search.config import EbookSearchConfig, RerankConfig, load_config, normalize_embedding_model from python.ebook_search.config import EbookSearchConfig, RerankConfig, load_config
from python.ebook_search.embeddings import MODEL_DIMENSIONS, ensure_embedding_models from python.ebook_search.embeddings import MODEL_DIMENSIONS, ensure_embedding_models
from python.ebook_search.ingest import chunk_text, find_existing_source from python.ebook_search.ingest import chunk_text, find_existing_source
from python.ebook_search.search import ( from python.ebook_search.search import (
@@ -452,24 +452,25 @@ def test_1024_embedding_table_has_cosine_hnsw_index() -> None:
def test_embedding_model_aliases_normalize_to_provider_names(mocker: MockerFixture) -> None: def test_embedding_model_aliases_normalize_to_provider_names(mocker: MockerFixture) -> None:
mocker.patch.dict(environ, {}, clear=False) mocker.patch.dict(environ, {}, clear=False)
environ.pop("EBOOK_SEARCH_EMBEDDING_MODEL", None)
assert normalize_embedding_model() == "qwen3-embedding-0.6b" assert load_config().embedding_model == "qwen3-embedding-0.6b"
environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding-0.6b" environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding-0.6b"
assert normalize_embedding_model() == "qwen3-embedding-0.6b" assert load_config().embedding_model == "qwen3-embedding-0.6b"
environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "Qwen3-Embedding-0.6B" environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "Qwen3-Embedding-0.6B"
assert normalize_embedding_model() == "qwen3-embedding-0.6b" assert load_config().embedding_model == "qwen3-embedding-0.6b"
environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "Qwen/Qwen3-Embedding-4B" environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "Qwen/Qwen3-Embedding-4B"
assert normalize_embedding_model() == "qwen3-embedding-4b" assert load_config().embedding_model == "qwen3-embedding-4b"
environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding:8b" environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding:8b"
assert normalize_embedding_model() == "qwen3-embedding-8b" assert load_config().embedding_model == "qwen3-embedding-8b"
environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding-8b" environ["EBOOK_SEARCH_EMBEDDING_MODEL"] = "qwen3-embedding-8b"
assert normalize_embedding_model() == "qwen3-embedding-8b" assert load_config().embedding_model == "qwen3-embedding-8b"
def test_answer_generation_is_enabled_by_default(mocker: MockerFixture) -> None: def test_answer_generation_is_enabled_by_default(mocker: MockerFixture) -> None:
+39
View File
@@ -10,6 +10,7 @@ import pytest
from python.ebook_search.answer import answer_query from python.ebook_search.answer import answer_query
from python.ebook_search.config import EbookSearchConfig, RerankConfig from python.ebook_search.config import EbookSearchConfig, RerankConfig
from python.ebook_search.embeddings import embed_texts from python.ebook_search.embeddings import embed_texts
from python.ebook_search.llm_interface import check_chat_endpoint, check_embedding_endpoint
from python.ebook_search.search import SearchResult from python.ebook_search.search import SearchResult
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -23,6 +24,44 @@ def make_async_client(mocker: MockerFixture, fake_post) -> httpx.AsyncClient:
return client return client
async def test_model_endpoint_checks_share_http_probe(mocker: MockerFixture) -> None:
client = mocker.MagicMock(spec=httpx.AsyncClient)
response = mocker.MagicMock(spec=httpx.Response)
client.get = mocker.AsyncMock(return_value=response)
config = EbookSearchConfig(
rerank=RerankConfig(enabled=False),
embedding_base_url="https://embedding.example/v1/",
vllm_base_url="https://chat.example/v1/",
vllm_api_key="secret",
)
assert await check_embedding_endpoint(client, config, timeout_seconds=2.0)
assert await check_chat_endpoint(client, config, timeout_seconds=3.0)
assert client.get.await_args_list == [
mocker.call("https://embedding.example/v1/models", headers={}, timeout=2.0),
mocker.call(
"https://chat.example/v1/models",
headers={"Authorization": "Bearer secret"},
timeout=3.0,
),
]
assert response.raise_for_status.call_count == 2
async def test_model_endpoint_checks_report_http_failures(mocker: MockerFixture) -> None:
client = mocker.MagicMock(spec=httpx.AsyncClient)
client.get = mocker.AsyncMock(
side_effect=[
httpx.ConnectError("embedding offline"),
httpx.ConnectError("chat offline"),
]
)
config = EbookSearchConfig(rerank=RerankConfig(enabled=False))
assert not await check_embedding_endpoint(client, config)
assert not await check_chat_endpoint(client, config)
async def test_answer_query_uses_httpx_chat_completions(mocker: MockerFixture) -> None: async def test_answer_query_uses_httpx_chat_completions(mocker: MockerFixture) -> None:
captured: dict[str, object] = {} captured: dict[str, object] = {}
+89 -23
View File
@@ -7,7 +7,7 @@ from datetime import UTC, datetime
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import pytest import pytest
from sqlalchemy import select from sqlalchemy import event, select
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, create_async_engine from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, create_async_engine
from python.ebook_search.config import EbookSearchConfig from python.ebook_search.config import EbookSearchConfig
@@ -33,15 +33,12 @@ from python.ebook_search.protected_phrases.matching import (
detect_protected_phrases_for_query, detect_protected_phrases_for_query,
index_chunk_phrase_mentions, index_chunk_phrase_mentions,
load_phrase_lookup, load_phrase_lookup,
phrase_hit_counts_for_chunks,
phrase_hits_for_chunks,
resolve_overlaps, resolve_overlaps,
) )
from python.ebook_search.protected_phrases.models import ( from python.ebook_search.protected_phrases.models import (
ChunkPhraseHit,
HydratedPhraseMatch,
LLMJudgment, LLMJudgment,
PhraseCandidate, PhraseCandidate,
PhraseMatch,
) )
from python.ebook_search.protected_phrases.store import book_ids_pending_first_judgment, corpus_phrase_stats from python.ebook_search.protected_phrases.store import book_ids_pending_first_judgment, corpus_phrase_stats
from python.ebook_search.protected_phrases.text_normalization import normalize_text, tokenize from python.ebook_search.protected_phrases.text_normalization import normalize_text, tokenize
@@ -187,26 +184,89 @@ def test_score_candidate_weights_metadata_source(config: EbookSearchConfig) -> N
assert score_candidate(metadata_only, config) == 2.0 + 0.5 assert score_candidate(metadata_only, config) == 2.0 + 0.5
async def test_detect_protected_phrases_hydrates_alias_matches( async def test_load_phrase_lookup_loads_phrases_and_aliases_with_one_query(
engine: AsyncEngine,
session: AsyncSession,
) -> None:
"""One joined query should load phrases with zero, one, or many aliases."""
source = await add_source(session)
aliased_phrase = await add_phrase(session, source.id, phrase_text="lock in", phrase_norm="lock in")
plain_phrase = await add_phrase(session, source.id, phrase_text="mage king", phrase_norm="mage king")
session.add_all(
[
EbookPhraseAlias(phrase_id=aliased_phrase.id, alias_text="locked in", alias_norm="locked in"),
EbookPhraseAlias(
phrase_id=aliased_phrase.id,
alias_text="locked completely in",
alias_norm="locked completely in",
),
]
)
await session.commit()
statements: list[str] = []
def record_statement(*args: object) -> None:
statements.append(str(args[2]))
event.listen(engine.sync_engine, "before_cursor_execute", record_statement)
try:
lookup = await load_phrase_lookup(session, EbookSearchConfig(phrase_max_tokens=2), book_id=source.id)
finally:
event.remove(engine.sync_engine, "before_cursor_execute", record_statement)
assert lookup.phrase_ids_by_norm == {
"lock in": (aliased_phrase.id,),
"locked completely in": (aliased_phrase.id,),
"locked in": (aliased_phrase.id,),
"mage king": (plain_phrase.id,),
}
assert lookup.phrases_by_id == {aliased_phrase.id: aliased_phrase, plain_phrase.id: plain_phrase}
assert lookup.max_tokens == 3
assert len(statements) == 1
async def test_detect_protected_phrases_queries_canonical_and_alias_matches_once(
engine: AsyncEngine,
session: AsyncSession, session: AsyncSession,
config: EbookSearchConfig, config: EbookSearchConfig,
) -> None: ) -> None:
"""Query detection should use RAM aliases and hydrate phrase metadata from the DB.""" """Query detection should hydrate repeated canonical and alias matches with one statement."""
source = await add_source(session) source = await add_source(session)
phrase = await add_phrase(session, source.id, phrase_text="lock in", phrase_norm="lock in") phrase = await add_phrase(session, source.id, phrase_text="lock in", phrase_norm="lock in")
await add_phrase(session, source.id, phrase_text="mage king", phrase_norm="mage king")
session.add(EbookPhraseAlias(phrase_id=phrase.id, alias_text="locked in", alias_norm="locked in")) session.add(EbookPhraseAlias(phrase_id=phrase.id, alias_text="locked in", alias_norm="locked in"))
await session.commit() await session.commit()
matches = await detect_protected_phrases_for_query(session, "what is locked-in", config) statements: list[str] = []
assert [(match.phrase_text, match.canonical_id, match.phrase_type) for match in matches] == [ def record_statement(*args: object) -> None:
("lock in", "condition:lock_in", "fictional_condition") statements.append(str(args[2]))
event.listen(engine.sync_engine, "before_cursor_execute", record_statement)
try:
matches = await detect_protected_phrases_for_query(
session,
"lock-in then locked-in and lock-in",
config,
)
empty_matches = await detect_protected_phrases_for_query(session, "", config)
finally:
event.remove(engine.sync_engine, "before_cursor_execute", record_statement)
assert [(match.phrase_text, match.matched_norm, match.start_token, match.end_token) for match in matches] == [
("lock in", "lock in", 0, 2),
("lock in", "locked in", 3, 5),
("lock in", "lock in", 6, 8),
] ]
assert empty_matches == []
assert len(statements) == 1
assert " UNION " in statements[0]
def test_resolve_overlaps_keeps_independent_nested_phrases() -> None: def test_resolve_overlaps_keeps_independent_nested_phrases() -> None:
"""Overlap resolution should keep useful nested concepts when metadata permits it.""" """Overlap resolution should keep useful nested concepts when metadata permits it."""
child = hydrated_match( child = phrase_match(
phrase_id=1, phrase_id=1,
phrase_text="mage king", phrase_text="mage king",
canonical_id="title:mage_king", canonical_id="title:mage_king",
@@ -214,7 +274,7 @@ def test_resolve_overlaps_keeps_independent_nested_phrases() -> None:
end_token=5, end_token=5,
allow_nested=True, allow_nested=True,
) )
parent = hydrated_match( parent = phrase_match(
phrase_id=2, phrase_id=2,
phrase_text="mage king of mars", phrase_text="mage king of mars",
canonical_id="entity:mage_king_of_mars", canonical_id="entity:mage_king_of_mars",
@@ -231,7 +291,7 @@ def test_resolve_overlaps_keeps_independent_nested_phrases() -> None:
def test_resolve_overlaps_suppresses_weaker_same_canonical_match() -> None: def test_resolve_overlaps_suppresses_weaker_same_canonical_match() -> None:
"""Same-canonical overlaps should keep the stronger evidence.""" """Same-canonical overlaps should keep the stronger evidence."""
weak = hydrated_match( weak = phrase_match(
phrase_id=1, phrase_id=1,
phrase_text="lock", phrase_text="lock",
canonical_id="condition:lock_in", canonical_id="condition:lock_in",
@@ -239,7 +299,7 @@ def test_resolve_overlaps_suppresses_weaker_same_canonical_match() -> None:
end_token=3, end_token=3,
importance=0.2, importance=0.2,
) )
strong = hydrated_match( strong = phrase_match(
phrase_id=2, phrase_id=2,
phrase_text="lock in", phrase_text="lock in",
canonical_id="condition:lock_in", canonical_id="condition:lock_in",
@@ -274,19 +334,25 @@ async def test_index_chunk_phrase_mentions_uses_normalized_window_lookup(
await session.commit() await session.commit()
lookup = await load_phrase_lookup(session, config, book_id=source.id) lookup = await load_phrase_lookup(session, config, book_id=source.id)
count = await index_chunk_phrase_mentions(session, chunk, lookup=lookup) statements: list[str] = []
def record_statement(*args: object) -> None:
statements.append(str(args[2]))
event.listen(session.bind.sync_engine, "before_cursor_execute", record_statement)
try:
count = index_chunk_phrase_mentions(session, chunk, lookup=lookup)
finally:
event.remove(session.bind.sync_engine, "before_cursor_execute", record_statement)
await session.commit() await session.commit()
assert statements == []
mention = await session.scalar(select(EbookChunkPhraseMention)) mention = await session.scalar(select(EbookChunkPhraseMention))
assert count == 1 assert count == 1
assert mention is not None assert mention is not None
assert mention.chunk_id == chunk.id assert mention.chunk_id == chunk.id
assert mention.phrase_id == phrase.id assert mention.phrase_id == phrase.id
assert chunk.text[mention.start_char : mention.end_char] == "lock-in" assert chunk.text[mention.start_char : mention.end_char] == "lock-in"
assert await phrase_hit_counts_for_chunks(session, chunk_ids=[chunk.id], phrase_ids=[phrase.id]) == {chunk.id: 1}
assert await phrase_hits_for_chunks(session, chunk_ids=[chunk.id], phrase_ids=[phrase.id]) == {
chunk.id: (ChunkPhraseHit(phrase_id=phrase.id, phrase_text="lock in", mention_count=1),)
}
@pytest.mark.usefixtures("worker_pool") @pytest.mark.usefixtures("worker_pool")
@@ -1060,7 +1126,7 @@ async def add_phrase(
return phrase return phrase
def hydrated_match( def phrase_match(
*, *,
phrase_id: int, phrase_id: int,
phrase_text: str, phrase_text: str,
@@ -1070,9 +1136,9 @@ def hydrated_match(
importance: float = 0.8, importance: float = 0.8,
allow_nested: bool = False, allow_nested: bool = False,
suppress_children: bool = True, suppress_children: bool = True,
) -> HydratedPhraseMatch: ) -> PhraseMatch:
"""Build a hydrated match for overlap tests.""" """Build a metadata-backed phrase match for overlap tests."""
return HydratedPhraseMatch( return PhraseMatch(
phrase_id=phrase_id, phrase_id=phrase_id,
matched_norm=phrase_text, matched_norm=phrase_text,
phrase_text=phrase_text, phrase_text=phrase_text,
+4 -3
View File
@@ -62,8 +62,9 @@ async def test_search_ebooks_runs_phrase_detection_in_parallel_with_retrieval(mo
assert await asyncio.to_thread(phrase_started.wait, 2) assert await asyncio.to_thread(phrase_started.wait, 2)
return [SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)] return [SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)]
async def fake_query_phrase_matches(_engine, _query, _config): async def fake_query_phrase_matches(_engine, _query, _config, *, phrase_matching):
"""Record that phrase detection started and return no matches.""" """Record that phrase detection started and return no matches."""
assert phrase_matching is True
phrase_started.set() phrase_started.set()
return [] return []
@@ -90,7 +91,7 @@ async def test_search_ebooks_skips_phrase_matching_when_disabled(mocker: MockerF
return_value=[SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)], return_value=[SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)],
) )
mocker.patch("python.ebook_search.search.bm25_candidates", return_value=[]) mocker.patch("python.ebook_search.search.bm25_candidates", return_value=[])
detect_mock = mocker.patch("python.ebook_search.search.query_phrase_matches") detect_mock = mocker.patch("python.ebook_search.search.detect_protected_phrases_for_query")
boost_mock = mocker.patch("python.ebook_search.search.apply_phrase_mention_boosts") boost_mock = mocker.patch("python.ebook_search.search.apply_phrase_mention_boosts")
config = EbookSearchConfig(rerank=RerankConfig(enabled=False)) config = EbookSearchConfig(rerank=RerankConfig(enabled=False))
@@ -115,7 +116,7 @@ async def test_search_ebooks_ignores_phrase_matching_when_config_disabled(mocker
return_value=[SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)], return_value=[SearchResult(chunk_id=1, text="vector", source_title="Vector", vector_score=0.9)],
) )
mocker.patch("python.ebook_search.search.bm25_candidates", return_value=[]) mocker.patch("python.ebook_search.search.bm25_candidates", return_value=[])
detect_mock = mocker.patch("python.ebook_search.search.query_phrase_matches") detect_mock = mocker.patch("python.ebook_search.search.detect_protected_phrases_for_query")
boost_mock = mocker.patch("python.ebook_search.search.apply_phrase_mention_boosts") boost_mock = mocker.patch("python.ebook_search.search.apply_phrase_mention_boosts")
config = EbookSearchConfig(rerank=RerankConfig(enabled=False), phrase_matching_enabled=False) config = EbookSearchConfig(rerank=RerankConfig(enabled=False), phrase_matching_enabled=False)