Spaces:
Sleeping
Sleeping
| """ | |
| Benchmark: Heuristic baseline vs Qwen3-0.6B (Ollama/Metal GPU) on LogSentinel v2. | |
| Also supports Groq via --provider groq --api-key <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":"<role>","incident_type":"<type>","evidence_indices":[0]} | |
| {"action_type":"assign_severity","agent_role":"<role>","incident_type":"<type>","severity":"P1"} | |
| {"action_type":"execute_mitigation","agent_role":"<role>","mitigation_id":"fix","evidence_indices":[0]} | |
| {"action_type":"verify_recovery","agent_role":"<role>","evidence_indices":[0]} | |
| {"action_type":"submit_joint_report","agent_role":"<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"<think>.*?</think>", "", 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() | |