Download tests/test_api_endpoints.py from DiabetesCareChatbot/dmChatbotBackend: direct link, hf CLI and curl.
- Browser
- Download file 9.33 kB
-
https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/tests/test_api_endpoints.py
- Command line
-
hf download hf://spaces/DiabetesCareChatbot/dmChatbotBackend/tests/test_api_endpoints.py
-
curl -L -o test_api_endpoints.py https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/tests/test_api_endpoints.py
9.33 kB
| 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) | |