SidhaGarg/Cloud-DevOps-RLEnv
0
1import asyncio2import json3import os4import sys5from typing import Any, Dict, List, Tuple6 7from openai import OpenAI8from pydantic import ValidationError9 10from env import CloudDevOpsEnv11from models import CloudAction12 13API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")14MODEL_NAME = os.getenv("MODEL_NAME", "google/gemma-4-26B-A4B-it")15HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("API_KEY")16 17BENCHMARK = "CloudDevOpsEnv"18MAX_STEPS = 1519MAX_TOTAL_REWARD = 1.020SCORE_MIN = 0.00121SCORE_MAX = 0.99922 23 24def log_start(task: str, env: str, model: str) -> None:25 print(f"[START] task={task} env={env} model={model}", flush=True)26 27 28def log_step(step: int, action: Any, reward: float, done: bool, error: Any) -> None:29 action_dict = action.model_dump() if hasattr(action, "model_dump") else str(action)30 if isinstance(action_dict, dict):31 action_str = json.dumps(action_dict, separators=(",", ":"))32 else:33 action_str = str(action_dict)34 action_str = action_str.replace("\n", " ").replace("\r", " ")35 36 error_str = "null" if not error else str(error).replace("\n", " ").replace("\r", " ")37 done_str = str(done).lower()38 print(39 f"[STEP] step={step} action={action_str} reward={reward:.2f} done={done_str} error={error_str}",40 flush=True,41 )42 43 44def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:45 rewards_str = ",".join(f"{r:.2f}" for r in rewards)46 success_str = str(success).lower()47 print(48 f"[END] success={success_str} steps={steps} score={score:.3f} rewards={rewards_str}",49 flush=True,50 )51 52 53def get_model_action(54 client: OpenAI,55 task_name: str,56 step: int,57 last_obs: str,58 last_error: str,59 history: List[Dict[str, str]],60) -> Tuple[CloudAction, str]:61 """Prompt the LLM and parse its response into a CloudAction."""62 system_prompt = (63 "You are an expert AI DevOps Engineer diagnosing a cloud infrastructure issue. "64 "You must respond ONLY with a raw JSON object matching this schema:\n"65 "{\n"66 ' "command": "list_resources" | "describe_resource" | "view_logs" | "query_metadata" | "update_security_group" | "restart_service" | "submit_solution",\n'67 ' "resource_id": "string (optional)",\n'68 ' "parameters": {"key": "value"} (optional)\n'69 "}\n"70 "Optimization objective: maximize reward by minimizing unnecessary actions because each step has a cost.\n"71 "Use parameters only when needed:\n"72 "- update_security_group: parameters must include port and action\n"73 "- query_metadata: parameters must include ip_address\n"74 "- list_resources / describe_resource / view_logs / restart_service / submit_solution: parameters should be omitted\n"75 "Task playbooks:\n"76 "- easy: identify sg-web and open port 80 using update_security_group with action=allow\n"77 "- medium: inspect i-api logs, resolve DB IP using query_metadata, then update sg-db port 5432 with action=allow\n"78 "- hard: inspect lb-main logs, resolve failing upstream IP via query_metadata, inspect i-web2, then restart i-web2\n"79 "When logs provide only IP addresses, use query_metadata with parameters.ip_address to resolve the resource_id before remediation.\n"80 "Do not include markdown blocks like ```json. Just output the JSON."81 )82 83 user_prompt = (84 f"Task: {task_name}\n"85 f"Step {step}.\n"86 f"Last Observation:\n{last_obs}\n"87 )88 if last_error:89 user_prompt += f"\nLast Error:\n{last_error}\n"90 user_prompt += "\nWhat is your next action JSON?"91 92 messages = [{"role": "system", "content": system_prompt}] + history + [93 {"role": "user", "content": user_prompt}94 ]95 96 try:97 response = client.chat.completions.create(98 model=MODEL_NAME,99 messages=messages,100 temperature=0.0,101 max_tokens=200,102 )103 raw_text = (response.choices[0].message.content or "").strip()104 105 if raw_text.startswith("```json"):106 raw_text = raw_text.replace("```json", "").replace("```", "").strip()107 108 action_dict = json.loads(raw_text)109 return CloudAction(**action_dict), raw_text110 except (json.JSONDecodeError, ValidationError) as exc:111 print(f"[DEBUG] Model parse failed: {exc}", file=sys.stderr, flush=True)112 return CloudAction(command="list_resources"), "failed_parse"113 except Exception as exc:114 print(f"[DEBUG] API request failed: {exc}", file=sys.stderr, flush=True)115 return CloudAction(command="list_resources"), "api_error"116 117 118async def run_task(task_name: str, client: OpenAI) -> None:119 env = CloudDevOpsEnv(task_name=task_name)120 121 history: List[Dict[str, str]] = []122 rewards: List[float] = []123 steps_taken = 0124 score = 0.0125 success = False126 127 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)128 129 try:130 result = await env.reset()131 last_obs = result.observation.output132 last_error = result.observation.error or ""133 134 for step in range(1, MAX_STEPS + 1):135 if result.done:136 break137 138 action, raw_response = get_model_action(139 client, task_name, step, last_obs, last_error, history140 )141 142 result = await env.step(action)143 obs = result.observation144 reward = result.reward or 0.0145 done = result.done146 error = obs.error147 148 rewards.append(reward)149 steps_taken = step150 last_obs = obs.output151 last_error = error or ""152 153 log_step(step=step, action=action, reward=reward, done=done, error=error)154 155 history.append({"role": "assistant", "content": raw_response})156 history.append(157 {158 "role": "user",159 "content": f"Observation: {last_obs}\nError: {last_error}",160 }161 )162 163 if done:164 break165 166 score = sum(rewards)167 # Keep score strictly in (0,1) after formatting to avoid validator endpoint failures.168 score = max(SCORE_MIN, min(score, SCORE_MAX))169 success = bool(result.info.get("resolved", False))170 171 finally:172 try:173 await env.close()174 except Exception as exc:175 print(f"[DEBUG] env.close() failed: {exc}", file=sys.stderr, flush=True)176 log_end(success=success, steps=steps_taken, score=score, rewards=rewards)177 178 179async def main() -> None:180 if not HF_TOKEN:181 print(182 "[WARN] HF_TOKEN (or API_KEY fallback) is not set. API calls will fail in remote evaluation.",183 file=sys.stderr,184 flush=True,185 )186 187 client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)188 189 tasks = ["easy", "medium", "hard"]190 for task in tasks:191 await run_task(task, client)192 193 194if __name__ == "__main__":195 asyncio.run(main())196 