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)