feat(kb/search): LLM-powered query expansion in /kb/search
Single-shot search ranks the source whose title literally contains the query terms way above everything else, so "מהי תקנה 37" returns 6 hits from the תקנה 37 circular + 2 from the law and zero from ספר הליקויים — even though the latter is the canonical medical reference for that regulation. The ask agent already worked around this by issuing several search_insurance_kb calls with different phrasings; that intelligence is now also available to the public /kb/search endpoint. Implementation: - _retrieve_rrf extracts vector + lexical retrieval (RRF fused, no rerank) into a reusable async helper. Each candidate now carries chunk_id so callers can dedup when merging multiple sub-searches. - _expand_query asks the configured Claude model (via ai-gateway) for up to 3 alternative phrasings, with a 12s timeout. On any failure the caller silently falls back to single-query mode. - search() gains an `expand: bool` parameter. When true, runs the original query plus all variants through _retrieve_rrf in parallel, merges by chunk_id keeping the best RRF score, and reranks the union against the ORIGINAL query so the final order reflects what the user asked, not the variants. - /kb/search defaults to expand=true. Cost per call: 1 extra LLM round-trip (~1-2s on Sonnet) and 3 extra Voyage embeds (~150ms each, parallel). Latency budget grows from ~700ms to ~3s, acceptable for a search the user explicitly triggered. Refs Task Master #1
This commit is contained in:
@@ -37,6 +37,10 @@ class SearchRequest(BaseModel):
|
|||||||
query: str = Field(..., min_length=1, max_length=500)
|
query: str = Field(..., min_length=1, max_length=500)
|
||||||
kind: Literal["any", "law", "regulation", "circular"] = "any"
|
kind: Literal["any", "law", "regulation", "circular"] = "any"
|
||||||
top_k: int = Field(8, ge=1, le=15)
|
top_k: int = Field(8, ge=1, le=15)
|
||||||
|
# Multi-query expansion: ask the LLM for alternative phrasings, retrieve
|
||||||
|
# for each, merge, then rerank against the original query. Recovers
|
||||||
|
# documents that don't share the user's exact wording.
|
||||||
|
expand: bool = True
|
||||||
|
|
||||||
|
|
||||||
@router.post("/search")
|
@router.post("/search")
|
||||||
@@ -47,6 +51,7 @@ async def search(body: SearchRequest, request: Request):
|
|||||||
query=body.query,
|
query=body.query,
|
||||||
kind=body.kind,
|
kind=body.kind,
|
||||||
top_k=body.top_k,
|
top_k=body.top_k,
|
||||||
|
expand=body.expand,
|
||||||
)
|
)
|
||||||
# Dates come back as datetime.date; serialize as ISO strings for the UI.
|
# Dates come back as datetime.date; serialize as ISO strings for the UI.
|
||||||
serialized = []
|
serialized = []
|
||||||
|
|||||||
+123
-20
@@ -1,12 +1,22 @@
|
|||||||
"""Hybrid search: vector + full-text, fused with Reciprocal Rank Fusion,
|
"""Hybrid search: vector + full-text, fused with Reciprocal Rank Fusion,
|
||||||
and finally reranked with a Voyage cross-encoder for precision."""
|
and finally reranked with a Voyage cross-encoder for precision.
|
||||||
|
|
||||||
|
`search()` supports an `expand=True` mode that asks the LLM for query
|
||||||
|
variants, retrieves candidates for each, merges them, and reranks the
|
||||||
|
union against the original query. Costs one extra LLM call + one Voyage
|
||||||
|
embed per variant, but recovers documents that don't share the user's
|
||||||
|
exact phrasing — e.g. ספר הליקויים when the user types "מהי תקנה 37".
|
||||||
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
|
from openai import AsyncOpenAI
|
||||||
|
|
||||||
from api.services.kb import voyage
|
from api.services.kb import voyage
|
||||||
from api.services.kb.db import get_pool
|
from api.services.kb.db import get_pool
|
||||||
|
|
||||||
@@ -17,17 +27,20 @@ _CANDIDATES_PER_SIDE = 30
|
|||||||
# Number of fused RRF candidates fed to the reranker. Higher = better
|
# Number of fused RRF candidates fed to the reranker. Higher = better
|
||||||
# recall before the final cut; capped at ~30 to keep latency reasonable.
|
# recall before the final cut; capped at ~30 to keep latency reasonable.
|
||||||
_RERANK_POOL = 30
|
_RERANK_POOL = 30
|
||||||
|
# How many alternative phrasings to ask the LLM to generate when expand=True.
|
||||||
|
_EXPANSION_VARIANTS = 3
|
||||||
|
_EXPANSION_TIMEOUT_S = 12.0
|
||||||
|
|
||||||
|
|
||||||
async def search(
|
async def _retrieve_rrf(
|
||||||
query: str,
|
query: str,
|
||||||
kind: Literal["law", "regulation", "circular", "any"] = "any",
|
kind: Literal["law", "regulation", "circular", "any"],
|
||||||
top_k: int = 8,
|
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
query = (query or "").strip()
|
"""Run vector + lexical retrieval, fuse with RRF. No rerank, no top_k cut.
|
||||||
if not query:
|
|
||||||
return []
|
|
||||||
|
|
||||||
|
Each candidate dict carries `chunk_id` (kb_chunk.id) so callers can
|
||||||
|
deduplicate when merging results from multiple queries.
|
||||||
|
"""
|
||||||
[vec] = await voyage.embed([query], input_type="query")
|
[vec] = await voyage.embed([query], input_type="query")
|
||||||
|
|
||||||
kind_filter = ""
|
kind_filter = ""
|
||||||
@@ -36,7 +49,6 @@ async def search(
|
|||||||
kind_filter = "AND s.kind = $4"
|
kind_filter = "AND s.kind = $4"
|
||||||
params.append(kind)
|
params.append(kind)
|
||||||
|
|
||||||
# Two CTEs — vector rank and lexical rank — fused via RRF.
|
|
||||||
sql = f"""
|
sql = f"""
|
||||||
WITH vec AS (
|
WITH vec AS (
|
||||||
SELECT c.id, ROW_NUMBER() OVER (ORDER BY c.embedding <=> $1::vector) AS rank
|
SELECT c.id, ROW_NUMBER() OVER (ORDER BY c.embedding <=> $1::vector) AS rank
|
||||||
@@ -67,6 +79,7 @@ async def search(
|
|||||||
) u GROUP BY id
|
) u GROUP BY id
|
||||||
)
|
)
|
||||||
SELECT
|
SELECT
|
||||||
|
c.id AS chunk_id,
|
||||||
s.id AS source_id, s.kind, s.title, s.identifier, s.source_url,
|
s.id AS source_id, s.kind, s.title, s.identifier, s.source_url,
|
||||||
s.published_at, s.effective_at, s.original_path,
|
s.published_at, s.effective_at, s.original_path,
|
||||||
c.heading_path, c.section_ref, c.content, c.page_number,
|
c.heading_path, c.section_ref, c.content, c.page_number,
|
||||||
@@ -81,33 +94,123 @@ async def search(
|
|||||||
pool = await get_pool()
|
pool = await get_pool()
|
||||||
async with pool.acquire() as conn:
|
async with pool.acquire() as conn:
|
||||||
rows = await conn.fetch(sql, *params)
|
rows = await conn.fetch(sql, *params)
|
||||||
|
return [dict(r) for r in rows]
|
||||||
|
|
||||||
candidates = [dict(r) for r in rows]
|
|
||||||
logger.info(
|
_EXPANSION_PROMPT = (
|
||||||
"[kb.search] query=%r kind=%s RRF_candidates=%d",
|
"אתה עוזר חיפוש בבסיס ידע משפטי של הביטוח הלאומי בישראל "
|
||||||
query[:80], kind, len(candidates),
|
"(החוק, התקנות, ספרי המבחנים והליקויים, וחוזרי המוסד). "
|
||||||
)
|
"המשתמש שואל שאלה. החזר עד {n} ניסוחים חלופיים שיעזרו לאחזר "
|
||||||
if not candidates:
|
"את אותו מידע ממסמכים משפטיים בעברית — תרגומים מסגנון שאלתי "
|
||||||
|
"לסגנון מסמך, מילים נרדפות (למשל 'תקנה 37' → 'בדיקה מחדש', "
|
||||||
|
"'נכות' → 'נכה'/'דרגת נכות'/'מוגבלות'), והוספת מונחים רפואיים/משפטיים "
|
||||||
|
"שלא מופיעים במפורש בשאלה. החזר ניסוח אחד לשורה, ללא מספור או הקדמה. "
|
||||||
|
"אל תחזור על הניסוח המקורי."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _expand_query(query: str) -> list[str]:
|
||||||
|
"""Ask the LLM for alternative phrasings. Empty list on failure (caller
|
||||||
|
just runs the original query)."""
|
||||||
|
base = (os.environ.get("AI_GATEWAY_URL") or "http://localhost:3000").rstrip("/")
|
||||||
|
api_key = os.environ.get("AI_GATEWAY_API_KEY", "")
|
||||||
|
if not api_key:
|
||||||
return []
|
return []
|
||||||
|
client = AsyncOpenAI(base_url=f"{base}/v1", api_key=api_key)
|
||||||
|
try:
|
||||||
|
resp = await asyncio.wait_for(
|
||||||
|
client.chat.completions.create(
|
||||||
|
model=os.environ.get("CLAUDE_MODEL", "sonnet"),
|
||||||
|
messages=[
|
||||||
|
{"role": "system", "content": _EXPANSION_PROMPT.format(n=_EXPANSION_VARIANTS)},
|
||||||
|
{"role": "user", "content": query},
|
||||||
|
],
|
||||||
|
max_tokens=256,
|
||||||
|
temperature=0.3,
|
||||||
|
),
|
||||||
|
timeout=_EXPANSION_TIMEOUT_S,
|
||||||
|
)
|
||||||
|
except (asyncio.TimeoutError, Exception) as e:
|
||||||
|
logger.warning("[kb.search] query expansion failed: %s", e)
|
||||||
|
return []
|
||||||
|
text = (resp.choices[0].message.content or "").strip()
|
||||||
|
variants: list[str] = []
|
||||||
|
for line in text.splitlines():
|
||||||
|
v = line.strip(" -•·*\t").strip()
|
||||||
|
if not v or v == query or v in variants:
|
||||||
|
continue
|
||||||
|
variants.append(v)
|
||||||
|
if len(variants) >= _EXPANSION_VARIANTS:
|
||||||
|
break
|
||||||
|
return variants
|
||||||
|
|
||||||
|
|
||||||
|
async def _rerank_or_truncate(
|
||||||
|
query: str, candidates: list[dict], top_k: int,
|
||||||
|
) -> list[dict]:
|
||||||
|
"""Cross-encoder rerank against `query`; fall back to RRF order on error."""
|
||||||
rerank_disabled = os.environ.get("KB_RERANK_DISABLED", "").lower() in ("1", "true")
|
rerank_disabled = os.environ.get("KB_RERANK_DISABLED", "").lower() in ("1", "true")
|
||||||
if rerank_disabled or len(candidates) <= top_k:
|
if rerank_disabled or len(candidates) <= top_k:
|
||||||
return candidates[:top_k]
|
return candidates[:top_k]
|
||||||
|
|
||||||
# Cross-encoder rerank for precision. On failure fall back to the RRF order
|
|
||||||
# rather than returning nothing — degraded precision is better than outage.
|
|
||||||
texts = [c["content"] for c in candidates]
|
texts = [c["content"] for c in candidates]
|
||||||
try:
|
try:
|
||||||
ranked = await voyage.rerank(query=query, documents=texts, top_k=top_k)
|
ranked = await voyage.rerank(query=query, documents=texts, top_k=top_k)
|
||||||
except voyage.VoyageError as e:
|
except voyage.VoyageError as e:
|
||||||
logger.warning("[kb.search] rerank failed, falling back to RRF: %s", e)
|
logger.warning("[kb.search] rerank failed, falling back to RRF: %s", e)
|
||||||
return candidates[:top_k]
|
return candidates[:top_k]
|
||||||
|
|
||||||
out: list[dict] = []
|
out: list[dict] = []
|
||||||
for r in ranked:
|
for r in ranked:
|
||||||
i = r["index"]
|
i = r["index"]
|
||||||
if 0 <= i < len(candidates):
|
if 0 <= i < len(candidates):
|
||||||
item = {**candidates[i], "rerank_score": r["relevance_score"]}
|
out.append({**candidates[i], "rerank_score": r["relevance_score"]})
|
||||||
out.append(item)
|
|
||||||
logger.info("[kb.search] rerank kept top %d / %d", len(out), len(candidates))
|
logger.info("[kb.search] rerank kept top %d / %d", len(out), len(candidates))
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
async def search(
|
||||||
|
query: str,
|
||||||
|
kind: Literal["law", "regulation", "circular", "any"] = "any",
|
||||||
|
top_k: int = 8,
|
||||||
|
expand: bool = False,
|
||||||
|
) -> list[dict]:
|
||||||
|
query = (query or "").strip()
|
||||||
|
if not query:
|
||||||
|
return []
|
||||||
|
|
||||||
|
if not expand:
|
||||||
|
candidates = await _retrieve_rrf(query, kind)
|
||||||
|
logger.info(
|
||||||
|
"[kb.search] query=%r kind=%s RRF_candidates=%d",
|
||||||
|
query[:80], kind, len(candidates),
|
||||||
|
)
|
||||||
|
if not candidates:
|
||||||
|
return []
|
||||||
|
return await _rerank_or_truncate(query, candidates, top_k)
|
||||||
|
|
||||||
|
# Multi-query expansion: variants run in parallel with the original.
|
||||||
|
variants = await _expand_query(query)
|
||||||
|
queries = [query] + variants
|
||||||
|
logger.info("[kb.search] expand=true variants=%d", len(variants))
|
||||||
|
|
||||||
|
results = await asyncio.gather(
|
||||||
|
*[_retrieve_rrf(q, kind) for q in queries],
|
||||||
|
return_exceptions=True,
|
||||||
|
)
|
||||||
|
merged: dict[int, dict] = {}
|
||||||
|
for r in results:
|
||||||
|
if isinstance(r, Exception):
|
||||||
|
logger.warning("[kb.search] sub-retrieval failed: %s", r)
|
||||||
|
continue
|
||||||
|
for c in r:
|
||||||
|
cid = c["chunk_id"]
|
||||||
|
if cid not in merged or c["score"] > merged[cid]["score"]:
|
||||||
|
merged[cid] = c
|
||||||
|
|
||||||
|
candidates = sorted(merged.values(), key=lambda x: -x["score"])[:_RERANK_POOL]
|
||||||
|
logger.info(
|
||||||
|
"[kb.search] merged candidates=%d (from %d queries)",
|
||||||
|
len(candidates), len(queries),
|
||||||
|
)
|
||||||
|
if not candidates:
|
||||||
|
return []
|
||||||
|
return await _rerank_or_truncate(query, candidates, top_k)
|
||||||
|
|||||||
Reference in New Issue
Block a user