tek-wizard/devops-incident-responder
2
1from __future__ import annotations2 3import argparse4import json5import random6import sys7from pathlib import Path8 9PROJECT_ROOT = Path(__file__).resolve().parents[1]10if str(PROJECT_ROOT) not in sys.path:11 sys.path.insert(0, str(PROJECT_ROOT))12 13from server.environment import DevOpsEnv14from server.models import IncidentAction15from server.seed_sets import get_seed_set16from server.tasks import TASKS17from scripts.heuristic_policy import choose_action as heuristic_action18from scripts.metrics import aggregate_results19from scripts.random_policy import choose_action as random_action20 21 22def _policy_fn(name: str):23 if name == "heuristic":24 return heuristic_action25 if name == "random":26 return random_action27 raise ValueError(f"Unknown policy '{name}'")28 29 30def run_eval(policy_name: str, seed_set_name: str) -> dict:31 policy = _policy_fn(policy_name)32 seeds = get_seed_set(seed_set_name)33 results = []34 35 for task_id in TASKS:36 for seed in seeds:37 env = DevOpsEnv()38 observation = env.reset_episode(task_id=task_id, seed=seed).model_dump()39 history: list[dict[str, str]] = []40 rng = random.Random(seed)41 42 while not env.done:43 action = policy(observation, history, rng=rng)44 history.append(action)45 observation_model, _, done, _ = env.step_episode(IncidentAction(**action))46 observation = observation_model.model_dump()47 if done:48 break49 50 grade = env.grade().model_dump()51 grade.update({"seed": seed, "task_id": task_id})52 results.append(grade)53 54 return {"policy": policy_name, "seed_set": seed_set_name, "summary": aggregate_results(results), "episodes": results}55 56 57def main() -> None:58 parser = argparse.ArgumentParser(description="Evaluate a baseline policy against the seeded incident environment.")59 parser.add_argument("--policy", choices=["heuristic", "random"], default="heuristic")60 parser.add_argument("--seed-set", choices=["train", "val", "holdout"], default="val")61 parser.add_argument("--output", type=Path, default=None)62 args = parser.parse_args()63 64 payload = run_eval(args.policy, args.seed_set)65 rendered = json.dumps(payload, indent=2)66 if args.output is not None:67 args.output.write_text(rendered + "\n", encoding="utf-8")68 print(rendered)69 70 71if __name__ == "__main__":72 main()73 