dmChatbotBackend / src /agents /agents.py
github-actions
Auto deploy from GitHub
84ae02f
Raw
History Blame Contribute Delete
21.4 kB
import os
import re
import json
import time
import asyncio
import functools
from langchain_core.messages import SystemMessage, HumanMessage, AIMessage
from src.utils.logger import setup_logger
logger = setup_logger("Agents")
def throttle_agent(func):
@functools.wraps(func)
async def wrapper(self, state, *args, **kwargs):
agent_name = getattr(self, "agent_name", self.__class__.__name__)
logger.info(f"[{agent_name}] >>> Start executing")
start_time = time.time()
try:
res = await func(self, state, *args, **kwargs)
elapsed = time.time() - start_time
logger.info(f"[{agent_name}] <<< Execution completed in {elapsed:.3f} seconds.")
if elapsed < 1.0:
delay = 1.5 - elapsed
logger.info(f"[{agent_name}] Execution was faster than 1.0s. Throttling: waiting {delay:.3f}s to reach 1.5s total time.")
await asyncio.sleep(delay)
logger.info(f"[{agent_name}] Throttling completed. Proceeding to next step.")
return res
except Exception as e:
logger.error(f"[{agent_name}] Exception during execution: {e}")
raise e
return wrapper
from src.agent_params import get_agent_params
from src.core.model_manager import model_manager
from src.core.state import AgentState
from src.core.evidence_models import ClinicalOutputWithEvidence, EvidenceCitation
from src.tools.web_tools import web_search_tool
from src.tools.dietary_tools import search_guidelines, get_nutritional_data, page_indexed_retrieval
from src.tools.patient_memory import save_patient_memory, get_patient_memory
from src.agents.role_utils import classify_role
from datetime import datetime
class BaseAgent:
def __init__(self, fallback_prompt: str, prompt_file: str = None, tools: list = None, agent_name: str = None):
self.agent_name = agent_name or self.__class__.__name__
self.fallback_prompt = fallback_prompt
self.prompt_file = prompt_file
self.tools = tools or []
self.params = get_agent_params(self.agent_name)
self.temperature = float(self.params.get("temperature", 0.0))
self.model_name = self.params.get("model_name")
self._refresh_llm()
def _refresh_llm(self):
llm = model_manager.get_llm(
temperature=self.temperature,
model_name=self.model_name,
)
if self.tools:
llm = llm.bind_tools(self.tools)
self.llm = llm
def parse_json_response(self, response_text: str):
if not response_text:
return {}
text = response_text.strip()
try:
return json.loads(text)
except json.JSONDecodeError:
match = re.search(r"\{.*\}", text, re.S)
if not match:
return {}
try:
return json.loads(match.group(0))
except json.JSONDecodeError:
return {}
@property
def system_prompt(self) -> str:
"""Dynamically load prompt from file if available, otherwise use fallback."""
if self.prompt_file:
current_dir = os.path.dirname(os.path.abspath(__file__))
prompt_path = os.path.abspath(os.path.join(current_dir, "..", "prompts", self.prompt_file))
try:
if os.path.exists(prompt_path):
with open(prompt_path, "r", encoding="utf-8") as handle:
prompt = handle.read().strip()
else:
prompt = self.fallback_prompt
except Exception:
prompt = self.fallback_prompt
else:
prompt = self.fallback_prompt
skip_confidence = self.agent_name in {"ResponseValidator", "SafetyCheck"}
confidence_instruction = (
"\n\nAt the end of your response, include a confidence score from 0.0 to 1.0 "
"in the format: Confidence: 0.8"
)
if not skip_confidence and "confidence" not in prompt.lower():
prompt = f"{prompt}{confidence_instruction}"
return prompt
@throttle_agent
async def run(self, state: AgentState, config=None):
"""Standard run method for graph nodes."""
messages = [SystemMessage(content=self.system_prompt)] + state["messages"]
logger.info(f"--- Sending {len(messages)} messages to LLM ({self.agent_name}) ---")
start_time = time.time()
try:
response = await self.llm.ainvoke(messages, config=config)
end_time = time.time()
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
confidence = None
if hasattr(response, "content") and isinstance(response.content, str):
match = re.search(r"Confidence(?:\s+Score)?:\s*([0-9.]+)", response.content, re.IGNORECASE)
if match:
try:
confidence = float(match.group(1))
except ValueError:
pass
metrics = {
"agent": self.agent_name,
"tokens": tokens,
"time": round(end_time - start_time, 3),
"confidence": confidence
}
return {"messages": [response], "metrics": [metrics]}
except Exception as e:
logger.error(f"Error in {self.agent_name}.run: {e}")
raise
class RoleClassifier(BaseAgent):
def __init__(self):
fallback_prompt = """You are a medical triage assistant.
Classify the user input into one of five roles: 'patient', 'caregiver', 'clinician', 'researcher', or 'dietary'. Return only the name."""
super().__init__(fallback_prompt, "RoleClassifier.txt")
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
messages = [SystemMessage(content=self.system_prompt)] + state["messages"]
logger.info(f"--- RoleClassifier: Sending {len(messages)} messages to LLM ---")
user_message = state["messages"][-1].content.lower() if state["messages"] else ""
role = classify_role(user_message)
start_time = time.time()
try:
if role is None:
response = await self.llm.ainvoke(messages)
end_time = time.time()
raw = response.content.lower()
if not raw.strip():
role = "patient"
else:
roles = ["patient", "caregiver", "clinician", "researcher", "dietary"]
role = next((r for r in roles if r in raw), "patient")
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
duration = round(end_time - start_time, 3)
else:
end_time = time.time()
tokens = 0
duration = round(end_time - start_time, 3)
metrics = {
"agent": "RoleClassifier",
"tokens": tokens,
"time": duration
}
return {"user_role": role, "metrics": [metrics]}
except Exception as e:
logger.error(f"Error in RoleClassifier.run: {e}")
raise
class PatientLLM(BaseAgent):
def __init__(self):
fallback_prompt = """You are a compassionate medical assistant for patients.
Provide helpful, empathetic, and medically sound advice."""
super().__init__(fallback_prompt, "PatientLLM.txt", tools=[web_search_tool, save_patient_memory, get_patient_memory])
class CaregiverLLM(BaseAgent):
def __init__(self):
fallback_prompt = """You are a supportive caregiver assistant for a diabetes management platform.
Help caregivers interpret symptoms, monitor treatment adherence, and know when to escalate to urgent care.
Frame advice as practical proxy guidance for a patient while remaining clear and compassionate."""
super().__init__(fallback_prompt, "CaregiverLLM.txt", tools=[web_search_tool, save_patient_memory, get_patient_memory])
class ResponseValidator(BaseAgent):
def __init__(self):
fallback_prompt = """You are a medical response validator.
Check if the last response is medically accurate and follows guidelines. Return JSON only."""
super().__init__(fallback_prompt, "ResponseValidator.txt")
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
last_message = state["messages"][-1].content
start_time = time.time()
response = await self.llm.ainvoke([
SystemMessage(content=self.system_prompt),
HumanMessage(content=f"Verify this response: {last_message}")
])
end_time = time.time()
parsed = self.parse_json_response(response.content)
decision = parsed.get("decision", "invalid").lower()
is_valid = decision == "valid"
if not parsed:
lower_response = response.content.lower()
if re.search(r"\binvalid\b", lower_response):
is_valid = False
elif re.search(r"\bvalid\b", lower_response):
is_valid = True
else:
is_valid = False
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
metrics = {
"agent": "ResponseValidator",
"tokens": tokens,
"time": round(end_time - start_time, 3)
}
return {"is_valid": is_valid, "metrics": [metrics]}
class SafetyCheck(BaseAgent):
def __init__(self):
fallback_prompt = """You are a medical safety officer.
Check if the response contains any dangerous advice or misinformation. Return JSON only."""
super().__init__(fallback_prompt, "SafetyCheck.txt")
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
last_message = state["messages"][-1].content
start_time = time.time()
response = await self.llm.ainvoke([
SystemMessage(content=self.system_prompt),
HumanMessage(content=f"Safety check on this: {last_message}")
])
end_time = time.time()
parsed = self.parse_json_response(response.content)
decision = parsed.get("decision", "unsafe").lower()
is_safe = decision == "safe"
if not parsed:
lower_response = response.content.lower()
if re.search(r"\bunsafe\b", lower_response):
is_safe = False
elif re.search(r"\bsafe\b", lower_response):
is_safe = True
else:
is_safe = False
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
metrics = {
"agent": "SafetyCheck",
"tokens": tokens,
"time": round(end_time - start_time, 3)
}
return {"is_safe": is_safe, "metrics": [metrics]}
class IntentClassifier(BaseAgent):
def __init__(self):
fallback_prompt = """You are a clinical intent classifier.
Classify into: 'diagnosis', 'treatment', 'monitoring', or 'general'."""
super().__init__(fallback_prompt, "IntentClassifier.txt")
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
start_time = time.time()
response = await self.llm.ainvoke([SystemMessage(content=self.system_prompt)] + state["messages"])
end_time = time.time()
intent = response.content.lower().strip()
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
metrics = {
"agent": "IntentClassifier",
"tokens": tokens,
"time": round(end_time - start_time, 3)
}
return {"intent_type": intent, "metrics": [metrics]}
class ClinicalSpecialist(BaseAgent):
"""
Clinical specialist agent with structured output including evidence citations.
Bug 12.3: Provides explainability through guideline sources and evidence levels.
"""
def __init__(self, specialty: str):
fallback_prompt = f"You are a clinical specialist in {specialty}. Provide expert medical support."
super().__init__(fallback_prompt, f"ClinicalSpecialist_{specialty}.txt")
self.specialty = specialty
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
"""
Run clinical specialist with structured output requiring evidence citations.
Returns both the text response and evidence citations in AgentState.
"""
messages = [SystemMessage(content=self.system_prompt)] + state["messages"]
logger.info(f"--- ClinicalSpecialist ({self.specialty}): Running with evidence structure ---")
start_time = time.time()
try:
# Use structured output with the LLM if available
try:
# Try to use with_structured_output for models that support it
llm_with_output = self.llm.with_structured_output(ClinicalOutputWithEvidence)
response = await llm_with_output.ainvoke(messages, config=config)
except (AttributeError, NotImplementedError):
# Fallback: regular invocation and manual extraction
logger.warning(f"Model does not support structured output, using fallback")
response = await self.llm.ainvoke(messages, config=config)
# Create a basic ClinicalOutputWithEvidence from the response
from src.core.evidence_models import Citation
response = ClinicalOutputWithEvidence(
recommendation=response.content[:200] if hasattr(response, 'content') else str(response),
explanation=response.content if hasattr(response, 'content') else str(response),
citations=[
Citation(
source_document="Knowledge Base",
evidence_level="C",
section="General"
)
],
confidence_score=0.7
)
end_time = time.time()
# Extract tokens
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif isinstance(response, dict) and "usage_metadata" in response:
tokens = response["usage_metadata"].get("total_tokens", 0)
elif hasattr(response, "response_metadata") and "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
confidence = 0.7
if isinstance(response, ClinicalOutputWithEvidence):
output_content = response.recommendation
citations = response.citations
confidence = getattr(response, "confidence_score", 0.7)
else:
output_content = response.content if hasattr(response, 'content') else str(response)
citations = []
if isinstance(output_content, str):
match = re.search(r"Confidence(?:\s+Score)?:\s*([0-9.]+)", output_content, re.IGNORECASE)
if match:
try:
confidence = float(match.group(1))
except ValueError:
pass
metrics = {
"agent": f"ClinicalSpecialist({self.specialty})",
"tokens": tokens,
"time": round(end_time - start_time, 3),
"confidence": confidence
}
# Build evidence citations from the structured output
evidence_citations = []
for idx, citation in enumerate(citations):
evidence_citation = {
"recommendation_id": f"{self.specialty}_{idx}",
"source_document": citation.source_document if hasattr(citation, 'source_document') else "Unknown",
"page_number": getattr(citation, 'page_number', None),
"evidence_level": getattr(citation, 'evidence_level', 'C'),
"agent_name": f"ClinicalSpecialist({self.specialty})",
"timestamp": datetime.utcnow().isoformat()
}
evidence_citations.append(evidence_citation)
from langchain_core.messages import AIMessage
return {
"messages": [AIMessage(content=output_content)],
"metrics": [metrics],
"evidence_citations": evidence_citations
}
except Exception as e:
logger.error(f"Error in ClinicalSpecialist({self.specialty}).run: {e}")
raise
class OutputMerger(BaseAgent):
def __init__(self):
fallback_prompt = "You are a clinical coordinator. Merge outputs into a single cohesive report."
super().__init__(fallback_prompt, "OutputMerger.txt")
@throttle_agent
async def run(self, state: AgentState, config=None, **kwargs):
latest_user_message = None
for message in reversed(state["messages"]):
if getattr(message, "type", None) == "human":
latest_user_message = message.content
break
human_messages = []
if latest_user_message:
human_messages.append(HumanMessage(content=f"Original user request:\n{latest_user_message}"))
clinician_outputs = state.get("clinician_outputs") or []
if clinician_outputs:
human_messages.append(
HumanMessage(content="Latest specialist outputs:\n" + "\n\n".join(clinician_outputs))
)
messages = [SystemMessage(content=self.system_prompt)] + human_messages
start_time = time.time()
response = None
full_content = ""
async for chunk in self.llm.astream(messages, config=config):
response = chunk
if chunk and hasattr(chunk, "content") and isinstance(chunk.content, str):
full_content += chunk.content
if response is None:
response = await self.llm.ainvoke(messages, config=config)
if response and hasattr(response, "content") and isinstance(response.content, str):
full_content = response.content
end_time = time.time()
tokens = 0
if hasattr(response, "usage_metadata") and response.usage_metadata:
tokens = response.usage_metadata.get("total_tokens", 0)
elif "token_usage" in response.response_metadata:
tokens = response.response_metadata["token_usage"].get("total_tokens", 0)
confidence = None
if full_content:
match = re.search(r"Confidence(?:\s+Score)?:\s*([0-9.]+)", full_content, re.IGNORECASE)
if match:
try:
confidence = float(match.group(1))
except ValueError:
pass
metrics = {
"agent": self.__class__.__name__,
"tokens": tokens,
"time": round(end_time - start_time, 3),
"confidence": confidence
}
return {"messages": [response], "metrics": [metrics]}
class ResearchAgent(BaseAgent):
def __init__(self):
fallback_prompt = """You are a medical research assistant. Provide detailed information for researchers."""
super().__init__(fallback_prompt, "ResearchAgent.txt", tools=[web_search_tool, page_indexed_retrieval])
class DietarySpecialist(BaseAgent):
def __init__(self):
fallback_prompt = """You are a certified dietary specialist. Provide advice based on guidelines."""
super().__init__(fallback_prompt, "DietarySpecialist.txt", tools=[search_guidelines, get_nutritional_data, page_indexed_retrieval, save_patient_memory, get_patient_memory])