dmChatbotBackend / src /core /graph.py
github-actions
Auto deploy from GitHub
23bfb17
Raw History Blame Contribute Delete
23.8 kB
import re
import time
from langgraph.graph import StateGraph, END
from src.core.state import AgentState
from src.agents.agents import RoleClassifier, ResponseValidator, SafetyCheck, IntentClassifier
from langchain_core.messages import AIMessage, ToolMessage
from langgraph.prebuilt import ToolNode
from src.tools.web_tools import web_search_tool
from src.utils.logger import setup_logger
from src.core.policy import PlatformGateService, audit_event, gate_update
from src.tools.fhir_memory import save_chat_as_fhir, _ensure_session_id
from src.runtime.execution_context import ExecutionContext
from src.runtime.manager_runtime import ManagerRuntime
from src.skills.registry import default_skill_registry
from src.skills.adapters import create_worker_adapter
from src.skills.runtime import SkillRuntime
from src.runtime.worker_runtime import WorkerRuntime
logger = setup_logger("MedicalPipeline")
WORKER_NODES = {
"patient-support": "patient_llm",
"caregiver-support": "caregiver_llm",
"researcher": "research_agent",
"diabetes-research": "research_agent",
"diabetes-dietary-coach": "dietary_assist",
"clinical-diagnosis-support": "diagnosis_assist",
"clinical-treatment-support": "treatment_assist",
"clinical-monitoring-support": "monitoring_assist",
"clinical-general-support": "general_assist",
}
manager_runtime = ManagerRuntime(default_skill_registry())
MAX_RECOVERY_ATTEMPTS = 3
EMERGENCY_PATTERNS = [
r"\b(unconscious|passed out|not breathing|seizure|seizing|stroke|heart attack|chest pain|coma)\b",
r"\b(dka|ketoacidosis|hypoglycemic emergency|hyperglycemic emergency|insulin overdose)\b",
]
def _platform_agent(agent_type):
return agent_type()
def log_step(name: str, output: str = None):
"""Utility to log both to terminal and return a state update for the logs list."""
logger.info(f"Executing: {name}")
log_msg = f"➔ Executing Node: {name}"
if output:
log_msg += f"\nOutput: {output}"
return {"logs": [log_msg]}
def _extract_message_content(message) -> str:
if getattr(message, "content", None):
return message.content
return str(getattr(message, "tool_calls", ""))
def _set_latest_clinician_output(res: dict):
latest_output = ""
if res.get("messages"):
latest_output = _extract_message_content(res["messages"][-1])
if latest_output:
res["clinician_outputs"] = [latest_output]
else:
res["clinician_outputs"] = []
return res
def detect_emergency(state: AgentState) -> bool:
last_message = state.get("messages")[-1] if state.get("messages") else None
if not last_message:
return False
content = getattr(last_message, "content", "")
if not content:
return False
normalized = content.lower()
return any(re.search(pattern, normalized) for pattern in EMERGENCY_PATTERNS)
def _manager_plan_for_state(state: AgentState):
"""Build a Manager route for direct routing calls without prior node state."""
existing_plan = state.get("manager_route")
if existing_plan:
return existing_plan
intent = state.get("intent_type")
context = ExecutionContext(
user_role=state.get("user_role") or ("clinician" if intent else None),
intent=intent,
user_query=_extract_message_content(state["messages"][-1]) if state.get("messages") else "",
)
return manager_runtime.plan_route(context)
async def role_classifier_node(state: AgentState):
res = await _platform_agent(RoleClassifier).run(state)
log = log_step("Role Classifier", f"Detected Role: {res.get('user_role')}")
res.update(log)
if "metrics" in res and res["metrics"]:
for m in res["metrics"]:
m["skill_id"] = "role-routing"
m["onloaded_skill"] = "role-routing"
m["offloaded_skill"] = "role-routing"
# Intent is only needed for the clinician pathway.
if res.get('user_role') != "clinician":
res['intent_type'] = "general"
route_context = ExecutionContext(
request_id=state.get("request_id") or "",
trace_id=state.get("trace_id") or "",
session_id=state.get("session_id") or "",
patient_id=state.get("patient_id"),
user_query=_extract_message_content(state["messages"][-1]) if state.get("messages") else "",
user_role=res.get("user_role"),
intent=res.get("intent_type"),
)
route_plan = manager_runtime.plan_route(route_context)
res["manager_route"] = route_plan
selected_skill = route_plan.get("manager_skill") or route_plan.get("worker_skill")
if selected_skill:
res["active_skill_id"] = selected_skill
res["active_skill_version"] = "1.0.0"
res["skill_plan"] = [{"skill_id": selected_skill, "reason": "manager route"}]
res["skill_trace"] = (state.get("skill_trace") or []) + [
audit_event({**state, "active_skill_id": "role-routing", "active_skill_version": "1.0.0"}, "skill_onloaded", agent="RoleClassifier"),
audit_event({**state, "active_skill_id": selected_skill, "active_skill_version": "1.0.0"}, "skill_selected"),
audit_event({**state, "active_skill_id": "role-routing", "active_skill_version": "1.0.0"}, "skill_offloaded", agent="RoleClassifier")
]
return res
from langchain_core.runnables.config import RunnableConfig
async def _worker_skill_node(state: AgentState, skill_id: str, config: RunnableConfig):
registry = default_skill_registry()
adapter = create_worker_adapter(skill_id, registry=registry)
if state.get("runtime_mode", "skill") == "agent":
result = await adapter.execute(state, config=config)
else:
worker = WorkerRuntime(registry)
try:
result = await worker.execute_adapter(adapter, state, config=config)
finally:
worker.flush_and_destroy()
output = result.output
messages = result.messages
latest = messages[-1] if messages else None
content = latest.content if getattr(latest, "content", "") else str(getattr(latest, "tool_calls", ""))
skill_event = audit_event({**state, "active_skill_id": result.skill_id, "active_skill_version": result.skill_version}, "skill_executed")
agent_metric = dict(result.metrics) if result.metrics else {}
agent_metric["skill_id"] = result.skill_id
agent_metric["onloaded_skill"] = result.skill_id
agent_metric["offloaded_skill"] = result.skill_id
if "agent" not in agent_metric or not agent_metric["agent"]:
agent_metric["agent"] = adapter.agent_name or skill_id
base = {
"messages": messages,
"logs": log_step(skill_id, content)["logs"],
"metrics": [agent_metric],
"active_skill_id": result.skill_id,
"active_skill_version": result.skill_version,
"skill_trace": (state.get("skill_trace") or []) + [
audit_event({**state, "active_skill_id": result.skill_id, "active_skill_version": result.skill_version}, "skill_onloaded", agent=agent_metric["agent"]),
skill_event,
audit_event({**state, "active_skill_id": result.skill_id, "active_skill_version": result.skill_version}, "skill_offloaded", agent=agent_metric["agent"]),
],
}
if skill_id == "patient-support":
clean_out = output or (latest.content if latest and not (hasattr(latest, "tool_calls") and latest.tool_calls) else "")
if clean_out and not (str(clean_out).strip().startswith('{"name":') and ('"parameters":' in str(clean_out) or '"arguments":' in str(clean_out))):
base["patient_response"] = clean_out
elif skill_id == "diabetes-research":
clean_out = output or (latest.content if latest and not (hasattr(latest, "tool_calls") and latest.tool_calls) else "")
if clean_out and not (str(clean_out).strip().startswith('{"name":') and ('"parameters":' in str(clean_out) or '"arguments":' in str(clean_out))):
base["research_output"] = clean_out
elif skill_id.startswith("clinical-"):
clinical_output = output if isinstance(output, list) else ([output] if output else [])
base["clinician_outputs"] = (state.get("clinician_outputs") or []) + clinical_output
return base
async def patient_llm_node(state: AgentState, config: RunnableConfig):
result = await _worker_skill_node(state, "patient-support", config)
return result
async def caregiver_llm_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "caregiver-support", config)
async def validator_node(state: AgentState):
res = await _platform_agent(ResponseValidator).run(state)
log = log_step("Response Validator", f"Valid: {res.get('is_valid')}")
res.update(log)
if "metrics" in res and res["metrics"]:
for m in res["metrics"]:
m["skill_id"] = "clinical-response-validation"
m["onloaded_skill"] = "clinical-response-validation"
m["offloaded_skill"] = "clinical-response-validation"
gate_result = PlatformGateService.record(state, "validation", bool(res.get("is_valid")))
res.update(gate_result)
return res
async def safety_check_node(state: AgentState):
res = await _platform_agent(SafetyCheck).run(state)
log = log_step("Safety Check", f"Safe: {res.get('is_safe')}")
res.update(log)
gate_result = PlatformGateService.record(state, "safety", bool(res.get("is_safe")))
res.update(gate_result)
return res
async def recovery_loop_node(state: AgentState):
attempts = state.get("attempts", 0) + 1
log = log_step("Recovery Loop", f"Attempt {attempts}")
return {
"attempts": attempts,
"messages": [AIMessage(content="[RECOVERY] Let me try rephrasing or improving my previous response.")],
"logs": log["logs"]
}
async def emergency_response_node(state: AgentState):
response = (
"This appears to be a possible medical emergency. Seek urgent medical assistance now "
"and do not delay care. If the person is unconscious, not breathing, or having a seizure, "
"call emergency services immediately."
)
log = log_step("Emergency Fast Path", response)
return {
"messages": [AIMessage(content=response)],
"logs": log["logs"],
"is_valid": True,
"is_safe": True,
"gates": gate_update(state, "emergency", "triggered", "deterministic emergency pattern matched"),
"skill_trace": (state.get("skill_trace") or []) + [audit_event(state, "gate", gate="emergency", decision="triggered")],
}
async def intent_classifier_node(state: AgentState):
res = await _platform_agent(IntentClassifier).run(state)
log = log_step("Intent Classifier", f"Intent: {res.get('intent_type')}")
res.update(log)
if "metrics" in res and res["metrics"]:
for m in res["metrics"]:
m["skill_id"] = "clinical-intent-routing"
m["onloaded_skill"] = "clinical-intent-routing"
m["offloaded_skill"] = "clinical-intent-routing"
route_context = ExecutionContext(
request_id=state.get("request_id") or "",
trace_id=state.get("trace_id") or "",
session_id=state.get("session_id") or "",
patient_id=state.get("patient_id"),
user_query=_extract_message_content(state["messages"][-1]) if state.get("messages") else "",
user_role=state.get("user_role"),
intent=res.get("intent_type"),
)
route_plan = manager_runtime.plan_route(route_context)
intent_skill = route_plan.get("worker_skill")
res["manager_route"] = route_plan
res["active_skill_id"] = intent_skill
res["active_skill_version"] = "1.0.0"
res["skill_plan"] = [{"skill_id": intent_skill, "reason": "clinical intent routing"}]
res["skill_trace"] = (state.get("skill_trace") or []) + [
audit_event({**state, "active_skill_id": "clinical-intent-routing", "active_skill_version": "1.0.0"}, "skill_onloaded", agent="IntentClassifier"),
audit_event({**state, "active_skill_id": intent_skill, "active_skill_version": "1.0.0"}, "skill_selected"),
audit_event({**state, "active_skill_id": "clinical-intent-routing", "active_skill_version": "1.0.0"}, "skill_offloaded", agent="IntentClassifier"),
]
return res
async def diagnosis_assist_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "clinical-diagnosis-support", config)
async def treatment_assist_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "clinical-treatment-support", config)
async def monitoring_assist_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "clinical-monitoring-support", config)
async def general_assist_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "clinical-general-support", config)
async def research_agent_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "diabetes-research", config)
async def dietary_assist_node(state: AgentState, config: RunnableConfig):
return await _worker_skill_node(state, "diabetes-dietary-coach", config)
try:
from langgraph.runtime import Runtime
except ImportError:
Runtime = None
def _ensure_tool_config(config: RunnableConfig = None) -> dict:
call_config = dict(config) if config else {}
configurable = dict(call_config.get("configurable", {}))
if "__pregel_runtime" not in configurable and Runtime is not None:
try:
configurable["__pregel_runtime"] = Runtime()
except Exception:
pass
call_config["configurable"] = configurable
return call_config
async def tool_node_with_logging(state: AgentState, config: RunnableConfig = None):
start_time = time.time()
call_config = _ensure_tool_config(config)
res = await tool_node.ainvoke(state, config=call_config)
end_time = time.time()
# Extract tool execution details with input arguments and responses
tool_messages = res.get("messages", []) if isinstance(res, dict) else (res if isinstance(res, list) else [])
last_ai_msg = None
for msg in reversed(state.get("messages", [])):
if getattr(msg, "tool_calls", None):
last_ai_msg = msg
break
tc_map = {}
if last_ai_msg and getattr(last_ai_msg, "tool_calls", None):
for tc in last_ai_msg.tool_calls:
tc_id = tc.get("id")
if tc_id:
tc_map[tc_id] = tc
tool_details = []
for tm in tool_messages:
if isinstance(tm, ToolMessage):
tool_name = getattr(tm, "name", None) or "Tool"
tool_call_id = getattr(tm, "tool_call_id", None)
matched_tc = tc_map.get(tool_call_id, {})
args = matched_tc.get("args") if matched_tc else None
out = str(getattr(tm, "content", ""))
tool_details.append({
"name": tool_name,
"input": args,
"output": out[:600] if out else "",
"time": round(end_time - start_time, 3)
})
names = [t["name"] for t in tool_details]
tools_str = ", ".join(dict.fromkeys(names)) if names else "Tools"
output_summary = f"{len(tool_messages)} tool(s) executed ({tools_str})." if tool_messages else "Tool executed."
log = log_step("Executing Tools (RAG/Web Search)", output_summary)
duration = round(end_time - start_time, 3)
metrics = {
"agent": "ToolsNode",
"tool_name": tools_str,
"tools": tool_details if tool_details else [{"name": tools_str, "time": duration}],
"tokens": 0,
"time": duration
}
if isinstance(res, list):
return {"messages": res, "logs": log["logs"], "metrics": [metrics]}
res.update(log)
res.update({"metrics": [metrics]})
return res
async def persistence_node(state: AgentState):
"""Save the current chat history to Supabase in FHIR format."""
patient_id = state.get("patient_id", "anonymous")
session_id = _ensure_session_id(patient_id, state.get("session_id"))
# Convert LangChain messages to a simple list of dicts for the tool
formatted_messages = []
for msg in state["messages"]:
role = "user" if msg.type == "human" else "assistant"
formatted_messages.append({"role": role, "content": msg.content})
# Persist patient-facing and caregiver-facing conversations for the active patient.
if state.get("user_role") in ("patient", "caregiver"):
start_time = time.time()
res = save_chat_as_fhir.invoke({
"patient_id": patient_id,
"messages": formatted_messages,
"session_id": session_id,
})
end_time = time.time()
log = log_step("FHIR Persistence", res)
metrics = {
"agent": "PersistenceNode",
"skill_id": "fhir-patient-memory",
"onloaded_skill": "fhir-patient-memory",
"offloaded_skill": "fhir-patient-memory",
"tokens": 0,
"time": round(end_time - start_time, 3)
}
return {
"session_id": session_id,
"logs": log["logs"],
"metrics": [metrics],
}
return {"session_id": session_id}
# Define routing functions
def route_after_role(state: AgentState):
role = state["user_role"]
if role in {"patient", "caregiver"} and detect_emergency(state):
return "emergency_response"
route_plan = _manager_plan_for_state(state)
if route_plan.get("manager_skill") == "clinical-intent-routing":
return "intent_classifier"
return WORKER_NODES.get(route_plan.get("worker_skill"), END)
def route_research_agent(state: AgentState):
last_message = state["messages"][-1] if state.get("messages") else None
if not last_message:
return END
tool_msgs = [m for m in state.get("messages", []) if isinstance(m, ToolMessage)]
if len(tool_msgs) >= 3:
return END
if hasattr(last_message, "tool_calls") and last_message.tool_calls:
return "tools_node"
return END
def route_patient_llm(state: AgentState):
last_message = state["messages"][-1] if state.get("messages") else None
if not last_message:
return "validator"
tool_msgs = [m for m in state.get("messages", []) if isinstance(m, ToolMessage)]
if len(tool_msgs) >= 3:
return "validator"
if hasattr(last_message, "tool_calls") and last_message.tool_calls:
return "tools_node"
return "validator"
def route_after_tools(state: AgentState):
role = state.get("user_role")
if role == "researcher":
return "research_agent"
elif role == "dietary":
return "dietary_assist"
elif role == "caregiver":
return "caregiver_llm"
return "patient_llm"
def route_after_validator(state: AgentState):
if state.get("is_valid", False):
return "safety_check"
return "recovery_loop"
def route_after_safety(state: AgentState):
if state.get("is_safe", False):
return "persistence_node"
return "recovery_loop"
def route_after_persistence(state: AgentState):
return END
def route_after_recovery(state: AgentState):
attempts = state.get("attempts", 0)
if attempts >= MAX_RECOVERY_ATTEMPTS:
return "persistence_node"
role = state.get("user_role")
if role == "dietary":
return "dietary_assist"
elif role == "caregiver":
return "caregiver_llm"
elif role == "patient":
return "patient_llm"
return "persistence_node"
def route_after_intent(state: AgentState):
return WORKER_NODES.get(_manager_plan_for_state(state).get("worker_skill"), "general_assist")
from src.tools.dietary_tools import page_indexed_retrieval, search_guidelines, get_nutritional_data
# Updated Tool Node to include page indexing RAG
tools = [web_search_tool, page_indexed_retrieval, search_guidelines, get_nutritional_data]
tool_node = ToolNode(tools)
# Build the graph
builder = StateGraph(AgentState)
# Add nodes
builder.add_node("role_classifier", role_classifier_node)
builder.add_node("patient_llm", patient_llm_node)
builder.add_node("caregiver_llm", caregiver_llm_node)
builder.add_node("validator", validator_node)
builder.add_node("safety_check", safety_check_node)
builder.add_node("recovery_loop", recovery_loop_node)
builder.add_node("emergency_response", emergency_response_node)
builder.add_node("intent_classifier", intent_classifier_node)
builder.add_node("diagnosis_assist", diagnosis_assist_node)
builder.add_node("treatment_assist", treatment_assist_node)
builder.add_node("monitoring_assist", monitoring_assist_node)
builder.add_node("general_assist", general_assist_node)
builder.add_node("research_agent", research_agent_node)
builder.add_node("tools_node", tool_node_with_logging)
builder.add_node("dietary_assist", dietary_assist_node)
builder.add_node("persistence_node", persistence_node)
# Set entry point
builder.set_entry_point("role_classifier")
# Define edges
builder.add_conditional_edges("role_classifier", route_after_role, {
"patient_llm": "patient_llm",
"caregiver_llm": "caregiver_llm",
"intent_classifier": "intent_classifier",
"research_agent": "research_agent",
"dietary_assist": "dietary_assist",
"emergency_response": "emergency_response",
END: END
})
# Patient and Caregiver Pathways
builder.add_conditional_edges("patient_llm", route_patient_llm, {
"tools_node": "tools_node",
"validator": "validator"
})
builder.add_conditional_edges("caregiver_llm", route_patient_llm, {
"tools_node": "tools_node",
"validator": "validator"
})
builder.add_conditional_edges("validator", route_after_validator, {
"safety_check": "safety_check",
"recovery_loop": "recovery_loop"
})
builder.add_conditional_edges("safety_check", route_after_safety, {
"persistence_node": "persistence_node",
"recovery_loop": "recovery_loop"
})
builder.add_edge("persistence_node", END)
builder.add_conditional_edges("recovery_loop", route_after_recovery, {
"dietary_assist": "dietary_assist",
"caregiver_llm": "caregiver_llm",
"patient_llm": "patient_llm",
"persistence_node": "persistence_node"
})
builder.add_edge("emergency_response", "persistence_node")
# Clinician Pathway
builder.add_conditional_edges("intent_classifier", route_after_intent, {
"diagnosis_assist": "diagnosis_assist",
"treatment_assist": "treatment_assist",
"monitoring_assist": "monitoring_assist",
"general_assist": "general_assist"
})
builder.add_edge("diagnosis_assist", END)
builder.add_edge("treatment_assist", END)
builder.add_edge("monitoring_assist", END)
builder.add_edge("general_assist", END)
# Researcher Pathway
builder.add_conditional_edges("research_agent", route_research_agent, {
"tools_node": "tools_node",
END: END
})
# Shared Tool Pathway
builder.add_conditional_edges("tools_node", route_after_tools, {
"research_agent": "research_agent",
"patient_llm": "patient_llm",
"caregiver_llm": "caregiver_llm",
"dietary_assist": "dietary_assist"
})
# Dietary Pathway
def route_dietary_assist(state: AgentState):
last_message = state["messages"][-1] if state.get("messages") else None
if not last_message:
return "validator"
tool_msgs = [m for m in state.get("messages", []) if isinstance(m, ToolMessage)]
if len(tool_msgs) >= 3:
return "validator"
if hasattr(last_message, "tool_calls") and last_message.tool_calls:
return "tools_node"
return "validator"
builder.add_conditional_edges("dietary_assist", route_dietary_assist, {
"tools_node": "tools_node",
"validator": "validator"
})
# Compile the graph
medical_pipeline = builder.compile()