File size: 6,174 Bytes
8df47e6 | 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 | """Quick validation tests for MedTriage environment endpoints."""
import httpx
import json
BASE = "http://127.0.0.1:7860"
ALL_PASS = True
def check(name, condition, detail=""):
global ALL_PASS
detail_str = f" ({detail})" if detail else ""
if condition:
print(f" PASS: {name}{detail_str}")
else:
print(f" FAIL: {name}{detail_str}")
ALL_PASS = False
print("\n=== MedTriage Validation Tests ===\n")
# ββ Health + Tasks ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
r = httpx.get(f"{BASE}/health", timeout=10)
check("Health endpoint returns 200", r.status_code == 200, r.json().get("status"))
r = httpx.get(f"{BASE}/tasks", timeout=10)
tasks = [t["id"] for t in r.json()["tasks"]]
check("3 tasks available", len(tasks) == 3, str(tasks))
check("vital-triage task exists", "vital-triage" in tasks)
check("differential-diagnosis task exists", "differential-diagnosis" in tasks)
check("treatment-safety task exists", "treatment-safety" in tasks)
# ββ Task 1 Tests βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("\nββ Task 1: vital-triage ββ")
r = httpx.post(f"{BASE}/reset", json={"task": "vital-triage"}, timeout=10)
check("Task 1 reset returns 200", r.status_code == 200)
obs = r.json()["observation"]
check("Observation has patient", "patient" in r.json()["observation"])
check("Observation has task_instruction", bool(obs.get("task_instruction")))
check("Observation has vitals", obs["patient"]["vitals"]["heart_rate"] is not None)
# Correct ESI-1
r = httpx.post(f"{BASE}/step", json={"esi_level": 1, "triage_reason": "hypotension severe"}, timeout=10)
check("Step returns 200", r.status_code == 200)
step = r.json()
check("Step returns reward in [0,1]", 0.0 <= step["reward"] <= 1.0, f"reward={step['reward']}")
check("Step done=True after action", step["done"] == True)
# ESI penalty test
httpx.post(f"{BASE}/reset", json={"task": "vital-triage"}, timeout=10)
r = httpx.post(f"{BASE}/step", json={"esi_level": 5, "triage_reason": "seems ok"}, timeout=10)
step = r.json()
check("Wrong ESI gets lower reward", step["reward"] <= 0.5, f"reward={step['reward']}")
# ββ Task 2 Tests βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("\nββ Task 2: differential-diagnosis ββ")
r = httpx.post(f"{BASE}/reset", json={"task": "differential-diagnosis"}, timeout=10)
check("Task 2 reset returns 200", r.status_code == 200)
r = httpx.post(f"{BASE}/step", json={
"diagnoses": ["STEMI", "NSTEMI", "Aortic dissection"],
"red_flags": ["ST elevation", "elevated troponin", "hypotension", "diaphoresis"],
"recommended_tests": ["12-lead ECG", "Troponin", "CXR"]
}, timeout=10)
step = r.json()
check("Task 2 step returns 200", r.status_code == 200)
check("Task 2 reward in [0,1]", 0.0 <= step["reward"] <= 1.0, f"reward={step['reward']}")
# ββ Task 3 Tests βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("\n-- Task 3: treatment-safety --")
r = httpx.post(f"{BASE}/reset", json={"task": "treatment-safety"}, timeout=10)
check("Task 3 reset returns 200", r.status_code == 200)
obs = r.json()["observation"]
check("Diagnosis context provided", "diagnosis_context" in obs.get("context", {}))
# Correct safe treatment β loop until we hit T001 (STEMI patient where aspirin is correct)
t3_correct_pass = False
t3_contra_pass = False
for _ in range(20):
r = httpx.post(f"{BASE}/reset", json={"task": "treatment-safety"}, timeout=10)
pid = r.json()["observation"]["patient"]["patient_id"]
if pid == "T001" and not t3_correct_pass:
r2 = httpx.post(f"{BASE}/step", json={
"diagnosis": "STEMI",
"drug_name": "aspirin",
"dose_mg": 300,
"route": "PO",
"rationale": "dual antiplatelet therapy"
}, timeout=10)
step = r2.json()
check("Task 3 correct drug (aspirin for STEMI) reward >= 0.6", step["reward"] >= 0.6, f"reward={step['reward']}")
t3_correct_pass = True
if pid == "T004" and not t3_contra_pass:
r2 = httpx.post(f"{BASE}/step", json={
"diagnosis": "Status asthmaticus",
"drug_name": "aspirin",
"dose_mg": 300,
"route": "PO",
"rationale": "anti-inflammatory"
}, timeout=10)
step = r2.json()
check("Contraindicated drug (aspirin in NSAID-allergy) returns 0.0", step["reward"] == 0.0, f"reward={step['reward']}")
t3_contra_pass = True
if t3_correct_pass and t3_contra_pass:
break
if not t3_correct_pass:
check("Task 3 correct drug test ran", False, "T001 patient not drawn in 20 tries")
if not t3_contra_pass:
check("Task 3 contraindication test ran", False, "T004 patient not drawn in 20 tries")
# ββ State endpoint ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("\nββ State endpoint ββ")
r = httpx.get(f"{BASE}/state", timeout=10)
check("State returns 200", r.status_code == 200)
state = r.json()
check("State has episode_id", bool(state.get("episode_id")))
check("State has step_count", "step_count" in state)
check("State has cumulative_reward", "cumulative_reward" in state)
# ββ Summary βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print()
if ALL_PASS:
print("β
ALL TESTS PASSED β Environment is working correctly!")
else:
print("β SOME TESTS FAILED β Check output above.")
|