pubmed-rag-bot / app.py
balade's picture
Upload app.py with huggingface_hub
15b8f0d verified
Raw
History Blame Contribute Delete
30.6 kB
"""PubMed RAG — pre-built FAISS index from hf.co/datasets/balade/pubmed-faiss-index"""
import os, sys, numpy as np, json, requests, logging, asyncio, time
from pathlib import Path
from dotenv import load_dotenv
from fastapi import FastAPI, Request
from evidence_hierarchy import get_evidence_level, grade_confidence
os.environ["TOKENIZERS_PARALLELISM"] = "false"
env_path = Path(__file__).parent / ".env"
if env_path.exists():
load_dotenv(env_path)
TOKEN = os.environ.get("TELEGRAM_BOT_TOKEN", "")
DEEPSEEK_KEY = os.environ.get("DEEPSEEK_API_KEY", "")
app = FastAPI()
rag = None
rag_error = None
INDEX_REPO = "balade/pubmed-faiss-index"
# Safe drug combos (standard of care — not interactions)
SAFE_COMBOS = {
("amlodipine", "lisinopril"): "Standard first-line antihypertensive combination (ACE inhibitor + CCB). Safe and guideline-recommended.",
("amlodipine", "enalapril"): "Standard antihypertensive combination. Safe.",
("lisinopril", "hydrochlorothiazide"): "Standard ACE inhibitor + thiazide combination. Safe.",
("metformin", "sitagliptin"): "Standard diabetes combination. Safe.",
("metformin", "glibenclamide"): "Standard combination. Monitor hypoglycemia risk.",
("metformin", "empagliflozin"): "Guide-recommended combination. Safe.",
("atorvastatin", "amlodipine"): "Common cardiovascular combination. Safe within dose limits.",
("aspirin", "atorvastatin"): "Standard secondary prevention. Safe.",
}
def check_safe_combo(q):
ql = q.lower()
for (da, db), msg in SAFE_COMBOS.items():
if da in ql and db in ql:
return msg
return ""
# Conversation memory per chat_id (last 3 turns)
chat_history = {}
def load_rag():
global rag, rag_error
if rag is not None:
return rag
try:
from huggingface_hub import snapshot_download
from sentence_transformers import SentenceTransformer
import faiss
print("Downloading FAISS index from HF dataset...")
data_dir = Path("/tmp/rag_data")
data_dir.mkdir(exist_ok=True)
snapshot_download(
repo_id=INDEX_REPO,
repo_type="dataset",
local_dir=str(data_dir),
local_dir_use_symlinks=False,
)
print("Download complete.")
print("Downloading FDA drug interactions...")
try:
from huggingface_hub import hf_hub_download
import shutil
for fname, label in [
("fda_drug_interactions.json", "label text"),
("prb_drug_products.json", "product registry"),
("prb_faers_compact.json", "FAERS stats"),
]:
try:
path = hf_hub_download("balade/chatbot-assets", fname, repo_type="dataset")
shutil.copy2(path, data_dir / fname)
except Exception as e:
print(f" {fname} ({label}) failed: {e}")
print("FDA assets downloaded.")
except Exception as e:
print(f"FDA download failed: {e}")
print("Loading embedding model (CPU)...")
model = SentenceTransformer("BAAI/bge-small-en-v1.5", device="cpu")
from sentence_transformers import CrossEncoder
print("Loading cross-encoder reranker...")
reranker = CrossEncoder("cross-encoder/ms-marco-MiniLM-L-6-v2", device="cpu")
print("Reranker loaded")
print("Loading FAISS indexes...")
NPROBE = 50
index_main = faiss.read_index(str(data_dir / "faiss.index"))
index_main.nprobe = NPROBE
ids_main = np.load(str(data_dir / "metadata_ids.npy"), mmap_mode="r")
texts_main = np.load(str(data_dir / "metadata_texts.npy"), mmap_mode="r")
try:
dois_main = np.load(str(data_dir / "metadata_dois.npy"), mmap_mode="r")
journals_main = np.load(str(data_dir / "metadata_journals.npy"), mmap_mode="r")
years_main = np.load(str(data_dir / "metadata_years.npy"), mmap_mode="r")
print(f"Main: {index_main.ntotal:,} vectors, metadata: {len(dois_main):,} entries")
except:
dois_main = journals_main = years_main = None
print(f"Main: {index_main.ntotal:,} vectors")
index_recent = None
ids_recent = texts_recent = None
dois_recent = journals_recent = years_recent = None
recent_path = data_dir / "faiss_recent.index"
if recent_path.exists():
index_recent = faiss.read_index(str(recent_path))
index_recent.nprobe = NPROBE
ids_recent = np.load(str(data_dir / "metadata_recent_ids.npy"), mmap_mode="r")
texts_recent = np.load(str(data_dir / "metadata_recent_texts.npy"), mmap_mode="r")
try:
dois_recent = np.load(str(data_dir / "metadata_recent_dois.npy"), mmap_mode="r")
journals_recent = np.load(str(data_dir / "metadata_recent_journals.npy"), mmap_mode="r")
years_recent = np.load(str(data_dir / "metadata_recent_years.npy"), mmap_mode="r")
except:
pass
print(f"Recent: {index_recent.ntotal:,} vectors")
else:
print("No recent index found")
# Load drug interactions
interactions = {}
di_path = data_dir / "drug_interactions.json"
if di_path.exists():
try:
for entry in json.loads(di_path.read_text()):
key = "_".join(sorted([d.lower() for d in entry["drugs"]]))
interactions[key] = entry
print(f"Drug interactions: {len(interactions)} pairs loaded")
except:
print("Failed to load drug interactions")
# Load FDA drug interactions (full label section 7 text)
fda_di = []
fda_path = data_dir / "fda_drug_interactions.json"
if fda_path.exists():
try:
payload = json.loads(fda_path.read_text())
fda_di = payload.get("entries", [])
print(f"FDA drug interactions: {len(fda_di):,} entries loaded")
except Exception as e:
print(f"Failed to load FDA interactions: {e}")
# Load Drug@FDA product registry (brand↔generic, formulations)
drug_products = {}
dp_path = data_dir / "prb_drug_products.json"
if dp_path.exists():
try:
dp_data = json.loads(dp_path.read_text())
drug_products = dp_data.get("entries", [])
print(f"Drug@FDA products: {len(drug_products):,} entries loaded")
except Exception as e:
print(f"Failed to load drug products: {e}")
# Load FAERS adverse event stats
faers_stats = {}
fa_path = data_dir / "prb_faers_compact.json"
if fa_path.exists():
try:
faers_stats = json.loads(fa_path.read_text())
n_drugs = len(faers_stats.get("per_drug", {}))
print(f"FAERS stats: {n_drugs} drugs loaded")
except Exception as e:
print(f"Failed to load FAERS stats: {e}")
# Load food database
food_db = {}
fd_path = data_dir / "food_db.json"
if fd_path.exists():
try:
food_db = json.loads(fd_path.read_text())
print(f"Food DB loaded: {len(food_db.get('usda', {}))} foods")
except:
print("Failed to load food DB")
# Load food-drug interactions
food_di = []
fdi_path = data_dir / "food_drug_interactions.json"
if fdi_path.exists():
try:
food_di = json.loads(fdi_path.read_text())
print(f"Food-drug interactions: {len(food_di)} entries loaded")
except:
print("Failed to load food-drug interactions")
# Load Epicure substitutes
epicure = {}
ep_path = data_dir / "epicure_substitutes.json"
if ep_path.exists():
try:
epicure = json.loads(ep_path.read_text())
print(f"Epicure substitutes: {len(epicure)} ingredients")
except:
print("Failed to load Epicure substitutes")
except Exception as e:
rag_error = f"{type(e).__name__}: {e}"
print(f"RAG LOAD FAILED: {rag_error}", flush=True)
import traceback; traceback.print_exc()
return
# Token quota — daily limit per user
quota = {}
def search_clinicaltrials(query):
try:
r = requests.get("https://clinicaltrials.gov/api/query/study_fields", params={
"expr": query, "fields": "NCTId,BriefTitle,Condition,OverallStatus,Phase",
"fmt": "json", "max_rnk": 5,
}, timeout=10)
studies = r.json().get("StudyFieldsResponse", {}).get("StudyFields", [])
results = []
for s in studies:
results.append({
"id": f"NCT:{s['NCTId'][0]}",
"text": f"{s['BriefTitle'][0]} | Condition: {', '.join(s.get('Condition', [''])[:3])} | Phase: {s.get('Phase', ['N/A'])[0]} | Status: {s.get('OverallStatus', ['N/A'])[0]}",
"trial": True,
})
return results
except:
return []
def search_pmc(query):
try:
r = requests.get("https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esearch.fcgi", params={
"db": "pmc", "term": query, "retmax": 5, "retmode": "json",
}, timeout=10)
pmids = r.json().get("esearchresult", {}).get("idlist", [])
if not pmids:
return []
r = requests.get("https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esummary.fcgi", params={
"db": "pmc", "id": ",".join(pmids), "retmode": "json",
}, timeout=10)
data = r.json().get("result", {})
results = []
for pid in pmids:
item = data.get(pid, {})
title = item.get("title", "")
source = item.get("source", "")
pubdate = item.get("pubdate", "")[:4]
results.append({
"id": f"PMC:{pid}",
"text": f"{title} | {source}, {pubdate}",
"pmc": True,
})
return results
except:
return []
def check_interactions(q):
words = q.lower().split()
results = []
for key, entry in interactions.items():
drugs = [d.lower() for d in entry["drugs"]]
if all(any(d in word or word in d for word in words) for d in drugs):
results.append(f"⚠️ {entry['drugs'][0]} + {entry['drugs'][1]}: {entry['severity'].upper()}{entry['effect']}")
return results
def search_fda_interactions(q):
ql = q.lower()
q_words = [w for w in ql.split() if len(w) > 2]
results = []
for entry in fda_di:
generic = (entry.get("generic_name") or "").lower()
brand = (entry.get("brand_name") or "").lower()
haystack = f"{generic} {brand}"
if not any(w in haystack for w in q_words):
continue
di_text = entry.get("drug_interactions", "")
if di_text and len(di_text) > 50:
label = brand or generic
results.append(f"[FDA] {label}{di_text[:1500]}")
if len(results) >= 3:
break
return results
def search_drug_products(q):
ql = q.lower()
q_words = [w for w in ql.split() if len(w) > 2]
results = []
for entry in drug_products:
products = entry.get("products", [])
for prod in products:
brand = (prod.get("brand_name") or "").lower()
ings = " ".join(i.get("name", "") for i in prod.get("active_ingredients", []))
if not any(w in f"{brand} {ings}" for w in q_words):
continue
ing_list = "; ".join(f"{i['name']} {i.get('strength','')}".strip() for i in prod.get("active_ingredients", []))
results.append(f"[Drug@FDA] {prod['brand_name']}{prod.get('dosage_form','')}, {prod.get('route','')} | Active: {ing_list}")
if len(results) >= 5:
break
if len(results) >= 5:
break
return results
def search_faers_stats(q):
ql = q.lower()
per_drug = faers_stats.get("per_drug", {})
results = []
for drugname, stats in per_drug.items():
if drugname not in ql and ql not in drugname:
# check if any query word matches
if not any(w in drugname for w in ql.split() if len(w) > 3):
continue
top_r = ", ".join(r["r"] for r in stats.get("top_reactions", [])[:5])
top_o = ", ".join(o["o"] for o in stats.get("top_outcomes", [])[:3])
parts = [f"[FAERS] {drugname}: {stats['reports']:,} reports ({stats['serious_pct']}% serious)"]
if top_r:
parts.append(f" Most common: {top_r}")
results.append("\n".join(parts))
if len(results) >= 3:
break
return results
def search_live_pubmed(q):
"""Live PubMed API — for drug pairs not in our JSON"""
all_drugs = set()
for d1, d2 in SAFE_COMBOS:
all_drugs.add(d1); all_drugs.add(d2)
for entry in interactions.values():
for d in entry.get("drugs", []):
all_drugs.add(d.lower())
ql = q.lower()
found = [d for d in all_drugs if d in ql]
if len(found) < 2:
return ""
found = found[:2]
try:
r = requests.get("https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esearch.fcgi",
params={"db": "pubmed", "term": f"({' AND '.join(found)}) AND interaction", "retmax": 3, "retmode": "json"}, timeout=10)
pmids = r.json().get("esearchresult", {}).get("idlist", [])
if not pmids:
return ""
r = requests.get("https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esummary.fcgi",
params={"db": "pubmed", "id": ",".join(pmids), "retmode": "json"}, timeout=10)
data = r.json().get("result", {})
snippets = []
for pid in pmids:
item = data.get(pid, {})
title = item.get("title", "")
source = item.get("source", "")
pubdate = item.get("pubdate", "")[:4]
if title:
snippets.append(f"{' AND '.join(found)} interaction: {title} | {source}, {pubdate} [Live PubMed]")
return "\n".join(snippets) if snippets else ""
except:
return ""
def search_food(q):
results = []
ql = q.lower()
for fid, item in food_db.get("usda", {}).items():
name = item.get("name", "").lower()
if ql in name or any(w in name for w in ql.split() if len(w) > 3):
results.append(f"🍽 {item['name']}")
if "nutrients" in item:
n = item["nutrients"][:5]
results.append(f" Nutrients: {', '.join([str(x['nutrient_id']) + '=' + str(x['amount']) for x in n])}")
if len(results) >= 3:
break
return results[:3]
def search_food_interactions(q):
ql = q.lower()
results = []
for entry in food_di:
names = [entry.get("name_id", "").lower(), entry.get("name_en", "").lower(), entry.get("latin", "").lower()]
match = False
for n in names:
if not n: continue
if n in ql: match = True; break
# Partial word match — "carambola" matches "averrhoa carambola"
if any(qw in n for qw in ql.split() if len(qw) > 2):
match = True; break
if any(nw in ql for nw in n.split() if len(nw) > 2):
match = True; break
if not match:
continue
for inter in entry.get("interactions", []):
with_drugs = ", ".join(inter.get("with", []))
results.append({
"item": entry["name_id"],
"drug": with_drugs,
"effect": inter["effect"],
"severity": inter["severity"],
"evidence": inter["evidence"],
"pmids": inter.get("pmids", []),
})
if results:
break
return results[:5]
def search_substitutes(q):
ql = q.lower()
for name, subs in epicure.items():
if name in ql or any(w in name for w in ql.split() if len(w) > 3):
return f"Substitutes for {name}: {', '.join(subs[:5])}"
return ""
class RAG:
def search(self, query, k=5):
ID_WORDS = {"di", "ke", "dari", "yang", "dan", "pada", "dengan", "atau", "ini", "itu",
"adalah", "untuk", "tidak", "akan", "bisa", "apakah", "bagaimana", "cara",
"kerja", "obat", "manfaat", "efek", "samping", "penggunaan", "tentang",
"sebagai", "dalam", "ada", "saya", "anda", "kami"}
is_id = sum(1 for w in query.lower().split() if w in ID_WORDS) >= 2
SEARCH_K = 50
RERANK_TRUNCATE = 512
q_emb = model.encode([query], normalize_embeddings=True).astype(np.float32)
def search_one(idx_obj, ids_arr, texts_arr, dois_arr=None, journals_arr=None, years_arr=None):
max_idx = len(ids_arr)
if dois_arr is not None:
max_idx = min(max_idx, len(dois_arr))
scores, idxs = idx_obj.search(q_emb, SEARCH_K)
results = []
for i, s in zip(idxs[0], scores[0]):
if i < 0 or i >= max_idx:
continue
pos = int(i)
pid = ids_arr[pos].decode("utf-8", errors="replace").replace("pmid_", "PMID:")
r = {"id": pid, "text": texts_arr[pos].decode("utf-8", errors="replace"), "score": float(s)}
if dois_arr is not None: r["doi"] = dois_arr[pos].decode("utf-8", errors="replace").strip("\x00").strip()
if journals_arr is not None: r["journal"] = journals_arr[pos].decode("utf-8", errors="replace").strip("\x00").strip()
if years_arr is not None: r["year"] = years_arr[pos].decode("utf-8", errors="replace").strip("\x00")
results.append(r)
return results
# Gather top candidates from both indexes
all_results = search_one(index_main, ids_main, texts_main, dois_main, journals_main, years_main)
if index_recent is not None:
all_results.extend(search_one(index_recent, ids_recent, texts_recent, dois_recent, journals_recent, years_recent))
# Dedup by ID
seen = set()
deduped = []
for r in sorted(all_results, key=lambda x: x["score"], reverse=True):
if r["id"] not in seen:
seen.add(r["id"])
deduped.append(r)
candidates = deduped[:SEARCH_K]
# Rerank with cross-encoder (English only — Indonesian skips)
if not is_id:
pairs = [[query, r["text"][:RERANK_TRUNCATE]] for r in candidates]
scores = reranker.predict(pairs, show_progress_bar=False)
for r, s in zip(candidates, scores):
r["rerank_score"] = float(s)
candidates.sort(key=lambda x: x["rerank_score"], reverse=True)
else:
candidates.sort(key=lambda x: x["score"], reverse=True)
final = candidates[:k]
for r in final:
r["evidence"] = get_evidence_level(["Journal Article"])
return final
def answer(self, q, k=5, chat_id=None):
from drug_synonyms import expand_query
import hashlib
# Conversation memory — store raw query for search, enriched for LLM
raw_q = q
if chat_id is not None and chat_id in chat_history:
prev_q, prev_a = chat_history[chat_id]
q = f"Context from earlier: user asked '{prev_q}' and was told '{prev_a[:200]}'. Now user asks: {q}"
del chat_history[chat_id]
# Always search with original query, not enriched
search_q = raw_q # consume history (only use last turn)
# Token quota
today = time.strftime("%Y-%m-%d")
quota_key = f"{today}_{hashlib.md5(q.encode()).hexdigest()[:8]}"
if quota_key not in quota:
quota[quota_key] = 0
quota[quota_key] += 1
if quota[quota_key] > 300:
return "Daily query limit reached (300/day). Upgrade for unlimited."
# Keyword filter — skip LLM for greetings
greeting_words = {"hi", "hello", "thanks", "thank", "halo", "hai", "assalamualaikum", "makasih", "test", "ping", "nice", "good", "great", "ok", "oke", "cool", "wow", "lol", "bye", "goodbye"}
if set(q.lower().split()) & greeting_words:
return "Halo! Tanya tentang obat atau penyakit, ya? Contoh: 'metformin diabetes' atau 'efek samping ibuprofen'"
hits = self.search(expand_query(search_q), k)
trials = search_clinicaltrials(search_q)
pmc_results = search_pmc(search_q)
live_di = search_live_pubmed(search_q)
food = search_food(search_q)
subs = search_substitutes(search_q)
foodi = search_food_interactions(search_q)
fda_di_results = search_fda_interactions(search_q)
drug_prod_results = search_drug_products(search_q)
faers_results = search_faers_stats(search_q)
# Safe regimen check — use raw query
safe_regimen = ""
extra_notes = []
ql = search_q.lower()
if "amlodipine" in ql and "lisinopril" in ql:
safe_regimen = "Standard first-line antihypertensive combination (ACE inhibitor + CCB). Safe and guideline-recommended."
if "metformin" in ql and "sitagliptin" in ql:
safe_regimen = "Standard diabetes combination. Safe."
# Known dangerous drug combos
if "simvastatin" in ql and "colchicine" in ql:
extra_notes.append("WARNING: Simvastatin + colchicine increases risk of myopathy and rhabdomyolysis (both CYP3A4 substrates, colchicine inhibits P-gp). Avoid combination or monitor closely.")
if "clarithromycin" in ql and "simvastatin" in ql:
extra_notes.append("WARNING: Clarithromycin + simvastatin severely increases statin levels (CYP3A4 inhibition). Risk of rhabdomyolysis. Avoid combination.")
# Common food-drug interaction checks
if "carambola" in ql or "belimbing" in ql or "star fruit" in ql:
extra_notes.append("CARAMBOLA (STAR FRUIT): Contraindicated in renal impairment. Contains neurotoxin caramboxin and oxalic acid. Interacts with antihypertensives. Avoid if kidney problems.")
if "kunyit" in ql or "turmeric" in ql:
extra_notes.append("KUNYIT (TURMERIC): May increase bleeding risk with warfarin, clopidogrel. CYP3A4 inhibitor.")
if "jahe" in ql or "ginger" in ql:
extra_notes.append("JAHE (GINGER): High doses may increase bleeding risk with anticoagulants. Culinary amounts safe.")
if "jambu" in ql or "guava" in ql:
extra_notes.append("JAMBU BIJI (GUAVA): Leaf tea may interact with warfarin. Fruit safe.")
if "rosella" in ql or "hibiscus" in ql:
extra_notes.append("ROSELLA (HIBISCUS): May interact with ACE inhibitors and diuretics.")
interactions = check_interactions(q)
ctx_parts = []
for h in hits:
ctx_parts.append(f"[{h['id']}] {h['text'][:800]}")
for t in trials:
ctx_parts.append(f"[{t['id']}] {t['text']}")
for p in pmc_results:
ctx_parts.append(f"[{p['id']}] {p['text']}")
for f in food:
ctx_parts.append(f"Food: {f}")
for fi in foodi:
ctx_parts.append(f"Herb-Drug Interaction: {fi['item']} + {fi['drug']}: {fi['effect']} ({fi['severity'].upper()})")
if subs:
ctx_parts.append(f"Ingredient Substitutes: {subs}")
if live_di:
ctx_parts.append(f"[Live PubMed] {live_di}")
for fd in fda_di_results:
ctx_parts.append(fd)
for dp in drug_prod_results:
ctx_parts.append(dp)
for fa in faers_results:
ctx_parts.append(fa)
if not ctx_parts:
msg = "No relevant results found."
if interactions:
msg += "\n\n" + "\n".join(interactions)
return msg
ctx_lines = list(ctx_parts)
if safe_regimen:
ctx_lines.insert(0, f"KNOWN SAFE COMBINATION: {safe_regimen}")
for note in extra_notes:
ctx_lines.append(f"SAFETY ALERT: {note}")
ctx = "\n\n".join(ctx_lines)
# Evidence + confidence (only from PubMed)
is_id_query = "rerank_score" not in hits[0] if hits else True
if hits:
best_score_raw = hits[0].get("rerank_score", hits[0]["score"])
best_score = best_score_raw if best_score_raw > 0 else hits[0]["score"]
if best_score < 0.30:
prefix = "⚠️ No strong match from PubMed. "
confidence = "LOW"
else:
prefix = ""
evidence_weights = [h["evidence"]["final_weight"] for h in hits]
avg_evidence = sum(evidence_weights) / len(evidence_weights)
combined = (best_score * 0.6) + (avg_evidence * 0.4)
if is_id_query:
confidence = "MODERATE" if combined > 0.35 else "LOW"
else:
confidence = "HIGH" if combined > 0.75 else "MODERATE" if combined > 0.55 else "LOW"
else:
prefix = ""
confidence = "LOW"
prompt = f"""Answer using ONLY the context below. Cite claims as [PMID:12345678] or [NCT01234567].
If insufficient, say so. No outside knowledge.
Question: {q}
Context:
{ctx}"""
try:
r = requests.post("https://api.deepseek.com/v1/chat/completions",
headers={"Authorization": f"Bearer {DEEPSEEK_KEY}", "Content-Type": "application/json"},
json={"model": "deepseek-chat", "messages": [{"role": "user", "content": prompt}],
"max_tokens": 500, "temperature": 0.3}, timeout=30)
answer = r.json()["choices"][0]["message"]["content"]
except Exception as e:
answer = f"[LLM error: {e}]"
sources = []
is_id_query = "rerank_score" not in hits[0] if hits else True
has_relevant = any(h.get("rerank_score", 0) > 0 for h in hits)
if has_relevant or is_id_query:
for h in hits:
rs = h.get("rerank_score", h["score"])
sc = rs if rs > 0 else h["score"]
sources.append(f"{h['id']} (score: {sc:.2f})")
for t in trials:
sources.append(f"{t['id']} (trial)")
for p in pmc_results:
sources.append(f"{p['id']} (full text)")
if food:
sources.append("Food data (USDA FDC)")
for fi in foodi:
pmids_str = ", ".join(fi.get("pmids", []))
sources.append(f"Herb-Drug: {fi['item']} + {fi['drug']} ({fi['severity'].upper()}) {pmids_str}")
if fda_di_results:
sources.append("FDA drug interaction labels (OpenFDA)")
body = f"{answer}\n\nConfidence: {confidence}"
if sources:
body += "\n\nSources:\n" + "\n".join(sources)
result = prefix + body
if chat_id is not None:
chat_history[chat_id] = (q, result)
return result
rag = RAG()
return rag
@app.get("/")
async def root():
status = "alive"
index_status = "loaded" if rag is not None else ("error" if rag_error else "loading")
resp = {
"status": status,
"model": "BAAI/bge-small-en-v1.5",
"index": f"{INDEX_REPO} ({index_status})",
}
if rag_error:
resp["error"] = rag_error
return resp
@app.get("/search")
async def search(q: str = "", k: int = 5):
if not q:
return {"error": "Provide ?q=query"}
if rag is None:
return {"error": "Still loading"}
try:
return {"results": rag.search(q, k)}
except Exception as e:
return {"error": f"{type(e).__name__}: {e}"}
@app.post("/webhook")
async def webhook(request: Request):
try:
data = await request.json()
msg = data.get("message", {})
text = msg.get("text", "").strip()
chat_id = msg.get("chat", {}).get("id")
if not text or not chat_id:
return {"ok": False}
if rag is None:
return {"method": "sendMessage", "chat_id": chat_id, "text": "Loading RAG..."}
answer = rag.answer(text, chat_id=chat_id)
return {"method": "sendMessage", "chat_id": chat_id, "text": answer}
except Exception as e:
logging.error(f"Webhook: {e}", exc_info=True)
import traceback
return {"ok": False, "error": f"{type(e).__name__}: {str(e)[:200]}"}
@app.on_event("startup")
async def startup():
asyncio.create_task(asyncio.to_thread(load_rag))
if __name__ == "__main__":
import uvicorn
port = int(os.environ.get("PORT", 7860))
uvicorn.run(app, host="0.0.0.0", port=port)