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