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()