import os
import re
import json
import requests
from langchain_community.llms import Ollama
from langchain_community.vectorstores import FAISS
from langchain.embeddings.base import Embeddings
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain.chains import RetrievalQA
from langchain.docstore.document import Document
from langchain.prompts import PromptTemplate
from IPython.display import display, HTML
class OllamaEmbeddings(Embeddings):
def __init__(self, model='nomic-embed-text', url='http://localhost:11434/api/embeddings'):
self.model = model
self.url = url
def embed(self, text):
headers = {'Content-Type': 'application/json'}
payload = {'model': self.model, 'prompt': text}
response = requests.post(self.url, headers=headers, data=json.dumps(payload))
if response.status_code == 200:
response_json = response.json()
if 'embedding' in response_json:
return response_json['embedding']
print("Key 'embedding' not found in response")
return None
raise Exception(f"Error {response.status_code}: {response.text}")
def embed_documents(self, texts):
return [self.embed(text) for text in texts]
def embed_query(self, text):
return self.embed(text)
llm = Ollama(
model="llama3.2",
temperature=0.1, # slight bump from 0: clinical summaries read better with a touch of fluency
)
def load_structured_txt(file_path):
with open(file_path, "r", encoding="utf-8") as f:
raw = f.read()
# Robust to "## Disease: X" or just "## X" (handles LLM formatting drift)
disease_blocks = re.split(r'(?=^## )', raw, flags=re.MULTILINE)
records = []
for block in disease_blocks:
block = block.strip()
if not block.startswith("## "):
continue
disease_name = block.splitlines()[0].replace("## Disease:", "").replace("##", "").strip()
section_splits = re.split(r'(?=### Section:)', block)
for sec in section_splits:
sec = sec.strip()
if not sec.startswith("### Section:"):
continue
lines = sec.splitlines()
section_name = lines[0].replace("### Section:", "").strip()
section_text = "\n".join(lines[1:]).strip()
records.append({
"disease": disease_name,
"section": section_name,
"text": section_text
})
return records
txt_path = "clinical_guidelines.txt"
records = load_structured_txt(txt_path)
print(f"Parsed {len(records)} disease/section blocks")
for r in records:
print(f" - {r['disease']} / {r['section']}")
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=400,
chunk_overlap=50,
separators=["\n\n", "\n", ". ", " "]
)
documents = []
for rec in records:
sub_chunks = text_splitter.split_text(rec["text"])
for chunk in sub_chunks:
documents.append(
Document(
page_content=chunk,
metadata={"disease": rec["disease"], "section": rec["section"]}
)
)
print(f"Created {len(documents)} chunks with citation metadata")
embeddings = OllamaEmbeddings()
knowledge_base = FAISS.from_documents(documents, embeddings)
test_docs = knowledge_base.similarity_search("What is the recommended treatment pathway for Hypertension and when should a specialist be invloved?", k=10)
for d in test_docs:
# print(d.metadata, "->", d.page_content[:80])
print(d.metadata)
clinical_prompt = PromptTemplate(
input_variables=["context", "question"],
template="""You are a clinical guideline assistant. Use ONLY the context below to answer.
If the context does not contain enough information, say so explicitly instead of guessing.
Context:
{context}
Question: {question}
Respond in exactly this structure:
1. Evidence-Backed Answer: <concise answer using only the context>
2. Treatment Pathway Summary: <numbered steps if applicable, otherwise "Not applicable">
3. Clinical Judgment Warning: <a short note that this is guideline-based information and a qualified clinician must confirm applicability to the specific patient>
"""
)
known_diseases = list({rec["disease"] for rec in records})
def detect_disease(question, diseases):
q_lower = question.lower()
for d in diseases:
if d.lower() in q_lower:
return d
return None
question = "What is the recommended treatment pathway for Hypertension and when should a specialist be involved?"
target_disease = detect_disease(question, known_diseases)
if target_disease:
retriever = knowledge_base.as_retriever(
search_kwargs={"k": 6, "filter": {"disease": target_disease}}
)
else:
retriever = knowledge_base.as_retriever(search_kwargs={"k": 8})
qa_chain = RetrievalQA.from_chain_type(
llm,
retriever=retriever,
chain_type_kwargs={"prompt": clinical_prompt},
return_source_documents=True
)
# question = "What is do u know about hypertension?"
response = qa_chain.invoke({"query": question})
print(response["result"])
print("\n--- Citations ---")
seen = set()
citations = []
for doc in response["source_documents"]:
key = (doc.metadata.get("disease"), doc.metadata.get("section"))
if key not in seen:
seen.add(key)
citations.append(key)
print(f"- {key[0]} -> {key[1]}")
citations_rows = "".join(
f"<tr><td><b>{d}</b></td><td>{s}</td></tr>" for d, s in citations
)
html_output = f"""
<h1>Clinical Guidelines Q&A Assistant</h1>
<h3>Question</h3>
<div>{question}</div>
<h3>Answer</h3>
<div style="white-space: pre-wrap; border:1px solid #ccc; padding:10px;">{response['result']}</div>
<h3>Citations</h3>
<table border="1" cellpadding="6" cellspacing="0">
<tr><th>Disease</th><th>Guideline Section</th></tr>
{citations_rows}
</table>
<div style="margin-top:15px; padding:10px; background:#fff3cd; border:1px solid #ffeeba;">
<b>Note:</b> This output is generated from synthetic training data and general guideline structure.
It is not a substitute for clinical judgment or current, source-verified medical guidance.
</div>
"""
display(HTML(html_output))1 views