Download src/core/graph.py from DiabetesCareChatbot/dmChatbotBackend: direct link, hf CLI and curl.
- Browser
- Download file 23.8 kB
-
https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/src/core/graph.py
- Command line
-
hf download hf://spaces/DiabetesCareChatbot/dmChatbotBackend/src/core/graph.py
-
curl -L -o graph.py https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/src/core/graph.py
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() | |