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