triage-flow / tests /run_integration_test.py
StrongCapybara's picture
feat: complete TriageFlow environment for Hackathon Submission
0cb452d
Raw
History Blame Contribute Delete
5.83 kB
"""Quick integration test for the TriageFlow environment."""
import sys
sys.path.insert(0, '.')
from triage_flow.environment import TriageEnvironment
from triage_flow.graders import grade_task
from models import TriageAction, ActionType, PriorityLevel, InfoField
def test_basic_triage():
env = TriageEnvironment()
obs = env.reset(task_name="basic-triage")
pid = obs.current_patient["patient_id"] if obs.current_patient else "none"
print(f"Reset OK - patients: {obs.queue_length}, current: {pid}")
# Correct assignments for all 3 patients
for pid, p in [("P001", PriorityLevel.IMMEDIATE), ("P002", PriorityLevel.NON_URGENT), ("P003", PriorityLevel.URGENT)]:
action = TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p)
obs = env.step(action)
print(f" Assigned {pid}={p.value} -> reward={obs.reward}, done={obs.done}")
state = env.state
print(f"State: assignments={state.assignments}, cleared={state.queue_cleared}")
score = grade_task("basic-triage", state.model_dump())
print(f"Grader score: {score}")
assert score == 1.0, f"Expected 1.0, got {score}"
print("BASIC TRIAGE: PASSED\n")
def test_wrong_assignments():
env = TriageEnvironment()
env.reset(task_name="basic-triage")
for pid, p in [("P001", PriorityLevel.NON_URGENT), ("P002", PriorityLevel.IMMEDIATE), ("P003", PriorityLevel.NON_URGENT)]:
env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
score = grade_task("basic-triage", env.state.model_dump())
print(f"All wrong score: {score}")
assert score == 0.0, f"Expected 0.0, got {score}"
print("WRONG ASSIGNMENTS: PASSED\n")
def test_incomplete_records():
env = TriageEnvironment()
obs = env.reset(task_name="incomplete-records-triage")
pid = obs.current_patient["patient_id"] if obs.current_patient else "none"
print(f"Incomplete records - patients: {obs.queue_length}, current: {pid}")
print(f" Missing fields: {obs.missing_fields}")
# Request info for P101 (vitals missing)
action = TriageAction(action_type=ActionType.REQUEST_INFO, patient_id="P101", info_field=InfoField.VITALS)
obs = env.step(action)
print(f" Info request -> reward={obs.reward}")
# Assign correct priorities
for pid, p in [("P101", PriorityLevel.IMMEDIATE), ("P102", PriorityLevel.LESS_URGENT),
("P103", PriorityLevel.URGENT), ("P104", PriorityLevel.NON_URGENT)]:
obs = env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
score = grade_task("incomplete-records-triage", env.state.model_dump())
print(f"Incomplete records score: {score}")
assert 0.5 <= score <= 1.0, f"Expected good score, got {score}"
print("INCOMPLETE RECORDS: PASSED\n")
def test_mass_casualty():
env = TriageEnvironment()
obs = env.reset(task_name="mass-casualty-triage")
print(f"Mass casualty - patients: {obs.queue_length}")
# Assign all 8 patients with correct priorities
correct = [
("P201", PriorityLevel.IMMEDIATE), ("P202", PriorityLevel.IMMEDIATE),
("P203", PriorityLevel.LESS_URGENT), ("P204", PriorityLevel.LESS_URGENT),
("P205", PriorityLevel.URGENT), ("P206", PriorityLevel.NON_URGENT),
("P207", PriorityLevel.URGENT), ("P208", PriorityLevel.NON_URGENT),
]
for pid, p in correct:
obs = env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
score = grade_task("mass-casualty-triage", env.state.model_dump())
print(f"Mass casualty perfect score: {score}")
assert score >= 0.8, f"Expected high score, got {score}"
print("MASS CASUALTY: PASSED\n")
def test_grader_variance():
"""Different behaviors must produce different scores."""
scores = set()
# Perfect
env = TriageEnvironment()
env.reset(task_name="basic-triage")
for pid, p in [("P001", PriorityLevel.IMMEDIATE), ("P002", PriorityLevel.NON_URGENT), ("P003", PriorityLevel.URGENT)]:
env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
scores.add(grade_task("basic-triage", env.state.model_dump()))
# All wrong
env = TriageEnvironment()
env.reset(task_name="basic-triage")
for pid, p in [("P001", PriorityLevel.NON_URGENT), ("P002", PriorityLevel.IMMEDIATE), ("P003", PriorityLevel.NON_URGENT)]:
env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
scores.add(grade_task("basic-triage", env.state.model_dump()))
# Partial
env = TriageEnvironment()
env.reset(task_name="basic-triage")
for pid, p in [("P001", PriorityLevel.IMMEDIATE), ("P002", PriorityLevel.IMMEDIATE), ("P003", PriorityLevel.IMMEDIATE)]:
env.step(TriageAction(action_type=ActionType.ASSIGN_PRIORITY, patient_id=pid, priority_level=p))
scores.add(grade_task("basic-triage", env.state.model_dump()))
# No actions
env = TriageEnvironment()
env.reset(task_name="basic-triage")
scores.add(grade_task("basic-triage", env.state.model_dump()))
print(f"Grader variance - unique scores: {scores}")
assert len(scores) >= 3, f"Expected 3+ different scores, got {scores}"
print("GRADER VARIANCE: PASSED\n")
def test_server_app_import():
"""Test that the server app can be imported."""
from server.app import app
print(f"Server app imported: {app.title}")
print("SERVER IMPORT: PASSED\n")
if __name__ == "__main__":
test_basic_triage()
test_wrong_assignments()
test_incomplete_records()
test_mass_casualty()
test_grader_variance()
test_server_app_import()
print("=" * 60)
print("ALL INTEGRATION TESTS PASSED!")
print("=" * 60)