File size: 6,455 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
155
156
157
158
159
160
161
162
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 _DummyMessage:
    def __init__(self, content):
        self.content = content


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


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

    dummy_fhir_module = types.ModuleType("src.tools.fhir_memory")
    dummy_fhir_module.save_chat_as_fhir = types.SimpleNamespace(invoke=lambda payload: {"saved": True})

    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


class TestGraphRouting(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.graph = _load_graph_module()

    def test_role_routes_caregiver_to_caregiver_llm(self):
        state = {"user_role": "caregiver"}
        self.assertEqual(self.graph.route_after_role(state), "caregiver_llm")

    def test_role_routes_clinician_to_intent_classifier(self):
        state = {"user_role": "clinician"}
        self.assertEqual(self.graph.route_after_role(state), "intent_classifier")

    def test_tools_route_returns_caregiver_llm_for_caregiver(self):
        state = {"user_role": "caregiver"}
        self.assertEqual(self.graph.route_after_tools(state), "caregiver_llm")

    def test_recovery_route_returns_caregiver_llm_for_caregiver(self):
        state = {"user_role": "caregiver"}
        self.assertEqual(self.graph.route_after_recovery(state), "caregiver_llm")

    def test_recovery_route_persists_after_max_attempts(self):
        state = {"user_role": "patient", "attempts": 3}
        self.assertEqual(self.graph.route_after_recovery(state), "persistence_node")

    def test_emergency_route_triggers_fast_path_for_patient(self):
        state = {
            "user_role": "patient",
            "messages": [_DummyMessage("The patient is unconscious and not breathing.")]
        }
        self.assertEqual(self.graph.route_after_role(state), "emergency_response")

    def test_intent_classifier_routes_general_to_general_assist(self):
        state = {"intent_type": "general"}
        self.assertEqual(self.graph.route_after_intent(state), "general_assist")


if __name__ == "__main__":
    unittest.main()