File size: 7,825 Bytes
f9cf02d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
339abf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f9cf02d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
339abf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
"""
Tests for the BugTriageEnv environment (server/environment.py).

Covers the reset/step/done lifecycle along with reward calculation and
metric tracking correctness.
"""

from __future__ import annotations

import pytest
from models import ActionModel
from server.environment import BugTriageEnv


@pytest.fixture()
def env() -> BugTriageEnv:
    return BugTriageEnv()


@pytest.fixture()
def easy_env(env: BugTriageEnv) -> BugTriageEnv:
    env.reset(task_id="bug_triage_easy", seed=42)
    return env


# ---------------------------------------------------------------------------
# Reset
# ---------------------------------------------------------------------------

class TestReset:
    def test_reset_returns_observation(self, env):
        obs = env.reset(task_id="bug_triage_easy", seed=42)
        assert obs.current_ticket is not None

    def test_reset_sets_steps_to_zero(self, easy_env):
        assert easy_env.steps_used == 0

    def test_reset_episode_not_done(self, easy_env):
        assert easy_env.episode_done is False

    def test_reset_clears_cumulative_reward(self, easy_env):
        assert easy_env.cumulative_reward == 0.0

    def test_reset_all_tasks(self, env):
        for task_id in ("bug_triage_easy", "bug_triage_medium", "bug_triage_hard"):
            obs = env.reset(task_id=task_id, seed=42)
            assert obs.current_ticket is not None

    def test_reset_unknown_task_raises(self, env):
        with pytest.raises((FileNotFoundError, ValueError)):
            env.reset(task_id="nonexistent_task")


# ---------------------------------------------------------------------------
# Step — basic lifecycle
# ---------------------------------------------------------------------------

class TestStep:
    def test_step_before_reset_raises(self, env):
        action = ActionModel(action_type="next_ticket", next_ticket={})
        with pytest.raises(RuntimeError):
            env.step(action)

    def test_step_returns_four_values(self, easy_env):
        action = ActionModel(action_type="next_ticket", next_ticket={})
        result = easy_env.step(action)
        assert len(result) == 4

    def test_step_increments_steps_used(self, easy_env):
        action = ActionModel(action_type="next_ticket", next_ticket={})
        easy_env.step(action)
        assert easy_env.steps_used == 1

    def test_step_reward_in_range(self, easy_env):
        action = ActionModel(action_type="next_ticket", next_ticket={})
        _, reward, _, _ = easy_env.step(action)
        assert 0.0 <= reward.step_reward <= 1.0

    def test_next_ticket_advances_index(self, easy_env):
        initial_ticket_id = easy_env.current_task.tickets[easy_env.current_ticket_index].ticket_id
        action = ActionModel(action_type="next_ticket", next_ticket={})
        obs, _, _, _ = easy_env.step(action)
        if obs.current_ticket:
            assert obs.current_ticket.ticket_id != initial_ticket_id


# ---------------------------------------------------------------------------
# Classify action
# ---------------------------------------------------------------------------

class TestClassifyAction:
    def _classify(self, env, severity, priority, component):
        return ActionModel(
            action_type="classify",
            classify={"severity": severity, "priority": priority, "component": component},
        )

    def test_classify_returns_valid_reward(self, easy_env):
        action = self._classify(easy_env, "sev2", "p2", "ios-app")
        _, reward, _, _ = easy_env.step(action)
        assert 0.0 <= reward.step_reward <= 1.0

    def test_classify_invalid_component_rejected(self, easy_env):
        action = self._classify(easy_env, "sev2", "p2", "nonexistent-component")
        _, _, _, info = easy_env.step(action)
        assert info.get("validation_error") is not None


class TestTriageSemantics:
    def test_mark_duplicate_marks_ticket_triaged(self, easy_env):
        easy_env.current_ticket_index = 3  # BUG-1004 duplicate of BUG-1001
        action = ActionModel(
            action_type="mark_duplicate",
            mark_duplicate={"canonical_ticket_id": "BUG-1001"},
        )
        easy_env.step(action)
        assert easy_env.ticket_states[3]["triaged"] is True

    def test_request_info_marks_ticket_triaged(self, easy_env):
        easy_env.current_ticket_index = 2  # BUG-1003 needs more info
        action = ActionModel(
            action_type="request_info",
            request_info={"info_type": "logs"},
        )
        easy_env.step(action)
        assert easy_env.ticket_states[2]["triaged"] is True


# ---------------------------------------------------------------------------
# Episode done conditions
# ---------------------------------------------------------------------------

class TestEpisodeDone:
    def test_episode_ends_when_budget_exhausted(self, env):
        obs = env.reset(task_id="bug_triage_easy", seed=42)
        budget = env.current_task.step_budget
        action = ActionModel(action_type="next_ticket", next_ticket={})
        done = False
        for _ in range(budget + 5):
            if done:
                break
            _, _, done, _ = env.step(action)
        assert env.episode_done is True

    def test_step_after_done_raises(self, easy_env):
        easy_env.episode_done = True
        action = ActionModel(action_type="next_ticket", next_ticket={})
        with pytest.raises(RuntimeError):
            easy_env.step(action)


# ---------------------------------------------------------------------------
# State
# ---------------------------------------------------------------------------

class TestState:
    def test_state_before_reset_raises(self, env):
        with pytest.raises(RuntimeError):
            env.state()

    def test_state_returns_correct_task_id(self, easy_env):
        state = easy_env.state()
        assert state.current_task_id == "bug_triage_easy"

    def test_state_total_tickets(self, easy_env):
        state = easy_env.state()
        assert state.total_tickets == len(easy_env.current_task.tickets)


# ---------------------------------------------------------------------------
# Partial score
# ---------------------------------------------------------------------------

class TestPartialScore:
    def test_partial_score_zero_before_steps(self, easy_env):
        assert easy_env._calculate_partial_score() == 0.0

    def test_partial_score_in_range_after_step(self, easy_env):
        action = ActionModel(action_type="next_ticket", next_ticket={})
        easy_env.step(action)
        score = easy_env._calculate_partial_score()
        assert 0.0 <= score <= 1.0


class TestTerminalRewards:
    def test_terminal_bonus_helper_returns_bonus_when_all_critical_are_triaged(self, easy_env):
        for index, ticket in enumerate(easy_env.current_task.tickets):
            ground_truth = easy_env.current_task.get_ground_truth(ticket.ticket_id)
            if ground_truth.true_severity in {"sev0", "sev1"}:
                easy_env.ticket_states[index]["triaged"] = True

        adjustment, breakdown = easy_env._terminal_step_adjustment()
        assert adjustment == easy_env.reward_calculator.ALL_CRITICAL_TRIAGED
        assert breakdown["all_critical_triaged_bonus"] == easy_env.reward_calculator.ALL_CRITICAL_TRIAGED

    def test_terminal_penalty_helper_returns_penalty_when_budget_ends_with_critical_remaining(self, easy_env):
        easy_env.steps_used = easy_env.current_task.step_budget

        adjustment, breakdown = easy_env._terminal_step_adjustment()
        assert adjustment == -easy_env.reward_calculator.BUDGET_EXHAUSTED_CRITICAL_REMAINING
        assert (
            breakdown["budget_exhausted_critical_remaining_penalty"]
            == -easy_env.reward_calculator.BUDGET_EXHAUSTED_CRITICAL_REMAINING
        )