dmChatbotBackend / tests /test_session_persistence.py
github-actions
Auto deploy from GitHub
b1198f0
Raw
History Blame Contribute Delete
5.92 kB
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()