"""Hybrid retrieval: pgvector cosine + Postgres full-text, fused via RRF. The bge model is loaded ONCE at import (module global), never per request. """ from __future__ import annotations from sentence_transformers import SentenceTransformer from sqlalchemy import text from app.db import SessionLocal MODEL_NAME = "BAAI/bge-small-en-v1.5" TOP_K = 5 CANDIDATES = 20 # per-arm candidate pool before RRF RRF_K = 60 # RRF damping constant POOL = 25 # fused rows fetched before diversification MAX_PER_CASE = 2 # cap chunks shown from any single judgment # Loaded once at import time. Reused for every query. _model = SentenceTransformer(MODEL_NAME) # Single-CTE RRF: rank vector arm + keyword arm, fuse by reciprocal rank. _HYBRID_SQL = text( """ WITH vec AS ( SELECT id, row_number() OVER (ORDER BY embedding <=> CAST(:qvec AS vector)) AS r FROM chunks ORDER BY embedding <=> CAST(:qvec AS vector) LIMIT :cand ), kw AS ( SELECT id, row_number() OVER ( ORDER BY ts_rank_cd(ts, plainto_tsquery('english', :q)) DESC) AS r FROM chunks WHERE ts @@ plainto_tsquery('english', :q) LIMIT :cand ), fused AS ( SELECT COALESCE(vec.id, kw.id) AS id, (COALESCE(1.0/(:rrfk + vec.r), 0) + COALESCE(1.0/(:rrfk + kw.r), 0)) AS rrf FROM vec FULL OUTER JOIN kw USING (id) ORDER BY rrf DESC LIMIT :topk ) SELECT c.id, c.case_name, c.content, fused.rrf FROM fused JOIN chunks c USING (id) ORDER BY fused.rrf DESC; """ ) def embed_query(query: str) -> list[float]: """Embed one query string. Normalized to match ingest (cosine).""" vec = _model.encode([query], normalize_embeddings=True, convert_to_numpy=True)[0] return vec.tolist() def _diversify(rows, top_k: int, max_per_case: int) -> list: """Cap chunks per case so the user sees several judgments, not one. Two-pass over RRF-ordered rows: first take up to `max_per_case` per case, then if short, backfill with leftover chunks (still RRF order). """ kept, overflow = [], [] seen: dict[str, int] = {} for r in rows: if seen.get(r.case_name, 0) < max_per_case: kept.append(r) seen[r.case_name] = seen.get(r.case_name, 0) + 1 else: overflow.append(r) if len(kept) >= top_k: break if len(kept) < top_k: kept.extend(overflow[: top_k - len(kept)]) return kept[:top_k] def hybrid_search(query: str, top_k: int = TOP_K) -> list[dict]: """Hybrid SQL -> diversify by case -> top-k [{case_name, content, score}].""" qvec = embed_query(query) db = SessionLocal() try: rows = db.execute( _HYBRID_SQL, { "qvec": str(qvec), # pgvector accepts the "[...]" string form "q": query, "cand": CANDIDATES, "rrfk": RRF_K, "topk": POOL, # fetch a wider pool; trim after diversifying }, ).all() finally: db.close() rows = _diversify(rows, top_k, MAX_PER_CASE) return [ {"case_name": r.case_name, "content": r.content, "score": float(r.rrf)} for r in rows ] if __name__ == "__main__": # Self-test: an exact statute term that pure-vector often misses but # full-text nails. Confirms the keyword arm contributes to fusion. q = "Section 302" print(f"Query: {q!r}\n") for i, hit in enumerate(hybrid_search(q), 1): present = "Section 302" in hit["content"] print(f"{i}. rrf={hit['score']:.5f} [exact-term:{present}] {hit['case_name']}") print(f" {hit['content'][:140]!r}\n")