Download agent/rag.py from SUPER321/anam_code: direct link, hf CLI and curl.
- Browser
- Download file 3.16 kB
-
https://huggingface.co/spaces/SUPER321/anam_code/resolve/main/agent/rag.py
- Command line
-
hf download hf://spaces/SUPER321/anam_code/agent/rag.py
-
curl -L -o rag.py https://huggingface.co/spaces/SUPER321/anam_code/resolve/main/agent/rag.py
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) | |