triage-flow / tests /test_models.py
StrongCapybara's picture
feat: complete TriageFlow environment for Hackathon Submission
0cb452d
Raw
History Blame Contribute Delete
4.19 kB
"""
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"