""" Benchmark: Heuristic baseline vs Qwen3-0.6B (Ollama/Metal GPU) on LogSentinel v2. Also supports Groq via --provider groq --api-key Usage: python3 benchmark.py # Ollama + Qwen3 (default) python3 benchmark.py --provider groq --api-key gsk_xxx """ from __future__ import annotations import argparse, json, os, re, sys, time from pathlib import Path sys.path.insert(0, str(Path(__file__).parent)) from openai import OpenAI from training.eval_baseline_vs_trained import heuristic_action from training.metrics import TrainingRun, extract_episode_metrics from training.plot_metrics import make_plots from environment import LogSentinelEnv # --------------------------------------------------------------------------- # Providers # --------------------------------------------------------------------------- PROVIDERS = { "ollama": {"base_url": "http://localhost:11434/v1", "api_key": "ollama", "model": "qwen3-soc"}, "groq": {"base_url": "https://api.groq.com/openai/v1", "model": "qwen/qwen3-32b"}, } # --------------------------------------------------------------------------- # Prompt # --------------------------------------------------------------------------- SYSTEM = """You are a SOC analyst. Reply with ONE JSON object only — no markdown, no explanation. Phase → action_type: detect → propose_incident triage → assign_severity mitigate → execute_mitigation verify → verify_recovery final_report → submit_joint_report Shapes: {"action_type":"propose_incident","agent_role":"","incident_type":"","evidence_indices":[0]} {"action_type":"assign_severity","agent_role":"","incident_type":"","severity":"P1"} {"action_type":"execute_mitigation","agent_role":"","mitigation_id":"fix","evidence_indices":[0]} {"action_type":"verify_recovery","agent_role":"","evidence_indices":[0]} {"action_type":"submit_joint_report","agent_role":"","report":{"incidents":[],"severity":"P1","summary":"done"}} Types: outage resource_exhaustion degradation security_breach config_error""" def build_user_msg(obs: dict) -> str: role = obs.get("agent_role", "incident_commander") phase = obs.get("current_phase", "detect") logs = obs.get("log_entries", [])[:6] lines = "\n".join(f"[{i}] {l.get('level','')}: {l.get('message','')[:100]}" for i, l in enumerate(logs)) return f"Role:{role} Phase:{phase.upper()}\nLOGS:\n{lines}\nJSON:" def parse_action(text: str) -> dict | None: text = re.sub(r".*?", "", text, flags=re.DOTALL).strip() if "```" in text: text = text.split("```")[1].lstrip("json").strip().rstrip("`") s, e = text.find("{"), text.rfind("}") + 1 if s != -1 and e > s: try: r = json.loads(text[s:e]) return r if isinstance(r, dict) else None except Exception: return None return None # --------------------------------------------------------------------------- # Episodes # --------------------------------------------------------------------------- def run_llm_episode(client: OpenAI, model: str, task: str, seed: int, max_steps: int = 8): env = LogSentinelEnv(seed=seed) result = env.reset(task_name=task, seed=seed) rewards, detected, phase_counts = [], [], {} for step in range(max_steps): if result.get("done"): break obs = result.get("observation", {}) try: resp = client.chat.completions.create( model=model, messages=[{"role":"system","content":SYSTEM}, {"role":"user", "content":build_user_msg(obs)}], max_tokens=150, temperature=0.2, ) action = parse_action(resp.choices[0].message.content or "") except Exception as ex: print(f" [API err] {ex}") action = None if not action: action = heuristic_action(obs, step, detected, phase_counts) result = env.step(action) rewards.append(float(result.get("reward") or 0.0)) if action.get("action_type") == "propose_incident" and action.get("incident_type"): detected.append(action["incident_type"]) return rewards, env.state def run_baseline_episode(task: str, seed: int): env = LogSentinelEnv(seed=seed) result = env.reset(task_name=task, seed=seed) rewards, detected, phase_counts, step = [], [], {}, 0 while not result.get("done") and step < 60: obs = result.get("observation", {}) action = heuristic_action(obs, step, detected, phase_counts) result = env.step(action) rewards.append(float(result.get("reward") or 0.0)) if action.get("action_type") == "propose_incident" and action.get("incident_type"): detected.append(action["incident_type"]) step += 1 return rewards, env.state # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(): parser = argparse.ArgumentParser() parser.add_argument("--provider", choices=["ollama","groq"], default="ollama") parser.add_argument("--api-key", default=os.environ.get("GROQ_API_KEY", "")) parser.add_argument("--episodes", type=int, default=10) parser.add_argument("--tasks", nargs="+", default=["soc_warroom_easy","soc_warroom_medium","soc_warroom_hard"]) parser.add_argument("--out-dir", type=Path, default=Path("assets")) args = parser.parse_args() cfg = PROVIDERS[args.provider] model = cfg["model"] key = args.api_key or cfg.get("api_key", "") if args.provider == "groq" and not key: print("ERROR: --api-key required for groq"); sys.exit(1) client = OpenAI(api_key=key, base_url=cfg["base_url"]) args.out_dir.mkdir(parents=True, exist_ok=True) print(f"Provider : {args.provider} | Model: {model}") print(f"Episodes : {args.episodes} × {len(args.tasks)} tasks = {args.episodes*len(args.tasks)} total") print("-" * 55) baseline_run = TrainingRun("baseline") llm_run = TrainingRun(f"{args.provider}_llm") total = args.episodes * len(args.tasks) idx = 0 for task in args.tasks: print(f"\nTask: {task}") for ep in range(args.episodes): seed = ep * 7 + hash(task) % 1000 t0 = time.time() r_b, s_b = run_baseline_episode(task, seed) r_l, s_l = run_llm_episode(client, model, task, seed) m_b = extract_episode_metrics(idx, task, r_b, s_b) m_l = extract_episode_metrics(idx, task, r_l, s_l) baseline_run.record(m_b) llm_run.record(m_l) idx += 1 print(f" [{idx:3d}/{total}] ep={ep:2d} | " f"baseline={m_b.total_reward:.3f} | " f"llm={m_l.total_reward:.3f} | " f"{time.time()-t0:.1f}s") baseline_run.save(args.out_dir / "baseline_metrics.json") llm_run.save(args.out_dir / "trained_metrics.json") print(f"\nMetrics saved → {args.out_dir}/") print("Generating plots...") make_plots( baseline_path=args.out_dir / "baseline_metrics.json", trained_path =args.out_dir / "trained_metrics.json", out_dir=args.out_dir, ) print("Done! Charts → assets/") print(f"\nBaseline : {json.dumps(baseline_run.summary(), indent=2)}") print(f"LLM : {json.dumps(llm_run.summary(), indent=2)}") if __name__ == "__main__": main()