File size: 5,380 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
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
"""
Module: test_environment.py
Purpose: Test core environment behavior — reset, step, state.
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 triage_flow.environment import TriageEnvironment
from models import TriageAction, ActionType, PriorityLevel, InfoField


class TestReset:
    """Test reset() behavior."""

    def test_reset_returns_observation(self):
        """reset() should return a TriageObservation."""
        env = TriageEnvironment()
        obs = env.reset(task_name="basic-triage")
        assert obs.done is False
        assert obs.current_patient is not None
        assert obs.queue_length > 0

    def test_reset_basic_triage_has_3_patients(self):
        """Basic triage should have 3 patients."""
        env = TriageEnvironment()
        obs = env.reset(task_name="basic-triage")
        assert obs.queue_length == 3

    def test_reset_clears_state(self):
        """reset() should clear all previous state."""
        env = TriageEnvironment()
        obs = env.reset(task_name="basic-triage")

        # Take an action
        action = TriageAction(
            action_type=ActionType.ASSIGN_PRIORITY,
            patient_id="P001",
            priority_level=PriorityLevel.IMMEDIATE
        )
        env.step(action)

        # Reset should clear
        obs = env.reset(task_name="basic-triage")
        assert env.state.step_count == 0
        assert env.state.assignments == {}

    def test_reset_all_tasks(self):
        """All 3 tasks should reset successfully."""
        env = TriageEnvironment()
        for task in ["basic-triage", "incomplete-records-triage", "mass-casualty-triage"]:
            obs = env.reset(task_name=task)
            assert obs.done is False
            assert obs.current_patient is not None


class TestStep:
    """Test step() behavior."""

    def test_step_valid_action(self):
        """step() with valid action should return updated observation."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        action = TriageAction(
            action_type=ActionType.ASSIGN_PRIORITY,
            patient_id="P001",
            priority_level=PriorityLevel.IMMEDIATE,
        )
        obs = env.step(action)
        assert obs.reward is not None
        assert obs.step_number == 1

    def test_step_correct_assignment_positive_reward(self):
        """Correct priority assignment should give positive reward."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        # P001 ground truth is IMMEDIATE
        action = TriageAction(
            action_type=ActionType.ASSIGN_PRIORITY,
            patient_id="P001",
            priority_level=PriorityLevel.IMMEDIATE,
        )
        obs = env.step(action)
        assert obs.reward > 0

    def test_step_incorrect_assignment_negative_reward(self):
        """Incorrect priority assignment should give negative or zero reward."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        # P001 ground truth is IMMEDIATE, assigning NON_URGENT is wrong
        action = TriageAction(
            action_type=ActionType.ASSIGN_PRIORITY,
            patient_id="P001",
            priority_level=PriorityLevel.NON_URGENT,
        )
        obs = env.step(action)
        assert obs.reward < 0

    def test_step_counts_increment(self):
        """Step count should increment with each step."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        action = TriageAction(action_type=ActionType.ADVANCE_QUEUE)
        obs = env.step(action)
        assert obs.step_number == 1

        obs = env.step(action)
        assert obs.step_number == 2

    def test_done_when_all_assigned(self):
        """Episode should end when all patients are assigned."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        # Assign all 3 patients
        for pid, priority in [
            ("P001", PriorityLevel.IMMEDIATE),
            ("P002", PriorityLevel.NON_URGENT),
            ("P003", PriorityLevel.URGENT),
        ]:
            action = TriageAction(
                action_type=ActionType.ASSIGN_PRIORITY,
                patient_id=pid,
                priority_level=priority,
            )
            obs = env.step(action)

        assert obs.done is True


class TestState:
    """Test state() behavior."""

    def test_state_serializable(self):
        """state() should return serializable data."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")
        state = env.state
        d = state.model_dump()
        assert isinstance(d, dict)
        assert "patients" in d
        assert "assignments" in d

    def test_state_tracks_assignments(self):
        """state() should reflect assignments made."""
        env = TriageEnvironment()
        env.reset(task_name="basic-triage")

        action = TriageAction(
            action_type=ActionType.ASSIGN_PRIORITY,
            patient_id="P001",
            priority_level=PriorityLevel.IMMEDIATE,
        )
        env.step(action)

        state = env.state
        assert "P001" in state.assignments
        assert state.assignments["P001"] == "immediate"