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