Team_Daemons / tests /test_reward_function.py
Rudransh-1508
feat: Implemented the engine
8aeb2b8
Raw
History Blame Contribute Delete
1.23 kB
from env.rca_env import RCAEnvironment
def test_correct_declaration_reward_and_done():
env = RCAEnvironment()
env.reset(seed=42, task="easy_linear_chain")
_, reward, done, info = env.step(
{
"action_type": "declare_root_cause",
"service": "Database",
"fault_type": "HTTP_500",
}
)
assert done is True
assert reward > 0.0
assert info["correct"] is True
def test_wrong_declaration_penalty():
env = RCAEnvironment()
env.reset(seed=42, task="easy_linear_chain")
_, reward, done, info = env.step(
{
"action_type": "declare_root_cause",
"service": "Auth_Service",
"fault_type": "Timeout",
}
)
assert done is True
assert reward == -0.3
assert info["correct"] is False
def test_timeout_penalty_when_limit_reached_without_declare():
env = RCAEnvironment()
obs = env.reset(seed=42, task="easy_linear_chain")
reward = 0.0
done = False
for _ in range(obs["step_limit"]):
_, reward, done, _ = env.step(
{"action_type": "query_traces", "service": "API_Gateway", "limit": 50}
)
assert done is True
assert reward <= 0.0