Spaces:
Sleeping
Sleeping
| """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) | |