import importlib.util import sys import types import unittest from pathlib import Path ROOT = Path(__file__).resolve().parents[2] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) class _DummyLogger: def info(self, *_args, **_kwargs): return None class _DummyStateGraph: def __init__(self, *args, **kwargs): pass def add_node(self, *args, **kwargs): return None def add_conditional_edges(self, *args, **kwargs): return None def add_edge(self, *args, **kwargs): return None def set_entry_point(self, *args, **kwargs): return None def compile(self): return self class DummyMessage: def __init__(self, message_type, content): self.type = message_type self.content = content def _load_graph_module(): langgraph = types.ModuleType("langgraph") graph_module = types.ModuleType("langgraph.graph") graph_module.StateGraph = _DummyStateGraph graph_module.END = "END" langgraph.graph = graph_module prebuilt_module = types.ModuleType("langgraph.prebuilt") prebuilt_module.ToolNode = lambda tools: tools langchain_core = types.ModuleType("langchain_core") messages_module = types.ModuleType("langchain_core.messages") class AIMessage: pass class ToolMessage: pass messages_module.AIMessage = AIMessage messages_module.ToolMessage = ToolMessage runnables_module = types.ModuleType("langchain_core.runnables") config_module = types.ModuleType("langchain_core.runnables.config") config_module.RunnableConfig = dict runnables_module.config = config_module langchain_core.messages = messages_module langchain_core.runnables = runnables_module sys.modules.setdefault("langgraph", langgraph) sys.modules.setdefault("langgraph.graph", graph_module) sys.modules.setdefault("langgraph.prebuilt", prebuilt_module) sys.modules.setdefault("langchain_core", langchain_core) sys.modules.setdefault("langchain_core.messages", messages_module) sys.modules.setdefault("langchain_core.runnables", runnables_module) sys.modules.setdefault("langchain_core.runnables.config", config_module) dummy_tools_module = types.ModuleType("src.tools.web_tools") dummy_tools_module.web_search_tool = object() captured = {} class DummySaveFHIR: def invoke(self, payload): captured.update(payload) return "saved" dummy_fhir_module = types.ModuleType("src.tools.fhir_memory") dummy_fhir_module.save_chat_as_fhir = DummySaveFHIR() logger_module = types.ModuleType("src.utils.logger") logger_module.setup_logger = lambda _name: _DummyLogger() dietary_module = types.ModuleType("src.tools.dietary_tools") dietary_module.page_indexed_retrieval = object() dietary_module.search_guidelines = object() dietary_module.get_nutritional_data = object() agent_instances_module = types.ModuleType("src.agents.agent_instances") agent_instances_module.role_classifier = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.patient_llm = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.caregiver_llm = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.validator = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.safety_check = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.intent_classifier = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.diagnosis_assist = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.treatment_assist = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.monitoring_assist = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.general_assist = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.output_merger = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.research_agent = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) agent_instances_module.dietary_assist = types.SimpleNamespace(run=lambda *_args, **_kwargs: None) sys.modules.setdefault("src.tools.web_tools", dummy_tools_module) sys.modules.setdefault("src.tools.fhir_memory", dummy_fhir_module) sys.modules.setdefault("src.utils.logger", logger_module) sys.modules.setdefault("src.tools.dietary_tools", dietary_module) sys.modules.setdefault("src.agents.agent_instances", agent_instances_module) spec = importlib.util.spec_from_file_location("graph_under_test", ROOT / "backend" / "src" / "core" / "graph.py") module = importlib.util.module_from_spec(spec) sys.modules[spec.name] = module spec.loader.exec_module(module) return module, captured class TestSessionPersistence(unittest.IsolatedAsyncioTestCase): async def test_persistence_node_passes_session_id_to_fhir(self): graph, captured = _load_graph_module() state = { "messages": [ DummyMessage("human", "What is my glucose trend?"), DummyMessage("ai", "Your glucose trend is stable."), ], "user_role": "patient", "patient_id": "patient-123", "session_id": "session-456", } result = await graph.persistence_node(state) self.assertEqual(captured["patient_id"], "patient-123") self.assertEqual(captured["session_id"], "session-456") self.assertEqual(captured["messages"][0]["role"], "user") self.assertEqual(captured["messages"][1]["role"], "assistant") self.assertEqual(result["metrics"][0]["agent"], "PersistenceNode") if __name__ == "__main__": unittest.main()