Spaces:
Sleeping
Sleeping
File size: 4,194 Bytes
0cb452d | 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 130 | """
Module: test_models.py
Purpose: Test all Pydantic models and enums validate correctly.
Part of: Medical Triage Assistant — OpenEnv Round 1
Author: Team Squirrel
"""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import pytest
from models import (
PriorityLevel,
ActionType,
InfoField,
PatientRecord,
TriageAction,
TriageObservation,
TriageState,
)
class TestEnums:
"""Test all enum definitions."""
def test_priority_levels(self):
"""Verify all 4 priority levels exist and have correct values."""
assert PriorityLevel.IMMEDIATE.value == "immediate"
assert PriorityLevel.URGENT.value == "urgent"
assert PriorityLevel.LESS_URGENT.value == "less_urgent"
assert PriorityLevel.NON_URGENT.value == "non_urgent"
def test_action_types(self):
"""Verify all 5 action types exist."""
assert len(ActionType) == 5
assert ActionType.ASSIGN_PRIORITY.value == "assign_priority"
assert ActionType.REQUEST_INFO.value == "request_info"
assert ActionType.ESCALATE.value == "escalate"
assert ActionType.DEFER.value == "defer"
assert ActionType.ADVANCE_QUEUE.value == "advance_queue"
def test_info_fields(self):
"""Verify all 5 info fields exist."""
assert len(InfoField) == 5
assert InfoField.VITALS.value == "vitals"
class TestTriageAction:
"""Test action model validation."""
def test_valid_assign_priority(self):
"""Test creating a valid assign_priority action."""
action = TriageAction(
action_type=ActionType.ASSIGN_PRIORITY,
patient_id="P001",
priority_level=PriorityLevel.IMMEDIATE,
)
assert action.action_type == ActionType.ASSIGN_PRIORITY
assert action.patient_id == "P001"
assert action.priority_level == PriorityLevel.IMMEDIATE
def test_valid_request_info(self):
"""Test creating a valid request_info action."""
action = TriageAction(
action_type=ActionType.REQUEST_INFO,
patient_id="P001",
info_field=InfoField.VITALS,
)
assert action.action_type == ActionType.REQUEST_INFO
assert action.info_field == InfoField.VITALS
def test_valid_advance_queue(self):
"""Test creating a valid advance_queue action."""
action = TriageAction(
action_type=ActionType.ADVANCE_QUEUE,
)
assert action.action_type == ActionType.ADVANCE_QUEUE
class TestTriageObservation:
"""Test observation model."""
def test_default_observation(self):
"""Test observation with default values."""
obs = TriageObservation(done=False, reward=None)
assert obs.done is False
assert obs.reward is None
assert obs.queue_length == 0
def test_full_observation(self):
"""Test observation with all fields."""
obs = TriageObservation(
done=False,
reward=0.50,
current_patient={"patient_id": "P001", "age": 62},
queue_length=3,
queue_position=1,
missing_fields=["vitals"],
previous_action_feedback="Episode started",
step_number=1,
task_name="basic-triage",
)
assert obs.current_patient["patient_id"] == "P001"
assert obs.queue_length == 3
assert obs.missing_fields == ["vitals"]
class TestTriageState:
"""Test state model."""
def test_default_state(self):
"""Test state with default values."""
state = TriageState()
assert state.step_count == 0
assert state.task_name == ""
assert state.patients == []
assert state.assignments == {}
def test_serializable(self):
"""Test state can be serialized to dict."""
state = TriageState(
task_name="basic-triage",
patients=[{"patient_id": "P001"}],
assignments={"P001": "immediate"},
)
d = state.model_dump()
assert d["task_name"] == "basic-triage"
assert d["assignments"]["P001"] == "immediate"
|