logsentinel / tests /test_env.py
Surya-sj's picture
feat: LogSentinel v2 — Multi-Agent SOC War-Room
f7d58be
Raw
History Blame Contribute Delete
33.8 kB
"""Comprehensive tests for LogSentinel v2 — Multi-Agent SOC War-Room."""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import pytest
from fastapi.testclient import TestClient
from server import app
from environment import LogSentinelEnv
from log_generator import (
generate_task1_logs,
generate_task2_logs,
generate_task3_logs,
generate_scenario,
filter_logs_for_role,
ROLE_SOURCE_VISIBILITY,
)
from graders import (
grade_action,
compute_episode_reward,
check_anti_hacking,
)
from models import (
Action,
AgentRole,
CurriculumState,
DifficultyLevel,
EpisodePhase,
GroundTruth,
ScenarioConfig,
)
# ===========================================================================
# 1. Log generator tests (legacy + new)
# ===========================================================================
class TestLogGenerator:
def test_task1_generates_10_logs(self):
logs, gt = generate_task1_logs()
assert len(logs) == 10
assert len(gt.log_classifications) == 10
def test_task1_deterministic(self):
logs1, gt1 = generate_task1_logs(seed=42)
logs2, gt2 = generate_task1_logs(seed=42)
assert [lg.message for lg in logs1] == [lg.message for lg in logs2]
assert gt1.log_classifications == gt2.log_classifications
def test_task2_generates_20_logs(self):
logs, gt = generate_task2_logs()
assert len(logs) == 20
assert len(gt.incidents) == 2
def test_task3_generates_30plus_logs(self):
logs, gt = generate_task3_logs()
assert len(logs) >= 30
assert len(gt.incidents) == 3
def test_task3_has_security_incident(self):
_, gt = generate_task3_logs()
types = [inc["type"] for inc in gt.incidents]
assert "security_breach" in types
def test_scenario_deterministic_with_seed(self):
cfg = ScenarioConfig(num_incidents=2, difficulty=DifficultyLevel.MEDIUM, seed=999)
logs1, gt1 = generate_scenario(cfg)
logs2, gt2 = generate_scenario(cfg)
assert len(logs1) == len(logs2)
assert [lg.message for lg in logs1] == [lg.message for lg in logs2]
assert gt1.incidents == gt2.incidents
def test_scenario_easy_has_one_incident(self):
cfg = ScenarioConfig(num_incidents=1, difficulty=DifficultyLevel.EASY, seed=1)
_, gt = generate_scenario(cfg)
assert len(gt.incidents) == 1
def test_scenario_hard_has_three_incidents(self):
cfg = ScenarioConfig(num_incidents=3, difficulty=DifficultyLevel.HARD, seed=2)
_, gt = generate_scenario(cfg)
assert len(gt.incidents) == 3
def test_scenario_noise_adds_normal_logs(self):
cfg = ScenarioConfig(
num_incidents=1, difficulty=DifficultyLevel.EASY,
confounding_noise_ratio=0.5, seed=3
)
logs, gt = generate_scenario(cfg)
normal_count = sum(1 for c in gt.log_classifications.values() if c == "normal")
assert normal_count > 0
def test_observability_drops_logs(self):
cfg_full = ScenarioConfig(
num_incidents=1, difficulty=DifficultyLevel.EASY,
observability_quality=1.0, confounding_noise_ratio=0.3, seed=77
)
cfg_low = ScenarioConfig(
num_incidents=1, difficulty=DifficultyLevel.EASY,
observability_quality=0.5, confounding_noise_ratio=0.3, seed=77
)
logs_full, _ = generate_scenario(cfg_full)
logs_low, _ = generate_scenario(cfg_low)
# Low observability should have <= full observability log count
assert len(logs_low) <= len(logs_full)
# ===========================================================================
# 2. Role-based partial observability tests
# ===========================================================================
class TestPartialObservability:
def test_filter_logs_for_app_sre(self):
logs, _ = generate_task3_logs(seed=42)
visible = filter_logs_for_role(logs, "app_sre")
allowed = set(ROLE_SOURCE_VISIBILITY["app_sre"])
for log in visible:
assert log.source in allowed
def test_filter_logs_for_db_sre(self):
logs, _ = generate_task3_logs(seed=42)
visible = filter_logs_for_role(logs, "db_sre")
allowed = set(ROLE_SOURCE_VISIBILITY["db_sre"])
for log in visible:
assert log.source in allowed
def test_filter_logs_for_security_analyst(self):
logs, _ = generate_task3_logs(seed=42)
visible = filter_logs_for_role(logs, "security_analyst")
allowed = set(ROLE_SOURCE_VISIBILITY["security_analyst"])
for log in visible:
assert log.source in allowed
def test_incident_commander_sees_more_logs_than_db_sre(self):
logs, _ = generate_task3_logs(seed=42)
cmd_visible = filter_logs_for_role(logs, "incident_commander")
db_visible = filter_logs_for_role(logs, "db_sre")
# IC sees all sources; DB SRE sees subset
assert len(cmd_visible) >= len(db_visible)
def test_env_step_returns_role_filtered_observation(self):
env = LogSentinelEnv()
result = env.reset(task_name="soc_warroom_easy", agent_role="db_sre", seed=10)
obs = result["observation"]
assert obs["agent_role"] == "db_sre"
# Logs visible to db_sre should only be from db_sre sources
allowed = set(ROLE_SOURCE_VISIBILITY["db_sre"])
for log in obs["log_entries"]:
assert log["source"] in allowed
def test_shared_board_updated_after_handoff(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_medium", agent_role="app_sre", seed=42)
env.step({
"action_type": "propose_incident",
"agent_role": "app_sre",
"incident_type": "outage",
"evidence_indices": [0, 1],
})
result = env.step({
"action_type": "request_handoff",
"agent_role": "app_sre",
"handoff_to": "db_sre",
"handoff_note": "Possible replication lag causing outage. Check pg replica.",
})
obs = result["observation"]
# Shared board should have the handoff recorded
assert obs.get("shared_board") is not None
board = obs["shared_board"]
assert any("handoff" in k for k in board)
# ===========================================================================
# 3. Phase transition tests
# ===========================================================================
class TestPhaseTransitions:
def test_initial_phase_is_detect(self):
env = LogSentinelEnv()
result = env.reset(task_name="soc_warroom_easy", seed=42)
assert result["observation"]["current_phase"] == "detect"
def test_propose_incident_triggers_triage_phase(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=42)
env.step({
"action_type": "propose_incident",
"agent_role": "app_sre",
"incident_type": "outage",
"evidence_indices": [0],
})
# Now take a triage action → should advance to triage
result = env.step({
"action_type": "assign_severity",
"agent_role": "app_sre",
"incident_type": "outage",
"severity": "P1",
})
obs = result["observation"]
assert obs["current_phase"] in ("triage", "detect") # may advance
def test_state_reflects_current_phase(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_medium", seed=10)
state = env.state
assert "current_phase" in state
assert state["current_phase"] in ("detect", "triage", "mitigate", "verify", "final_report")
def test_phase_tracks_in_world_state(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=1)
state = env.state
assert state["world_state"]["phase"] == "detect"
def test_execute_mitigation_advances_phase(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=5)
# Detect
env.step({"action_type": "propose_incident", "agent_role": "app_sre",
"incident_type": "outage", "evidence_indices": [0]})
# Triage
env.step({"action_type": "assign_severity", "agent_role": "app_sre",
"incident_type": "outage", "severity": "P2"})
# Mitigate → should advance phase
result = env.step({
"action_type": "execute_mitigation",
"agent_role": "app_sre",
"mitigation_id": "restart_app_servers",
"evidence_indices": [0, 1],
})
state = env.state
assert state["current_phase"] in ("triage", "mitigate", "verify")
# ===========================================================================
# 4. Delayed mitigation effects tests
# ===========================================================================
class TestDelayedMitigationEffects:
def test_mitigation_improves_service_health(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=7)
state_before = env.state
health_before = state_before["world_state"]["service_health"]
env.step({"action_type": "propose_incident", "agent_role": "app_sre",
"incident_type": "resource_exhaustion", "evidence_indices": [0]})
env.step({
"action_type": "execute_mitigation",
"agent_role": "db_sre",
"mitigation_id": "scale_connection_pool",
"evidence_indices": [0, 1, 2],
})
state_after = env.state
health_after = state_after["world_state"]["service_health"]
# Health should not decrease after relevant mitigation
assert health_after >= health_before - 0.01 # allow tiny float drift
def test_verify_recovery_sets_containment(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=8)
env.step({"action_type": "propose_incident", "agent_role": "security_analyst",
"incident_type": "security_breach", "evidence_indices": [0]})
env.step({"action_type": "execute_mitigation", "agent_role": "security_analyst",
"mitigation_id": "block_attacker_ip", "evidence_indices": [0]})
env.step({
"action_type": "verify_recovery",
"agent_role": "incident_commander",
"evidence_indices": [0, 1],
})
state = env.state
assert state["world_state"]["containment_status"] is True
def test_unsafe_mitigation_tracked(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=9)
env.step({
"action_type": "execute_mitigation",
"agent_role": "app_sre",
"mitigation_id": "restart_everything",
# No evidence_indices → unsafe
})
state = env.state
assert state["anti_hacking"]["unsafe_mitigations"] >= 1
# ===========================================================================
# 5. Reward component correctness tests
# ===========================================================================
class TestRewardComponents:
def _make_gt(self) -> GroundTruth:
return GroundTruth(
log_classifications={0: "error", 1: "warning", 2: "security"},
incidents=[
{"type": "outage", "severity": "P1", "correlated_indices": [0, 1], "description": "test"},
{"type": "security_breach", "severity": "P1", "correlated_indices": [2], "description": "test"},
],
)
def test_correct_classification_rewards(self):
_, gt = generate_task1_logs()
idx = 0
expected_class = gt.log_classifications[idx]
action = Action(action_type="classify_log", target_log_indices=[idx], classification=expected_class)
reward = grade_action(action, gt)
assert reward > 0
def test_wrong_classification_no_reward(self):
_, gt = generate_task1_logs()
idx = 0
wrong_class = "security" if gt.log_classifications[idx] != "security" else "normal"
action = Action(action_type="classify_log", target_log_indices=[idx], classification=wrong_class)
reward = grade_action(action, gt)
assert reward == 0.0
def test_incident_detection_reward(self):
_, gt = generate_task2_logs()
itype = gt.incidents[0]["type"]
action = Action(action_type="detect_incident", incident_type=itype)
reward = grade_action(action, gt)
assert reward > 0
def test_correlation_reward(self):
_, gt = generate_task2_logs()
indices = gt.incidents[0]["correlated_indices"]
action = Action(action_type="correlate_logs", correlated_indices=indices)
reward = grade_action(action, gt)
assert reward > 0
def test_report_grading(self):
_, gt = generate_task2_logs()
action = Action(
action_type="submit_report",
report={
"incidents": [{"type": "resource_exhaustion"}] * 2,
"severity": "P2",
"summary": "Database connection exhaustion causing cascading failures and disk issues",
},
)
reward = grade_action(action, gt)
assert reward > 0
def test_unknown_action_zero_reward(self):
_, gt = generate_task1_logs()
# action_type validation will raise; catch it
with pytest.raises(Exception):
Action(action_type="totally_unknown_xyz")
def test_observe_logs_zero_reward(self):
_, gt = generate_task1_logs()
action = Action(action_type="observe_logs", agent_role="app_sre")
assert grade_action(action, gt) == 0.0
def test_episode_reward_breakdown_structure(self):
gt = self._make_gt()
rb = compute_episode_reward(
detected_incident_types=["outage"],
severity_votes=[("outage", "P1")],
total_steps=10, max_steps=30,
noop_count=2, handoff_count=1, useful_handoffs=1, redundant_actions=0,
service_restored=True, breach_contained=True,
ground_truth=gt,
)
assert 0.0 <= rb.r_outcome <= 1.0
assert 0.0 <= rb.r_detection_f1 <= 1.0
assert 0.0 <= rb.r_severity_accuracy <= 1.0
assert 0.0 <= rb.r_efficiency <= 1.0
assert -1.0 <= rb.total <= 1.0
def test_episode_reward_r_outcome_full(self):
gt = self._make_gt()
rb = compute_episode_reward(
detected_incident_types=["outage", "security_breach"],
severity_votes=[("outage", "P1"), ("security_breach", "P1")],
total_steps=5, max_steps=30, noop_count=0,
handoff_count=2, useful_handoffs=2, redundant_actions=0,
service_restored=True, breach_contained=True,
ground_truth=gt,
)
assert rb.r_outcome == pytest.approx(1.0)
assert rb.total > 0.5
def test_mitigation_reward_positive(self):
_, gt = generate_task2_logs()
action = Action(
action_type="execute_mitigation",
agent_role="db_sre",
mitigation_id="scale_connection_pool",
evidence_indices=[0, 1],
)
reward = grade_action(action, gt)
assert reward >= 0.0 # 0 for no-match tasks, positive for relevant
def test_handoff_reward_positive(self):
_, gt = generate_task3_logs()
action = Action(
action_type="request_handoff",
agent_role="app_sre",
handoff_to="db_sre",
handoff_note="Observed high replication lag, needs DB SRE investigation.",
)
reward = grade_action(action, gt)
assert reward >= 0.05
def test_severity_near_miss_partial_credit(self):
_, gt = generate_task3_logs()
# P1 incident exists (outage), vote P2 → partial credit
action = Action(action_type="assign_severity", agent_role="app_sre",
incident_type="outage", severity="P2")
reward = grade_action(action, gt)
assert reward > 0 # partial credit
def test_severity_exact_match_full_credit(self):
_, gt = generate_task3_logs()
action = Action(action_type="assign_severity", agent_role="app_sre",
incident_type="outage", severity="P1")
reward = grade_action(action, gt)
assert reward >= 0.15
# ===========================================================================
# 6. Anti-hacking penalty tests
# ===========================================================================
class TestAntiHacking:
def test_duplicate_incident_proposal_penalty(self):
action = Action(action_type="propose_incident", agent_role="app_sre", incident_type="outage")
history = ["propose:outage", "propose:outage"]
penalty, reason = check_anti_hacking(action, history, ["outage"])
assert penalty > 0
assert "duplicate" in reason
def test_no_penalty_first_proposal(self):
action = Action(action_type="propose_incident", agent_role="app_sre", incident_type="outage")
penalty, reason = check_anti_hacking(action, [], [])
assert penalty == 0.0
assert reason == "ok"
def test_repeated_noop_penalty(self):
action = Action(action_type="observe_logs", agent_role="app_sre")
history = ["observe_logs", "observe_logs", "observe_logs"]
penalty, reason = check_anti_hacking(action, history, [])
assert penalty > 0
assert "noop" in reason
def test_unsafe_mitigation_penalty(self):
action = Action(action_type="execute_mitigation", agent_role="app_sre",
mitigation_id="restart_db") # no evidence_indices
penalty, reason = check_anti_hacking(action, [], [])
assert penalty >= 0.1
assert "unsafe_mitigation" in reason
def test_safe_mitigation_no_penalty(self):
action = Action(action_type="execute_mitigation", agent_role="app_sre",
mitigation_id="restart_db", evidence_indices=[0, 1])
penalty, reason = check_anti_hacking(action, [], [])
assert penalty == 0.0
def test_report_before_detection_penalty(self):
action = Action(action_type="submit_report",
report={"incidents": [], "severity": "P4", "summary": "Nothing found"})
penalty, reason = check_anti_hacking(action, [], []) # no detected types
assert penalty >= 0.1
assert "report_before_detection" in reason
def test_env_unsafe_mitigation_tracked_in_state(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=42)
env.step({
"action_type": "execute_mitigation",
"agent_role": "app_sre",
"mitigation_id": "blind_restart",
})
state = env.state
assert state["anti_hacking"]["unsafe_mitigations"] == 1
def test_duplicate_proposal_tracked_in_state(self):
env = LogSentinelEnv()
env.reset(task_name="soc_warroom_easy", seed=42)
for _ in range(3):
env.step({"action_type": "propose_incident", "agent_role": "app_sre",
"incident_type": "outage", "evidence_indices": [0]})
state = env.state
assert state["anti_hacking"]["duplicate_proposals"] >= 1
# ===========================================================================
# 7. Curriculum difficulty adjustment tests
# ===========================================================================
class TestCurriculumDifficulty:
def test_initial_difficulty_is_easy(self):
curriculum = CurriculumState()
assert curriculum.current_difficulty == DifficultyLevel.EASY
def test_high_success_promotes_difficulty(self):
curriculum = CurriculumState()
# Record many successes
for _ in range(10):
curriculum.record(0.9)
new_diff = curriculum.adjust_difficulty()
assert new_diff == DifficultyLevel.MEDIUM
def test_low_success_keeps_easy(self):
curriculum = CurriculumState()
for _ in range(10):
curriculum.record(0.1)
new_diff = curriculum.adjust_difficulty()
assert new_diff == DifficultyLevel.EASY # can't go below easy
def test_medium_high_success_promotes_to_hard(self):
curriculum = CurriculumState(current_difficulty=DifficultyLevel.MEDIUM)
for _ in range(10):
curriculum.record(0.9)
new_diff = curriculum.adjust_difficulty()
assert new_diff == DifficultyLevel.HARD
def test_medium_low_success_demotes_to_easy(self):
curriculum = CurriculumState(current_difficulty=DifficultyLevel.MEDIUM)
for _ in range(10):
curriculum.record(0.1)
new_diff = curriculum.adjust_difficulty()
assert new_diff == DifficultyLevel.EASY
def test_curriculum_adjusts_scenario_difficulty(self):
env = LogSentinelEnv()
# Simulate 10 successful episodes
for _ in range(10):
env._curriculum.record(0.9)
env._curriculum.adjust_difficulty()
# Reset adaptive curriculum → should generate medium difficulty scenario
result = env.reset(task_name="adaptive_curriculum", seed=42)
state = env.state
assert state["curriculum_difficulty"] in ("easy", "medium", "hard")
def test_curriculum_window_size_capped(self):
curriculum = CurriculumState(window_size=5)
for _ in range(20):
curriculum.record(0.9)
assert len(curriculum.success_history) <= 5
# ===========================================================================
# 8. Deterministic generation with seed tests
# ===========================================================================
class TestDeterministicGeneration:
def test_reset_same_seed_same_logs(self):
env = LogSentinelEnv()
r1 = env.reset(task_name="soc_warroom_easy", seed=42)
logs1 = [lg["message"] for lg in r1["observation"]["log_entries"]]
r2 = env.reset(task_name="soc_warroom_easy", seed=42)
logs2 = [lg["message"] for lg in r2["observation"]["log_entries"]]
assert logs1 == logs2
def test_reset_different_seeds_different_logs(self):
env = LogSentinelEnv()
r1 = env.reset(task_name="soc_warroom_medium", seed=1)
logs1 = [lg["message"] for lg in r1["observation"]["log_entries"]]
r2 = env.reset(task_name="soc_warroom_medium", seed=2)
logs2 = [lg["message"] for lg in r2["observation"]["log_entries"]]
# Very likely to differ
assert logs1 != logs2
def test_scenario_generator_seed_reproducible(self):
cfg = ScenarioConfig(num_incidents=2, difficulty=DifficultyLevel.MEDIUM, seed=12345)
logs_a, gt_a = generate_scenario(cfg)
logs_b, gt_b = generate_scenario(cfg)
assert [lg.timestamp for lg in logs_a] == [lg.timestamp for lg in logs_b]
assert len(gt_a.incidents) == len(gt_b.incidents)
# ===========================================================================
# 9. Reset / step / state API compatibility tests
# ===========================================================================
class TestEnvironment:
def setup_method(self):
self.env = LogSentinelEnv()
def test_reset_returns_observation(self):
result = self.env.reset(task_name="log_classification")
assert "observation" in result
assert "reward" in result
assert "done" in result
assert result["done"] is False
assert result["reward"] is None
def test_reset_has_log_entries(self):
result = self.env.reset(task_name="log_classification")
obs = result["observation"]
assert len(obs["log_entries"]) == 10
def test_step_returns_reward(self):
self.env.reset(task_name="log_classification")
result = self.env.step({
"action_type": "classify_log",
"target_log_indices": [0],
"classification": "normal",
})
assert "reward" in result
assert isinstance(result["reward"], float)
def test_step_returns_info(self):
self.env.reset(task_name="soc_warroom_easy", seed=42)
result = self.env.step({
"action_type": "observe_logs",
"agent_role": "app_sre",
})
assert "info" in result
assert "phase" in result["info"]
def test_submit_report_ends_episode(self):
self.env.reset(task_name="log_classification")
result = self.env.step({
"action_type": "submit_report",
"report": {"incidents": [], "severity": "P4", "summary": "No incidents found"},
})
assert result["done"] is True
def test_max_steps_ends_episode(self):
self.env.reset(task_name="log_classification")
for _ in range(20):
result = self.env.step({
"action_type": "classify_log",
"target_log_indices": [0],
"classification": "normal",
})
if result["done"]:
break
assert result["done"] is True
def test_state_tracks_progress(self):
self.env.reset(task_name="log_classification")
state = self.env.state
assert state["step_count"] == 0
assert state["episode_id"] is not None
self.env.step({"action_type": "classify_log", "target_log_indices": [0], "classification": "normal"})
state = self.env.state
assert state["step_count"] == 1
def test_state_has_world_state(self):
self.env.reset(task_name="soc_warroom_easy", seed=1)
state = self.env.state
assert "world_state" in state
ws = state["world_state"]
assert "service_health" in ws
assert "error_rate" in ws
def test_state_has_reward_breakdown_after_done(self):
self.env.reset(task_name="soc_warroom_easy", seed=1)
self.env.step({"action_type": "propose_incident", "agent_role": "app_sre",
"incident_type": "outage", "evidence_indices": [0]})
self.env.step({
"action_type": "submit_joint_report",
"agent_role": "incident_commander",
"report": {"incidents": [{"type": "outage"}], "severity": "P1",
"summary": "Outage detected, mitigation applied"},
})
state = self.env.state
rb = state["reward_breakdown"]
assert "r_outcome" in rb
assert "r_detection_f1" in rb
assert "total" in rb
def test_all_legacy_tasks_work(self):
for task in ["log_classification", "incident_detection", "full_triage"]:
result = self.env.reset(task_name=task)
assert result["done"] is False
assert len(result["observation"]["log_entries"]) > 0
def test_all_soc_tasks_work(self):
for task in ["soc_warroom_easy", "soc_warroom_medium", "soc_warroom_hard"]:
result = self.env.reset(task_name=task, seed=42)
assert result["done"] is False
obs = result["observation"]
assert "current_phase" in obs
assert obs["current_phase"] == "detect"
def test_invalid_action_type_returns_negative_reward(self):
self.env.reset(task_name="soc_warroom_easy", seed=42)
result = self.env.step({"action_type": "not_a_real_action"})
# Should return negative reward due to validation error
assert result["reward"] is not None
assert result["reward"] <= 0
def test_episode_after_done_returns_zero_reward(self):
self.env.reset(task_name="log_classification")
self.env.step({"action_type": "submit_report",
"report": {"incidents": [], "severity": "P4", "summary": "done"}})
result = self.env.step({"action_type": "classify_log",
"target_log_indices": [0], "classification": "normal"})
assert result["reward"] == 0.0
assert result["done"] is True
# ===========================================================================
# 10. FastAPI endpoint integration tests
# ===========================================================================
class TestServer:
def setup_method(self):
self.client = TestClient(app)
def test_health(self):
r = self.client.get("/health")
assert r.status_code == 200
assert r.json()["status"] == "healthy"
def test_tasks_include_soc_tasks(self):
r = self.client.get("/tasks")
assert r.status_code == 200
tasks = r.json()["tasks"]
names = {t["name"] for t in tasks}
assert "log_classification" in names
assert "soc_warroom_easy" in names
assert "soc_warroom_hard" in names
assert "adaptive_curriculum" in names
def test_reset(self):
r = self.client.post("/reset", json={})
assert r.status_code == 200
data = r.json()
assert "observation" in data
assert data["done"] is False
def test_reset_with_legacy_task(self):
r = self.client.post("/reset", json={"task_name": "incident_detection"})
assert r.status_code == 200
obs = r.json()["observation"]
assert len(obs["log_entries"]) == 20
def test_reset_with_soc_task(self):
r = self.client.post("/reset", json={"task_name": "soc_warroom_easy", "seed": 42})
assert r.status_code == 200
data = r.json()
assert data["done"] is False
assert data["observation"]["current_phase"] == "detect"
def test_step(self):
self.client.post("/reset", json={})
r = self.client.post("/step", json={
"action": {
"action_type": "classify_log",
"target_log_indices": [0],
"classification": "normal",
}
})
assert r.status_code == 200
data = r.json()
assert "reward" in data
assert "done" in data
assert "info" in data
def test_state(self):
self.client.post("/reset", json={})
r = self.client.get("/state")
assert r.status_code == 200
state = r.json()
assert "episode_id" in state
assert "step_count" in state
assert "world_state" in state
def test_full_legacy_episode(self):
self.client.post("/reset", json={"task_name": "log_classification"})
r = self.client.post("/step", json={
"action": {"action_type": "classify_log", "target_log_indices": [0, 1, 2],
"classification": "normal"}
})
assert r.status_code == 200
assert r.json()["done"] is False
r = self.client.post("/step", json={
"action": {"action_type": "submit_report",
"report": {"incidents": [], "severity": "P4",
"summary": "Classified logs, no major incidents"}}
})
assert r.status_code == 200
assert r.json()["done"] is True
def test_full_soc_episode(self):
self.client.post("/reset", json={"task_name": "soc_warroom_easy", "seed": 42})
# Detect
self.client.post("/step", json={"action": {
"action_type": "propose_incident",
"agent_role": "app_sre",
"incident_type": "resource_exhaustion",
"evidence_indices": [0, 1],
}})
# Triage
self.client.post("/step", json={"action": {
"action_type": "assign_severity",
"agent_role": "app_sre",
"incident_type": "resource_exhaustion",
"severity": "P2",
}})
# Mitigate
self.client.post("/step", json={"action": {
"action_type": "execute_mitigation",
"agent_role": "db_sre",
"mitigation_id": "scale_connection_pool",
"evidence_indices": [0, 1],
}})
# Verify
self.client.post("/step", json={"action": {
"action_type": "verify_recovery",
"agent_role": "incident_commander",
"evidence_indices": [0],
}})
# Final report
r = self.client.post("/step", json={"action": {
"action_type": "submit_joint_report",
"agent_role": "incident_commander",
"report": {
"incidents": [{"type": "resource_exhaustion", "severity": "P2"}],
"severity": "P2",
"summary": "DB connection pool exhaustion detected and mitigated. Service restored.",
},
}})
assert r.status_code == 200
assert r.json()["done"] is True
def test_step_with_role_in_action(self):
self.client.post("/reset", json={"task_name": "soc_warroom_medium", "seed": 10})
r = self.client.post("/step", json={"action": {
"action_type": "observe_logs",
"agent_role": "security_analyst",
}})
assert r.status_code == 200
obs = r.json()["observation"]
assert obs["agent_role"] == "security_analyst"
def test_state_has_curriculum_difficulty(self):
self.client.post("/reset", json={"task_name": "adaptive_curriculum", "seed": 42})
r = self.client.get("/state")
assert r.status_code == 200
state = r.json()
assert "curriculum_difficulty" in state
assert state["curriculum_difficulty"] in ("easy", "medium", "hard")