| from langgraph.graph import StateGraph, END |
| from src.core.state import AgentState |
| from src.agents.agent_instances import ( |
| health_coach, trend_analyzer, validator, safety_check |
| ) |
| from src.tools.fhir_memory import ( |
| get_patient_summary_fhir, save_observation, ingest_fhir_bundle, |
| get_medications_by_patient, save_medication |
| ) |
| from src.tools.web_tools import web_search_tool |
| from langgraph.prebuilt import ToolNode |
| import time |
| from src.utils.logger import setup_logger |
|
|
| logger = setup_logger("CDMPipeline") |
|
|
| def log_step(name: str, output: str = None): |
| logger.info(f"Executing: {name}") |
| log_msg = f"➔ CDM Node: {name}" |
| if output: |
| log_msg += f"\nOutput: {output}" |
| return {"logs": [log_msg]} |
|
|
| async def data_fetch_node(state: AgentState): |
| """ |
| Fetch FHIR data and medications for CDM context. |
| Bug 7.3: Include medication context in CDM pipeline. |
| """ |
| patient_id = state.get("patient_id", "unknown") |
| if patient_id == "unknown": |
| return log_step("Data Fetch", "No Patient ID provided.") |
| |
| start_time = time.time() |
| |
| |
| summary = get_patient_summary_fhir.invoke({"patient_id": patient_id}) |
| |
| |
| medications = get_medications_by_patient.invoke({"patient_id": patient_id}) |
| |
| |
| if isinstance(summary, dict) and "active_medications" in summary: |
| summary["active_medications"] = medications if medications else [] |
| |
| end_time = time.time() |
| |
| log_output = f"Retrieved summary and medications for {patient_id}" |
| if medications: |
| log_output += f" ({len(medications)} active medications)" |
| |
| log = log_step("FHIR Data Fetch", log_output) |
| |
| metrics = { |
| "agent": "DataFetchNode", |
| "tokens": 0, |
| "time": round(end_time - start_time, 3) |
| } |
| |
| return { |
| "fhir_data": [summary] if isinstance(summary, dict) else [], |
| "current_medications": medications if isinstance(medications, list) else [], |
| "logs": log["logs"], |
| "metrics": [metrics] |
| } |
|
|
| async def trend_analyzer_node(state: AgentState): |
| res = await trend_analyzer.run(state) |
| log = log_step("Trend Analyzer", res.get("trend_analysis", "")) |
| res.update(log) |
| return res |
|
|
| async def health_coach_node(state: AgentState): |
| res = await health_coach.run(state) |
| last_msg = res["messages"][-1] |
| content = last_msg.content if getattr(last_msg, "content", "") else "Tool calls generated." |
| log = log_step("Health Coach", content) |
| res.update(log) |
| return res |
|
|
| async def validator_node(state: AgentState): |
| from src.agents.agent_instances import validator |
| res = await validator.run(state) |
| log = log_step("Response Validator", f"Valid: {res.get('is_valid')}") |
| res.update(log) |
| return res |
|
|
| async def safety_check_node(state: AgentState): |
| from src.agents.agent_instances import safety_check |
| res = await safety_check.run(state) |
| log = log_step("Safety Check", f"Safe: {res.get('is_safe')}") |
| res.update(log) |
| return res |
|
|
| |
| def route_health_coach(state: AgentState): |
| last_message = state["messages"][-1] |
| if hasattr(last_message, "tool_calls") and last_message.tool_calls: |
| return "tools_node" |
| return "validator" |
|
|
| def route_after_tools(state: AgentState): |
| return "health_coach" |
|
|
| def route_after_validator(state: AgentState): |
| if state.get("is_valid", False): |
| return "safety_check" |
| return "health_coach" |
|
|
| |
| cdm_tools = [web_search_tool, save_observation, ingest_fhir_bundle, get_patient_summary_fhir, get_medications_by_patient, save_medication] |
| tool_node = ToolNode(cdm_tools) |
|
|
| async def tools_node_with_metrics(state: AgentState): |
| start_time = time.time() |
| res = await tool_node.ainvoke(state) |
| end_time = time.time() |
| |
| metrics = { |
| "agent": "CDMToolsNode", |
| "tokens": 0, |
| "time": round(end_time - start_time, 3) |
| } |
| |
| |
| return {"messages": res, "metrics": [metrics]} |
|
|
| |
| builder = StateGraph(AgentState) |
|
|
| builder.add_node("data_fetch", data_fetch_node) |
| builder.add_node("trend_analyzer", trend_analyzer_node) |
| builder.add_node("health_coach", health_coach_node) |
| builder.add_node("tools_node", tools_node_with_metrics) |
| builder.add_node("validator", validator_node) |
| builder.add_node("safety_check", safety_check_node) |
|
|
| builder.set_entry_point("data_fetch") |
| builder.add_edge("data_fetch", "trend_analyzer") |
| builder.add_edge("trend_analyzer", "health_coach") |
|
|
| builder.add_conditional_edges("health_coach", route_health_coach, { |
| "tools_node": "tools_node", |
| "validator": "validator" |
| }) |
|
|
| builder.add_edge("tools_node", "health_coach") |
|
|
| builder.add_conditional_edges("validator", route_after_validator, { |
| "safety_check": "safety_check", |
| "health_coach": "health_coach" |
| }) |
|
|
| builder.add_conditional_edges("safety_check", lambda x: "end" if x.get("is_safe") else "health_coach", { |
| "end": END, |
| "health_coach": "health_coach" |
| }) |
|
|
| cdm_pipeline = builder.compile() |
|
|