File size: 3,156 Bytes
6754826
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
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)