File size: 5,918 Bytes
b1198f0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
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()