| """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 |
| RRF_K = 60 |
| POOL = 25 |
| MAX_PER_CASE = 2 |
|
|
| |
| _model = SentenceTransformer(MODEL_NAME) |
|
|
| |
| _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), |
| "q": query, |
| "cand": CANDIDATES, |
| "rrfk": RRF_K, |
| "topk": POOL, |
| }, |
| ).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__": |
| |
| |
| 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") |
|
|