tek-wizard/devops-incident-responder
2
1from __future__ import annotations2 3import json4from pathlib import Path5 6 7NOTEBOOK_PATH = Path("/home/prateeksingh/Desktop/devops_responder/devops_sft_showcase.ipynb")8 9 10def md(source: str) -> dict:11 return {12 "cell_type": "markdown",13 "metadata": {},14 "source": [line + "\n" for line in source.strip("\n").split("\n")],15 }16 17 18def code(source: str) -> dict:19 return {20 "cell_type": "code",21 "execution_count": None,22 "metadata": {},23 "outputs": [],24 "source": [line + "\n" for line in source.strip("\n").split("\n")],25 }26 27 28def main() -> None:29 notebook = {30 "cells": [31 md(32 """33# DevOps SFT Showcase34 35This notebook is the cleanest path to a strong, defensible training result for the hackathon.36 37Instead of optimizing the full 6-task benchmark with unstable GRPO, it trains a compact incident-response policy with supervised fine-tuning on 3 showcase tasks:38 39- `bad_auth_deploy`40- `db_pool_exhaustion`41- `network_partition_failover`42 43It compares:44- `random`45- `zero-shot base model`46- `trained SFT model`47- `heuristic expert`48 49Outputs:50- validation tables51- a comparison plot for the README/blog52- saved CSV summaries53 """54 ),55 md("## 1. Install Dependencies"),56 code(57 """58%%capture59!pip install "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"60!pip install --upgrade "trl>=0.15.0" transformers datasets httpx matplotlib pandas61 """62 ),63 md("## 2. Load Model"),64 code(65 """66from unsloth import FastLanguageModel67import torch68 69BASE_MODEL = "Qwen/Qwen2.5-3B-Instruct"70MAX_SEQ_LENGTH = 76871MAX_COMPLETION_TOKENS = 3272MAX_PROMPT_TOKENS = MAX_SEQ_LENGTH - MAX_COMPLETION_TOKENS - 3273 74model, tokenizer = FastLanguageModel.from_pretrained(75 model_name=BASE_MODEL,76 max_seq_length=MAX_SEQ_LENGTH,77 load_in_4bit=True,78 fast_inference=False,79)80 81if tokenizer.pad_token is None:82 tokenizer.pad_token = tokenizer.eos_token83if hasattr(model, "generation_config"):84 model.generation_config.max_length = None85 86model = FastLanguageModel.get_peft_model(87 model,88 r=16,89 target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],90 lora_alpha=16,91 lora_dropout=0,92 bias="none",93 use_gradient_checkpointing="unsloth",94 random_state=42,95)96 97print("Model loaded")98print("BASE_MODEL =", BASE_MODEL)99print("MAX_SEQ_LENGTH =", MAX_SEQ_LENGTH)100 """101 ),102 md("## 3. Connect To The Deployed Environment"),103 code(104 """105import json106import random107import re108from typing import Any109 110import httpx111import matplotlib.pyplot as plt112import pandas as pd113from datasets import Dataset114 115ENV_URL = "https://tek-wizard-devops-incident-responder.hf.space"116SHOWCASE_TASKS = [117 "bad_auth_deploy",118 "db_pool_exhaustion",119 "network_partition_failover",120]121 122with httpx.Client(base_url=ENV_URL, timeout=20) as http:123 health = http.get("/health")124 assert health.status_code == 200, f"Server not reachable: {health.status_code}"125 catalog = http.get("/tasks")126 catalog.raise_for_status()127 catalog_payload = catalog.json()128 129TASKS = {task["id"]: task for task in catalog_payload["tasks"] if task["id"] in SHOWCASE_TASKS}130SEED_SETS = catalog_payload["seed_sets"]131 132 133def env_reset(task_id: str, seed: int) -> dict:134 with httpx.Client(base_url=ENV_URL, timeout=30) as http:135 response = http.post("/reset", params={"task_id": task_id, "seed": int(seed)})136 response.raise_for_status()137 return response.json()138 139 140def env_step(session_id: str, action: dict) -> dict:141 payload = dict(action)142 payload["session_id"] = session_id143 with httpx.Client(base_url=ENV_URL, timeout=30) as http:144 response = http.post("/step", json=payload)145 response.raise_for_status()146 return response.json()147 148 149def env_grade(session_id: str) -> dict:150 with httpx.Client(base_url=ENV_URL, timeout=30) as http:151 response = http.get("/grader", params={"session_id": session_id})152 response.raise_for_status()153 return response.json()154 155 156probe_obs = env_reset(SHOWCASE_TASKS[0], SEED_SETS["train"][0])157SERVICE_NAMES = sorted(probe_obs["service_status"].keys())158KNOWN_COMMANDS = sorted(probe_obs["available_commands"])159KNOWN_TARGETS = ["system", *SERVICE_NAMES]160 161print("Connected to environment:", ENV_URL)162print("Showcase tasks:", SHOWCASE_TASKS)163print("Validation seeds:", SEED_SETS["val"])164 """165 ),166 md("## 4. Policies And Prompt Format"),167 code(168 """169COMMAND_ALIASES = {170 "restart": "restart_service",171 "scale": "scale_service",172 "get_logs": "query_logs",173 "check_metrics": "get_metrics",174}175TARGET_ALIASES = {176 "db": "main-db",177 "main_db": "main-db",178}179VALID_TARGETS = {180 "get_recent_deploys": {"system"},181 "post_status_update": {"system"},182 "query_logs": set(SERVICE_NAMES),183 "get_metrics": set(SERVICE_NAMES),184 "inspect_config": set(SERVICE_NAMES),185 "check_dependency": {"auth-api", "billing-api", "worker"},186 "run_smoke_test": {"auth-api", "billing-api", "worker"},187 "restart_service": {"auth-api", "billing-api", "worker"},188 "rollback_deploy": {"auth-api", "billing-api"},189 "disable_feature_flag": {"auth-api", "billing-api"},190 "scale_service": {"worker", "main-db"},191 "clear_connections": {"main-db"},192 "drain_queue": {"queue"},193 "failover_db": {"main-db"},194}195 196SYSTEM_PROMPT = f\"\"\"You are a senior DevOps incident commander.197Return ONLY compact JSON with exactly these keys:198{{"command":"...", "target":"..."}}199 200Rules:201- Pick exactly one next action.202- Use only these commands: {", ".join(KNOWN_COMMANDS)}203- Use only these targets: {", ".join(KNOWN_TARGETS)}204- Prefer investigation before mitigation unless the evidence is already strong.205- Do not include markdown, prose, or extra fields.206\"\"\"207 208 209def compact_json(value) -> str:210 return json.dumps(value, separators=(",", ":"), sort_keys=True)211 212 213def compact_services(service_status: dict) -> dict:214 summary = {}215 for name, service in service_status.items():216 summary[name] = {217 "status": service["status"],218 "err": round(float(service["error_rate"]), 3),219 "lat": round(float(service["latency"]), 3),220 "sat": round(float(service["saturation"]), 3),221 "replicas": int(service["replicas"]),222 }223 return summary224 225 226def fit_prompt(prompt: str) -> str:227 token_ids = tokenizer(prompt, add_special_tokens=False)["input_ids"]228 if len(token_ids) <= MAX_PROMPT_TOKENS:229 return prompt230 trimmed_ids = token_ids[-MAX_PROMPT_TOKENS:]231 return tokenizer.decode(trimmed_ids, skip_special_tokens=False)232 233 234def make_prompt(observation: dict, history: list[dict]) -> str:235 history_tail = history[-5:]236 prompt = (237 f"<|im_start|>system\\n{SYSTEM_PROMPT}<|im_end|>\\n"238 f"<|im_start|>user\\n"239 f"Task: {observation['task_id']}\\n"240 f"Incident brief: {observation['incident_brief']}\\n"241 f"Step: {observation['step_count']}\\n"242 f"Remaining steps: {observation['remaining_steps']}\\n"243 f"Global signals: {compact_json(observation['global_signals'])}\\n"244 f"Services: {compact_json(compact_services(observation['service_status']))}\\n"245 f"Recent events: {compact_json(observation['recent_events'][-4:])}\\n"246 f"Discovered facts: {compact_json(observation['discovered_facts'][-6:])}\\n"247 f"Recent action history: {compact_json(history_tail)}\\n"248 "Return the next action JSON now.<|im_end|>\\n"249 f"<|im_start|>assistant\\n"250 )251 return fit_prompt(prompt)252 253 254def _seen(history: list[dict[str, str]], command: str, target: str) -> bool:255 return any(item.get("command") == command and item.get("target") == target for item in history)256 257 258def _queue_depth(observation: dict[str, Any]) -> float:259 return float(observation.get("global_signals", {}).get("queue_depth", 0.0))260 261 262def _text(observation: dict[str, Any]) -> str:263 facts = observation.get("discovered_facts", [])264 events = observation.get("recent_events", [])265 return " ".join([*facts, *events]).lower()266 267 268def heuristic_action(observation: dict[str, Any], history: list[dict[str, str]] | None = None) -> dict[str, str]:269 history = history or []270 task_id = observation.get("task_id")271 text = _text(observation)272 273 if task_id == "bad_auth_deploy":274 if not _seen(history, "get_recent_deploys", "system"):275 return {"command": "get_recent_deploys", "target": "system"}276 if not _seen(history, "query_logs", "auth-api"):277 return {"command": "query_logs", "target": "auth-api"}278 if not _seen(history, "rollback_deploy", "auth-api"):279 return {"command": "rollback_deploy", "target": "auth-api"}280 if not _seen(history, "run_smoke_test", "auth-api"):281 return {"command": "run_smoke_test", "target": "auth-api"}282 return {"command": "post_status_update", "target": "system"}283 284 if task_id == "db_pool_exhaustion":285 if "pool" not in text and not _seen(history, "get_metrics", "main-db"):286 return {"command": "get_metrics", "target": "main-db"}287 if "stale" not in text and not _seen(history, "query_logs", "main-db"):288 return {"command": "query_logs", "target": "main-db"}289 if not _seen(history, "scale_service", "main-db"):290 return {"command": "scale_service", "target": "main-db"}291 if not _seen(history, "clear_connections", "main-db"):292 return {"command": "clear_connections", "target": "main-db"}293 return {"command": "run_smoke_test", "target": "auth-api"}294 295 if task_id == "network_partition_failover":296 if "partition" not in text and not _seen(history, "check_dependency", "auth-api"):297 return {"command": "check_dependency", "target": "auth-api"}298 if "zone" not in text and not _seen(history, "query_logs", "auth-api"):299 return {"command": "query_logs", "target": "auth-api"}300 if not _seen(history, "failover_db", "main-db"):301 return {"command": "failover_db", "target": "main-db"}302 if not _seen(history, "run_smoke_test", "auth-api"):303 return {"command": "run_smoke_test", "target": "auth-api"}304 return {"command": "post_status_update", "target": "system"}305 306 return {"command": "query_logs", "target": "auth-api"}307 308 309def random_action(observation: dict[str, Any], history: list[dict[str, str]] | None = None, seed: int = 0) -> dict[str, str]:310 rng = random.Random(seed + len(history or []))311 return {312 "command": rng.choice(observation["available_commands"]),313 "target": rng.choice(["system", *observation["service_status"].keys()]),314 }315 316 317def strip_fences(text: str) -> str:318 cleaned = text.strip()319 if cleaned.startswith("```"):320 cleaned = re.sub(r"^```(?:json)?\\s*", "", cleaned, flags=re.IGNORECASE)321 cleaned = re.sub(r"\\s*```$", "", cleaned)322 return cleaned.strip()323 324 325def parse_action(completion: str) -> tuple[dict | None, str]:326 cleaned = strip_fences(completion)327 match = re.search(r"\\{.*\\}", cleaned, re.DOTALL)328 blob = match.group(0) if match else cleaned329 try:330 payload = json.loads(blob)331 except json.JSONDecodeError:332 return None, "invalid_json"333 if not isinstance(payload, dict):334 return None, "not_object"335 336 command = COMMAND_ALIASES.get(str(payload.get("command", "")).strip(), str(payload.get("command", "")).strip())337 target = TARGET_ALIASES.get(str(payload.get("target", "")).strip(), str(payload.get("target", "")).strip())338 if command not in KNOWN_COMMANDS:339 return None, "unknown_command"340 if target not in KNOWN_TARGETS:341 return None, "unknown_target"342 if target not in VALID_TARGETS.get(command, set()):343 return None, "invalid_target_for_command"344 return {"command": command, "target": target}, "json"345 """346 ),347 md("## 5. Build Training Dataset From Teacher Rollouts"),348 code(349 """350def rollout_teacher(task_id: str, seed: int) -> list[dict]:351 observation = env_reset(task_id, int(seed))352 session_id = observation["session_id"]353 history = []354 rows = []355 356 while True:357 teacher_action = heuristic_action(observation, history)358 rows.append(359 {360 "prompt": make_prompt(observation, history),361 "completion": compact_json(teacher_action),362 "task_id": task_id,363 "seed": int(seed),364 }365 )366 history.append(teacher_action)367 payload = env_step(session_id, teacher_action)368 observation = payload["observation"]369 if payload.get("done", False):370 break371 372 return rows373 374 375train_rows = []376for task_id in SHOWCASE_TASKS:377 for seed in SEED_SETS["train"]:378 train_rows.extend(rollout_teacher(task_id, seed))379 380# Repeat rows to strengthen a tiny dataset and stabilize SFT.381augmented_rows = train_rows * 4382random.Random(42).shuffle(augmented_rows)383 384sft_rows = [{"text": row["prompt"] + row["completion"]} for row in augmented_rows]385sft_dataset = Dataset.from_list(sft_rows)386 387print("Original training states:", len(train_rows))388print("Augmented SFT rows:", len(sft_rows))389print("Sample completion:", train_rows[0]["completion"])390 """391 ),392 md("## 6. Evaluation Helpers"),393 code(394 """395FastLanguageModel.for_inference(model)396 397 398def generate_model_action(observation: dict, history: list[dict], *, temperature: float = 0.0, do_sample: bool = False) -> tuple[str, dict | None, str]:399 prompt = make_prompt(observation, history)400 inputs = tokenizer(401 prompt,402 return_tensors="pt",403 truncation=True,404 max_length=MAX_PROMPT_TOKENS,405 ).to(model.device)406 407 with torch.no_grad():408 output = model.generate(409 **inputs,410 max_new_tokens=MAX_COMPLETION_TOKENS,411 temperature=temperature,412 do_sample=do_sample,413 pad_token_id=tokenizer.eos_token_id,414 max_length=None,415 )416 417 generated = output[0][inputs["input_ids"].shape[1]:]418 completion = tokenizer.decode(generated, skip_special_tokens=True).strip()419 action, parse_mode = parse_action(completion)420 return completion, action, parse_mode421 422 423def run_policy_episode(task_id: str, seed: int, policy_name: str) -> dict:424 observation = env_reset(task_id, int(seed))425 session_id = observation["session_id"]426 history = []427 traces = []428 429 while True:430 if policy_name == "random":431 action = random_action(observation, history, seed=int(seed))432 parse_mode = "random"433 completion = compact_json(action)434 elif policy_name == "heuristic":435 action = heuristic_action(observation, history)436 parse_mode = "heuristic"437 completion = compact_json(action)438 elif policy_name in {"base_model", "trained_model"}:439 completion, action, parse_mode = generate_model_action(440 observation,441 history,442 temperature=0.0,443 do_sample=False,444 )445 if action is None:446 break447 else:448 raise ValueError(policy_name)449 450 traces.append({"completion": completion, "action": action, "parse_mode": parse_mode})451 history.append(action)452 payload = env_step(session_id, action)453 observation = payload["observation"]454 if payload.get("done", False):455 break456 457 grade = env_grade(session_id)458 return {"task_id": task_id, "seed": int(seed), "grade": grade, "trace": traces}459 460 461def evaluate_policy(policy_name: str) -> pd.DataFrame:462 rows = []463 for task_id in SHOWCASE_TASKS:464 for seed in SEED_SETS["val"]:465 result = run_policy_episode(task_id, seed, policy_name)466 grade = result["grade"]467 rows.append(468 {469 "policy": policy_name,470 "task_id": task_id,471 "seed": int(seed),472 "score": float(grade["score"]),473 "resolved": bool(grade["resolved"]),474 "steps_taken": int(grade["steps_taken"]),475 "unsafe_actions": int(grade["unsafe_actions"]),476 "investigation_coverage": float(grade["investigation_coverage"]),477 }478 )479 return pd.DataFrame(rows)480 """481 ),482 md("## 7. Evaluate Zero-Shot Base Model"),483 code(484 """485base_model_results = evaluate_policy("base_model")486base_model_summary = (487 base_model_results.groupby("policy", as_index=False)488 .agg(avg_score=("score", "mean"), solve_rate=("resolved", "mean"))489)490display(base_model_summary.round(4))491 """492 ),493 md("## 8. Train With SFT"),494 code(495 """496from trl import SFTConfig, SFTTrainer497 498sft_config = SFTConfig(499 output_dir="./sft_showcase_output",500 dataset_text_field="text",501 max_seq_length=MAX_SEQ_LENGTH,502 packing=True,503 per_device_train_batch_size=2,504 gradient_accumulation_steps=4,505 learning_rate=8e-6,506 num_train_epochs=4,507 logging_steps=10,508 save_strategy="no",509 report_to="none",510)511 512sft_trainer = SFTTrainer(513 model=model,514 train_dataset=sft_dataset,515 args=sft_config,516 processing_class=tokenizer,517)518 519print("Running SFT...")520sft_trainer.train()521 """522 ),523 md("## 9. Compare Policies After Training"),524 code(525 """526trained_model_results = evaluate_policy("trained_model")527random_results = evaluate_policy("random")528heuristic_results = evaluate_policy("heuristic")529 530results_df = pd.concat(531 [random_results, base_model_results, trained_model_results, heuristic_results],532 ignore_index=True,533)534 535summary_df = (536 results_df.groupby("policy", as_index=False)537 .agg(538 avg_score=("score", "mean"),539 solve_rate=("resolved", "mean"),540 avg_steps=("steps_taken", "mean"),541 avg_unsafe_actions=("unsafe_actions", "mean"),542 avg_investigation_coverage=("investigation_coverage", "mean"),543 )544)545 546order = ["random", "base_model", "trained_model", "heuristic"]547summary_df["policy"] = pd.Categorical(summary_df["policy"], categories=order, ordered=True)548summary_df = summary_df.sort_values("policy")549display(summary_df.round(4))550 """551 ),552 md("## 10. Save Comparison Plot"),553 code(554 """555plot_df = summary_df.copy()556 557fig, axes = plt.subplots(1, 2, figsize=(13, 5))558 559axes[0].bar(plot_df["policy"], plot_df["avg_score"], color=["#94a1b2", "#7f5af0", "#2cb67d", "#ef4565"])560axes[0].set_title("Average Score")561axes[0].set_ylabel("Validation score")562axes[0].set_ylim(0, 1.0)563for idx, value in enumerate(plot_df["avg_score"]):564 axes[0].text(idx, value + 0.02, f"{value:.2f}", ha="center", fontsize=10)565 566axes[1].bar(plot_df["policy"], plot_df["solve_rate"], color=["#94a1b2", "#7f5af0", "#2cb67d", "#ef4565"])567axes[1].set_title("Solve Rate")568axes[1].set_ylabel("Fraction resolved")569axes[1].set_ylim(0, 1.0)570for idx, value in enumerate(plot_df["solve_rate"]):571 axes[1].text(idx, value + 0.02, f"{value:.2f}", ha="center", fontsize=10)572 573fig.suptitle("DevOps SFT Showcase: 3-Task Comparison", fontsize=14)574fig.tight_layout()575plt.savefig("devops_sft_showcase_comparison.png", dpi=160, bbox_inches="tight")576plt.show()577 578print("Saved devops_sft_showcase_comparison.png")579 """580 ),581 md("## 11. Save CSV Outputs"),582 code(583 """584results_df.to_csv("devops_sft_showcase_results.csv", index=False)585summary_df.to_csv("devops_sft_showcase_summary.csv", index=False)586 587print("Saved devops_sft_showcase_results.csv")588print("Saved devops_sft_showcase_summary.csv")589 """590 ),591 ],592 "metadata": {593 "kernelspec": {594 "display_name": "Python 3",595 "language": "python",596 "name": "python3",597 },598 "language_info": {599 "name": "python",600 "version": "3.12",601 },602 },603 "nbformat": 4,604 "nbformat_minor": 5,605 }606 607 NOTEBOOK_PATH.write_text(json.dumps(notebook, indent=2) + "\n", encoding="utf-8")608 print(f"Wrote {NOTEBOOK_PATH}")609 610 611if __name__ == "__main__":612 main()613 