gsearch-api / llm_client.py
tanmayapna's picture
initial deploy
6fc17e7
Raw
History Blame Contribute Delete
7.38 kB
"""
GCAS Search Engine – Provider-agnostic LLM reranker
Supported providers
-------------------
"openai" – uses chat completions (GPT-4o-mini by default)
"anthropic" – uses Messages API (claude-haiku by default)
The LLM receives a batch of FAISS candidates and returns a ranked subset
with relevance scores and short explanations.
"""
from __future__ import annotations
import json
import logging
import re
from typing import Any, Dict, List, Optional
from config import settings
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Prompts
# ---------------------------------------------------------------------------
SYSTEM_PROMPT = """\
You are a precise search-result ranker for an Indian college admissions database.
You receive:
1. A user's natural-language search query.
2. A numbered list of candidate database records.
Your task:
- Identify which records genuinely match the user's intent.
- Score each matched record from 0.0 (irrelevant) to 1.0 (perfect match).
- Discard records that are clearly irrelevant.
- Return at most the requested number of results.
OUTPUT FORMAT – respond with ONLY a JSON array, no prose:
[
{"index": <integer>, "score": <0.0–1.0>, "reason": "<one short sentence>"},
...
]
"""
# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------
def rerank_with_llm(
query: str,
candidates: List[Dict[str, Any]],
top_k: int,
*,
provider: Optional[str] = None,
model: Optional[str] = None,
api_key: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""
Ask the LLM to pick and score the best *top_k* records from *candidates*.
Parameters
----------
query : original user query
candidates : list of candidate dicts (from FAISS search)
top_k : max results to keep
provider : "openai" | "anthropic" – falls back to settings
model : model name – falls back to settings
api_key : API key override
Returns
-------
List of dicts: [{"index": int, "score": float, "reason": str}, ...]
In case of any error the original embedding scores are returned as-is.
"""
provider = provider or settings.llm_provider
model = model or settings.llm_model
user_message = _build_user_message(query, candidates, top_k)
try:
if provider == "openai":
raw = _call_openai(user_message, model, api_key)
elif provider == "anthropic":
raw = _call_anthropic(user_message, model, api_key)
else:
raise ValueError(f"Unknown LLM provider: '{provider}'")
return _parse_llm_response(raw, len(candidates), top_k)
except Exception:
logger.exception(
"LLM reranking failed (provider=%s model=%s). "
"Returning raw embedding scores.",
provider, model,
)
# Graceful degradation: pass through top-k by embedding score
return [
{"index": i, "score": c.get("score", 0.0), "reason": ""}
for i, c in enumerate(candidates[:top_k])
]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _build_user_message(
query: str,
candidates: List[Dict[str, Any]],
top_k: int,
) -> str:
"""Serialise candidates into a compact numbered list for the LLM prompt."""
lines: List[str] = [
f'Query: "{query}"',
f"Return the top {top_k} most relevant records from the list below.",
"",
]
for i, c in enumerate(candidates):
lines.append(f"[{i}] Table: {c['table']}")
# Include up to 20 non-null fields per record to stay within token budget
shown = 0
for k, v in c["data"].items():
if v is not None and str(v).strip() not in ("", "nan", "NaN", "None"):
lines.append(f" {k}: {v}")
shown += 1
if shown >= 20:
break
lines.append("") # blank line between records
return "\n".join(lines)
def _call_openai(user_message: str, model: str, api_key: Optional[str]) -> str:
try:
from openai import OpenAI
except ImportError as exc:
raise ImportError("openai package required: pip install openai") from exc
key = api_key or settings.openai_api_key
if not key:
raise ValueError("OPENAI_API_KEY is not configured.")
client = OpenAI(api_key=key)
# Use json_object mode only for models that support it
extra: Dict[str, Any] = {}
if any(tag in model for tag in ("gpt-4o", "gpt-4-turbo", "gpt-3.5-turbo")):
extra["response_format"] = {"type": "json_object"}
response = client.chat.completions.create(
model=model,
messages=[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_message},
],
temperature=0,
**extra,
)
return response.choices[0].message.content or ""
def _call_anthropic(user_message: str, model: str, api_key: Optional[str]) -> str:
try:
import anthropic
except ImportError as exc:
raise ImportError("anthropic package required: pip install anthropic") from exc
key = api_key or settings.anthropic_api_key
if not key:
raise ValueError("ANTHROPIC_API_KEY is not configured.")
client = anthropic.Anthropic(api_key=key)
response = client.messages.create(
model=model,
max_tokens=2048,
system=SYSTEM_PROMPT,
messages=[{"role": "user", "content": user_message}],
)
return response.content[0].text if response.content else ""
def _parse_llm_response(
raw: str,
num_candidates: int,
top_k: int,
) -> List[Dict[str, Any]]:
"""
Extract a JSON array from the LLM output.
Handles cases where the model wraps the array in an object.
"""
raw = raw.strip()
# 1. Try direct parse
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
# 2. Extract the first [...] block via regex
match = re.search(r"\[.*\]", raw, re.DOTALL)
if not match:
raise ValueError(f"No JSON array found in LLM output: {raw[:200]}")
parsed = json.loads(match.group())
# 3. If the model returned an object wrapping the array, unwrap it
if isinstance(parsed, dict):
for v in parsed.values():
if isinstance(v, list):
parsed = v
break
else:
raise ValueError(f"Unexpected JSON object from LLM: {list(parsed.keys())}")
# 4. Validate and normalise
results: List[Dict[str, Any]] = []
for item in parsed:
idx = item.get("index")
score = float(item.get("score", 0.0))
reason = str(item.get("reason", ""))
if idx is None or not isinstance(idx, int):
continue
if not (0 <= idx < num_candidates):
continue
results.append({"index": idx, "score": max(0.0, min(1.0, score)), "reason": reason})
# Sort by LLM score, take top_k
results.sort(key=lambda x: x["score"], reverse=True)
return results[:top_k]