import pytest import json import os from unittest.mock import patch from src.utils.auth import create_dev_token from main import convert_to_langchain_messages from langchain_core.messages import HumanMessage, AIMessage def test_health_check(client_main): response = client_main.get("/") assert response.status_code == 200 assert response.json() == {"status": "healthy", "message": "Medical AI Backend is running"} def test_history_conversion_ignores_unknown_roles_and_preserves_supported_roles(): messages = convert_to_langchain_messages( [ {"role": "user", "content": "hello"}, {"role": "assistant", "content": "hi"}, {"role": "system", "content": "ignore"}, ] ) assert [type(message) for message in messages] == [HumanMessage, AIMessage] assert [message.content for message in messages] == ["hello", "hi"] def test_process_rejects_missing_prompt(client_main): response = client_main.post("/process", json={"mode": "Standard Triage"}) assert response.status_code == 422 def test_seed_patient_data(client_main, clean_db): patient_id = "test-patient" response = client_main.post(f"/patient/seed?patient_id={patient_id}") assert response.status_code == 200 assert response.json()["status"] == "success" # Verify seed data exists in DB patient = clean_db.db["patients"].get(patient_id) assert patient is not None observations = clean_db._get_data("observations", {"patient_id": patient_id}) assert len(observations) == 3 def test_get_patient_summary(client_main, clean_db): patient_id = "summary-patient" # Seed data client_main.post(f"/patient/seed?patient_id={patient_id}") # Fetch summary response = client_main.get(f"/patient/{patient_id}") assert response.status_code == 200 summary = response.json() assert summary["patient"] == "Demo Patient" assert len(summary["recent_observations"]) == 3 def test_set_llm_provider(client_main): response = client_main.post("/config/llm?provider=openrouter") assert response.status_code == 200 data = response.json() assert data["status"] == "success" assert data["provider"] == "openrouter" def test_runtime_mode_toggle_requires_authentication(client_app): response = client_app.get("/config/runtime") assert response.status_code == 401 def test_runtime_mode_toggle_switches_and_is_reported(client_app, clean_db): token = create_dev_token("runtime-mode-user") headers = {"Authorization": f"Bearer {token}"} current = client_app.get("/config/runtime", headers=headers) assert current.status_code == 200 assert current.json()["mode"] == "skill" changed = client_app.post( "/config/runtime", json={"mode": "agent"}, headers=headers, ) assert changed.status_code == 200 assert changed.json() == {"mode": "agent"} seeded = client_app.post("/patient/seed", headers=headers) assert seeded.status_code == 200 processed = client_app.post( "/process", json={"prompt": "I need support."}, headers=headers, ) assert processed.status_code == 200 assert processed.json()["final_state"]["runtime_mode"] == "agent" client_app.post( "/config/runtime", json={"mode": "skill"}, headers=headers, ) def test_runtime_mode_toggle_rejects_unknown_mode(client_app): token = create_dev_token("runtime-mode-invalid") response = client_app.post( "/config/runtime", json={"mode": "unsupported"}, headers={"Authorization": f"Bearer {token}"}, ) assert response.status_code == 422 def test_process_pipeline_standard(client_main, clean_db): patient_id = "patient-standard" client_main.post(f"/patient/seed?patient_id={patient_id}") payload = { "prompt": "Hi, I have a slight headache.", "patient_id": patient_id, "mode": "Standard Triage" } response = client_main.post("/process", json=payload) assert response.status_code == 200 data = response.json() assert "messages" in data assert "final_state" in data assert "session_id" in data # Verify communications were saved (persistence node) session_id = data["session_id"] comms = clean_db._get_data("communications", {"session_id": session_id}) assert len(comms) > 0 def test_process_pipeline_cdm(client_main, clean_db): patient_id = "patient-cdm" client_main.post(f"/patient/seed?patient_id={patient_id}") payload = { "prompt": "How are my blood sugar trends looking?", "patient_id": patient_id, "mode": "CDM Proactive" } response = client_main.post("/process", json=payload) assert response.status_code == 200 data = response.json() assert "messages" in data assert "final_state" in data assert data["final_state"]["trend_analysis"] != "" def test_process_stream(client_main, clean_db): patient_id = "patient-stream" client_main.post(f"/patient/seed?patient_id={patient_id}") payload = { "prompt": "Hello", "patient_id": patient_id, "mode": "Standard Triage" } response = client_main.post("/process_stream", json=payload) assert response.status_code == 200 assert "text/event-stream" in response.headers["content-type"] # Read SSE events events = [] for line in response.iter_lines(): if line.startswith("data: "): event_data = json.loads(line[6:]) events.append(event_data) assert len(events) > 0 node_events = [e for e in events if e["type"] == "node"] assert len(node_events) > 0 end_events = [e for e in events if e["type"] == "end"] assert len(end_events) == 1 assert "final_state" in end_events[0] def test_mcp_endpoint(client_main, clean_db): # Test MCP list_tools req = {"method": "list_tools", "params": {}, "id": 1} response = client_main.post("/mcp", json=req) assert response.status_code == 200 res_data = response.json() assert "result" in res_data tools = res_data["result"] assert any(t["name"] == "analyze_health_trends" for t in tools) # Test MCP call_tool patient_id = "mcp-patient" client_main.post(f"/patient/seed?patient_id={patient_id}") call_req = { "method": "call_tool", "params": {"name": "analyze_health_trends", "arguments": {"patient_id": patient_id}}, "id": 2 } response2 = client_main.post("/mcp", json=call_req) assert response2.status_code == 200 res_data2 = response2.json() assert "result" in res_data2 analysis, _ = res_data2["result"] assert "Trend Analysis Report" in analysis def test_fhir_rest_endpoints(client_main, clean_db): patient_id = "fhir-rest-patient" client_main.post(f"/patient/seed?patient_id={patient_id}") # Patient res = client_main.get(f"/Patient/{patient_id}") assert res.status_code == 200 assert res.json()["resourceType"] == "Patient" # Observation res_obs = client_main.get(f"/Observation?patient={patient_id}") assert res_obs.status_code == 200 assert isinstance(res_obs.json(), list) assert len(res_obs.json()) == 3 # App.py endpoints (Authenticated) def test_authenticated_endpoints(client_app, clean_db): patient_id = "auth-patient-123" token = create_dev_token(patient_id) headers = {"Authorization": f"Bearer {token}"} # 1. Seed data res_seed = client_app.post("/patient/seed", headers=headers) assert res_seed.status_code == 200 assert res_seed.json()["status"] == "success" # 2. Get Summary res_sum = client_app.get("/patient/summary", headers=headers) assert res_sum.status_code == 200 assert res_sum.json()["patient"] == "Authenticated Patient" # 3. Process request payload = { "prompt": "I need some support.", "mode": "Standard Triage" } res_proc = client_app.post("/process", json=payload, headers=headers) assert res_proc.status_code == 200 proc_data = res_proc.json() session_id = proc_data["session_id"] assert session_id is not None # 4. List Sessions res_sess = client_app.get("/sessions", headers=headers) assert res_sess.status_code == 200 assert len(res_sess.json()) > 0 assert any(s["id"] == session_id for s in res_sess.json()) # 5. Get Session History res_hist = client_app.get(f"/sessions/{session_id}/history", headers=headers) assert res_hist.status_code == 200 assert len(res_hist.json()) > 0 def test_clinician_feedback_score(client_app): token = create_dev_token("feedback-user") headers = {"Authorization": f"Bearer {token}"} payload = { "trace_id": "test-trace-123", "score": 5, "feedback": "Excellent response quality." } csv_file = "clinician_scores.csv" if os.path.exists(csv_file): os.remove(csv_file) try: response = client_app.post("/feedback/score", json=payload, headers=headers) assert response.status_code == 200 assert response.json() == {"status": "success", "message": "Score saved successfully"} assert os.path.exists(csv_file) finally: if os.path.exists(csv_file): os.remove(csv_file)