anam_code / agent /rag.py
Vaibhav Kathait
Initial deploy
6754826
Raw History Blame Contribute Delete
3.16 kB
import json
import math
from pathlib import Path
from openai import AsyncOpenAI
from config import OPENAI_API_KEY
DATA_DIR = Path(__file__).parent / "data"
KB_FILE = DATA_DIR / "supplements_kb.json"
EMBEDDINGS_FILE = DATA_DIR / "supplements_embeddings.json"
EMBEDDING_MODEL = "text-embedding-3-small"
# Embeddings always go to OpenAI directly — LLM_BASE_URL may point at a
# non-OpenAI provider (Groq/Ollama/etc.) that doesn't serve this model.
_embed_client: AsyncOpenAI | None = None
def _get_embed_client() -> AsyncOpenAI:
global _embed_client
if _embed_client is None:
_embed_client = AsyncOpenAI(api_key=OPENAI_API_KEY, base_url="https://api.openai.com/v1")
return _embed_client
def _load_kb() -> list[dict]:
return json.loads(KB_FILE.read_text())
def _load_embeddings() -> dict:
if EMBEDDINGS_FILE.exists():
try:
return json.loads(EMBEDDINGS_FILE.read_text())
except Exception:
pass
return {}
def _save_embeddings(embeddings: dict) -> None:
EMBEDDINGS_FILE.write_text(json.dumps(embeddings))
def _chunk_text(entry: dict) -> str:
return (
f"Marker: {', '.join(entry['marker_names'])}\n"
f"Reference range: {entry['reference_range']}\n"
f"Deficiency symptoms: {entry['deficiency_symptoms']}\n"
f"Recommendations: {entry['recommendations']}\n"
f"Interaction cautions: {entry.get('interaction_cautions', '')}\n"
f"Note: {entry['consult_note']}"
)
def _cosine_similarity(a: list[float], b: list[float]) -> float:
dot = sum(x * y for x, y in zip(a, b))
mag_a = math.sqrt(sum(x * x for x in a))
mag_b = math.sqrt(sum(y * y for y in b))
if mag_a == 0 or mag_b == 0:
return 0.0
return dot / (mag_a * mag_b)
async def _ensure_embeddings() -> dict:
"""Return {id: {chunk, embedding}}, embedding any KB entries missing from the cache."""
kb = _load_kb()
embeddings = _load_embeddings()
missing = [e for e in kb if e["id"] not in embeddings]
if missing:
client = _get_embed_client()
for entry in missing:
chunk = _chunk_text(entry)
resp = await client.embeddings.create(model=EMBEDDING_MODEL, input=chunk)
embeddings[entry["id"]] = {"chunk": chunk, "embedding": resp.data[0].embedding}
_save_embeddings(embeddings)
return embeddings
async def search_supplements(query: str, top_k: int = 3) -> str:
"""Embed `query` and return the top-k most relevant supplement KB chunks."""
embeddings = await _ensure_embeddings()
if not embeddings:
return "The supplement knowledge base is empty."
client = _get_embed_client()
resp = await client.embeddings.create(model=EMBEDDING_MODEL, input=query)
q_vec = resp.data[0].embedding
scored = [
(_cosine_similarity(q_vec, data["embedding"]), kb_id, data["chunk"])
for kb_id, data in embeddings.items()
]
scored.sort(key=lambda x: x[0], reverse=True)
blocks = [f"--- {kb_id} ---\n{chunk}" for _, kb_id, chunk in scored[:top_k]]
return "\n\n".join(blocks)