# mcp/orchestrator.py import asyncio from typing import Dict, Any from mcp.arxiv import fetch_arxiv from mcp.pubmed import fetch_pubmed from mcp.nlp import extract_umls_concepts from mcp.umls_rel import fetch_relations from mcp.openfda import fetch_drug_safety from mcp.ncbi import search_gene, get_mesh_definition from mcp.disgenet import disease_to_genes from mcp.clinicaltrials import search_trials from mcp.mygene import mygene from mcp.opentargets import ot from mcp.cbio import cbio from mcp.openai_utils import ai_summarize, ai_qa from mcp.gemini import gemini_summarize, gemini_qa def _get_llm(llm: str): return (gemini_summarize, gemini_qa) if llm.lower() == "gemini" else (ai_summarize, ai_qa) async def orchestrate_search(query: str, llm: str = "openai") -> Dict[str, Any]: # 1) Parallel literature pulls arxiv_t, pubmed_t = fetch_arxiv(query), fetch_pubmed(query) papers = [] for res in await asyncio.gather(arxiv_t, pubmed_t, return_exceptions=True): if isinstance(res, list): papers.extend(res) # 2) SpaCy→UMLS concept linking blob = " ".join(p.get("summary","") for p in papers) umls = await extract_umls_concepts(blob) # 3) Fetch UMLS relations in parallel rels = await asyncio.gather( *[fetch_relations(c["cui"]) for c in umls], return_exceptions=True ) # 4) Enrich: OpenFDA, NCBI, DisGeNET, Trials, OpenTargets, cBioPortal keys = [c["name"] for c in umls] fda_tasks = [fetch_drug_safety(k) for k in keys] gene_task = search_gene(keys[0]) if keys else asyncio.sleep(0, result=[]) mesh_task = get_mesh_definition(keys[0]) if keys else asyncio.sleep(0, result="") dis_task = disease_to_genes(keys[0]) if keys else asyncio.sleep(0, result=[]) trials_task = search_trials(query) ot_task = ot.fetch(keys[0]) if keys else asyncio.sleep(0, result=[]) cbio_task = cbio.fetch_variants(keys[0]) if keys else asyncio.sleep(0, result=[]) fda, gene, mesh, dis, trials, ot_assoc, variants = await asyncio.gather( asyncio.gather(*fda_tasks, return_exceptions=True), gene_task, mesh_task, dis_task, trials_task, ot_task, cbio_task, return_exceptions=False ) # 5) AI summary summarize, _ = _get_llm(llm) try: ai_summary = await summarize(blob) except Exception: ai_summary = "LLM summary failed." return { "papers": papers, "umls": umls, "umls_relations": rels, "drug_safety": fda, "genes": [gene], "mesh_defs": [mesh], "gene_disease": dis, "clinical_trials": trials, "ot_associations": ot_assoc, "variants": variants, "ai_summary": ai_summary, "llm_used": llm.lower() } async def answer_ai_question(question: str, context: str = "", llm: str = "openai"): _, qa_fn = _get_llm(llm) try: answer = await qa_fn(question, context) except Exception: answer = "LLM follow-up failed." return {"answer": answer}