""" Query interface for the retrieval layer. ``query()`` is the single entry point the CLI (and any caller) uses. It: 1. embeds the natural-language query, 2. runs an ANN search against the FAISS index, 3. attaches title / year / category / URL / abstract snippet to each hit, 4. applies the requested filters (year range, primary category, survey-only), 5. returns a ranked list of ``Result`` objects. If no index exists on disk yet (or building it failed), ``query()`` falls back to a **pure-keyword** search that streams the metadata and matches on title/ abstract with a TF-style score. This guarantees the CLI always returns something, even before an embedding index has been built. """ from __future__ import annotations import datetime import re from dataclasses import dataclass, field import numpy as np from data import DataStore, SURVEY_KEYWORDS # Cap the working set for the keyword fallback so it stays fast on the 3.15M-row # metadata file (term-overlap scoring doesn't need the full corpus). _KEYWORD_WORKING_SET = 200_000 @dataclass class Result: paper_id: str title: str year: int category: str url: str snippet: str score: float method: str # "semantic" | "keyword" @dataclass class QueryOptions: top_k: int = 10 year_min: int | None = None year_max: int | None = None primary_category: str | None = None survey_only: bool = False embedder = None # set by caller; optional def _snippet(abstract: str, maxlen: int = 200) -> str: text = re.sub(r"\s+", " ", abstract or "").strip() if len(text) <= maxlen: return text return text[:maxlen].rsplit(" ", 1)[0] + " …" def _normalize(text: str) -> str: return re.sub(r"[^a-z0-9]+", " ", (text or "").lower()).strip() def _keyword_score(text: str, terms: list[str]) -> float: """TF-style overlap score between a doc and the query terms.""" words = set(text.split()) score = sum(1 for t in terms if t in words) # small bonus for title matches return score def _keyword_search( query: str, store: DataStore, opts: QueryOptions, ) -> list[Result]: """ Pure-keyword fallback: stream matching rows, score by term overlap, filter, and return the top-k. Used when no ANN index is available. """ terms = [t for t in _normalize(query).split() if len(t) > 1] if not terms: return [] # Pull a working set of rows (push the cheap filters into the scan). scan = store.scan_filtered( year_min=opts.year_min, year_max=opts.year_max, primary_category=opts.primary_category, survey_only=opts.survey_only, ) rows = ( scan.select( ["paper_id", "title", "primary_category", "first_version_date", "arxiv_abs_url", "abstract"] ) .limit(_KEYWORD_WORKING_SET) .collect() .to_dicts() ) if not rows: return [] scored: list[Result] = [] for r in rows: rec = dict(r) title = str(rec["title"] or "") abstract = str(rec["abstract"] or "") body = f"{title} {abstract}" score = _keyword_score(body.lower(), terms) if score <= 0: continue fvd = rec["first_version_date"] year = fvd.year if isinstance(fvd, __import__("datetime").datetime) else (int(fvd) if fvd is not None else 0) scored.append( Result( paper_id=str(rec["paper_id"]), title=title, year=year, category=str(rec["primary_category"] or ""), url=str(rec["arxiv_abs_url"] or ""), snippet=_snippet(abstract), score=score, method="keyword", ) ) scored.sort(key=lambda x: x.score, reverse=True) return scored[: opts.top_k] def query( query_text: str, store: DataStore, ann=None, embedder=None, opts: QueryOptions | None = None, ) -> list[Result]: """ Run a search. Tries the ANN index first; falls back to keyword search. Returns a list of ranked ``Result`` objects. """ opts = opts or QueryOptions() query_text = query_text or "" # --- semantic path --- if ann is not None and embedder is not None: try: return _semantic_query(query_text, store, ann, embedder, opts) except Exception as exc: # pragma: no cover - defensive print(f"[query] ANN search failed ({exc}); falling back to keyword search.") # --- keyword fallback --- return _keyword_search(query_text, store, opts) def _semantic_query( query_text: str, store: DataStore, ann, embedder, opts: QueryOptions, ) -> list[Result]: import faiss qvec = np.ascontiguousarray(embedder.embed_one(query_text), dtype=np.float32) k = min(opts.top_k * 8, len(ann), ann.index.ntotal) k = max(k, 1) dists, idx = ann.index.search(qvec.reshape(1, -1), k) dists = dists[0] idx = idx[0] # Gather candidate hits and apply filters. candidates: list[Result] = [] for d, i in zip(dists.tolist(), idx.tolist()): if i < 0 or i >= len(ann._ids): continue pid = ann._ids[i] year = ann._years[i] cat = ann._cats[i] if opts.year_min is not None and year < opts.year_min: continue if opts.year_max is not None and year > opts.year_max: continue if opts.primary_category and cat != opts.primary_category: continue if opts.survey_only: if not any(kw in ann._abstracts[i].lower() for kw in SURVEY_KEYWORDS): continue # FAISS L2 distance -> similarity in [0,1]; closer => higher. sim = 1.0 / (1.0 + float(d)) candidates.append( Result( paper_id=pid, title=ann._titles[i], year=year, category=cat, url=ann._urls[i], snippet=_snippet(ann._abstracts[i]), score=sim, method="semantic", ) ) candidates.sort(key=lambda x: x.score, reverse=True) return candidates[: opts.top_k]