akkki012/sysadmin_env
1
1#!/usr/bin/env python32"""3Baseline inference runner for the sysadmin OpenEnv environment.4 5Required environment variables:6 - API_BASE_URL: OpenAI-compatible API base URL for the LLM7 - MODEL_NAME: model identifier used for planning commands8 - HF_TOKEN: API token for the LLM endpoint9 10Optional:11 - ENV_BASE_URL: environment server URL (default: http://localhost:7860)12 - TASK_IDS: comma-separated task ids to run (default: 3 tasks)13 - MAX_STEPS: step cap per task (default: 8)14"""15 16import argparse17import asyncio18import json19import os20from typing import List, Tuple21 22from openai import OpenAI23from sysadmin_env import SysadminAction, SysadminEnv24 25 26# Explicit alias to match submission wording that requests OpenAIClient usage.27OpenAIClient = OpenAI28 29DEFAULT_TASKS = ["ssh_hardening", "open_port", "malicious_cron"]30SAFETY_BLOCKLIST = [31 "rm -rf /",32 "mkfs",33 "dd if=/dev/zero of=/dev/",34 "> /dev/sda",35]36 37 38def _required_env(name: str) -> str:39 value = os.getenv(name, "").strip()40 if not value:41 raise RuntimeError(f"Missing required environment variable: {name}")42 return value43 44 45def _extract_json(text: str) -> dict:46 text = text.strip()47 try:48 return json.loads(text)49 except json.JSONDecodeError:50 start = text.find("{")51 end = text.rfind("}")52 if start == -1 or end == -1 or end <= start:53 return {}54 try:55 return json.loads(text[start : end + 1])56 except json.JSONDecodeError:57 return {}58 59 60def _choose_action(61 client: OpenAIClient,62 model_name: str,63 task_description: str,64 stdout: str,65 stderr: str,66 history: List[Tuple[str, str, int]],67) -> Tuple[str, str]:68 history_lines = []69 for i, (command, explanation, exit_code) in enumerate(history[-6:], 1):70 history_lines.append(71 f"{i}. command={command!r}, explanation={explanation!r}, exit_code={exit_code}"72 )73 history_text = "\n".join(history_lines) if history_lines else "No previous steps."74 75 prompt = (76 "You are a Linux sysadmin assistant solving one task.\n"77 "Return ONLY valid JSON with keys: command, explanation.\n"78 "Rules:\n"79 "- Use short, safe bash commands.\n"80 "- If task looks complete, return command \"done\".\n"81 "- Never use destructive commands.\n\n"82 f"Task description:\n{task_description}\n\n"83 f"Recent history:\n{history_text}\n\n"84 f"Latest stdout:\n{stdout[:1000]}\n\n"85 f"Latest stderr:\n{stderr[:1000]}\n"86 )87 88 response = client.chat.completions.create(89 model=model_name,90 temperature=0.0,91 messages=[{"role": "user", "content": prompt}],92 )93 content = response.choices[0].message.content or ""94 payload = _extract_json(content)95 command = str(payload.get("command", "done")).strip() or "done"96 explanation = str(payload.get("explanation", "Proceeding to completion.")).strip()97 if any(blocked in command for blocked in SAFETY_BLOCKLIST):98 command = "done"99 explanation = "Stopping to avoid unsafe command."100 return command, explanation101 102 103async def _run_task(104 env_url: str,105 task_id: str,106 client: OpenAIClient,107 model_name: str,108 max_steps: int,109) -> float:110 async with SysadminEnv(base_url=env_url) as env:111 result = await env.reset(task_id=task_id)112 obs = result.observation113 print(f"[START] task_id={task_id} step=0")114 115 history: List[Tuple[str, str, int]] = []116 final_reward = 0.0117 118 for _ in range(max_steps):119 command, explanation = _choose_action(120 client=client,121 model_name=model_name,122 task_description=obs.task_description,123 stdout=obs.stdout,124 stderr=obs.stderr,125 history=history,126 )127 128 step_result = await env.step(129 SysadminAction(command=command, explanation=explanation)130 )131 obs = step_result.observation132 reward = step_result.reward133 if reward is not None:134 final_reward = float(reward)135 136 history.append((command, explanation, obs.exit_code))137 reward_text = "None" if reward is None else f"{float(reward):.3f}"138 print(139 "[STEP] "140 f"task_id={task_id} "141 f"step={obs.step} "142 f"command={json.dumps(command)} "143 f"explanation={json.dumps(explanation)} "144 f"exit_code={obs.exit_code} "145 f"done={obs.done} "146 f"reward={reward_text}"147 )148 149 if obs.done:150 break151 152 if not obs.done:153 step_result = await env.step(154 SysadminAction(155 command="done",156 explanation="Reached local max steps, requesting grading.",157 )158 )159 if step_result.reward is not None:160 final_reward = float(step_result.reward)161 print(162 "[STEP] "163 f"task_id={task_id} "164 f"step={step_result.observation.step} "165 f"command={json.dumps('done')} "166 f"explanation={json.dumps('Reached local max steps, requesting grading.')} "167 f"exit_code={step_result.observation.exit_code} "168 f"done={step_result.observation.done} "169 f"reward={float(step_result.reward or 0.0):.3f}"170 )171 172 print(173 f"[END] task_id={task_id} final_reward={final_reward:.3f} total_steps={obs.step}"174 )175 return final_reward176 177 178async def _amain(args: argparse.Namespace) -> None:179 api_base_url = _required_env("API_BASE_URL")180 model_name = _required_env("MODEL_NAME")181 hf_token = _required_env("HF_TOKEN")182 183 env_url = args.env_url or os.getenv("ENV_BASE_URL", "http://localhost:7860")184 max_steps = int(os.getenv("MAX_STEPS", args.max_steps))185 186 if args.tasks:187 task_ids = [t.strip() for t in args.tasks.split(",") if t.strip()]188 else:189 task_ids = [190 t.strip()191 for t in os.getenv("TASK_IDS", ",".join(DEFAULT_TASKS)).split(",")192 if t.strip()193 ]194 195 client = OpenAIClient(base_url=api_base_url, api_key=hf_token)196 197 rewards: List[float] = []198 for task_id in task_ids:199 try:200 reward = await _run_task(201 env_url=env_url,202 task_id=task_id,203 client=client,204 model_name=model_name,205 max_steps=max_steps,206 )207 except Exception as exc:208 print(209 "[END] "210 f"task_id={task_id} "211 "final_reward=0.000 "212 "total_steps=0 "213 f"error={json.dumps(str(exc))}"214 )215 reward = 0.0216 rewards.append(reward)217 218 mean_reward = sum(rewards) / len(rewards) if rewards else 0.0219 print(f"[END] aggregate_mean_reward={mean_reward:.3f} tasks={len(rewards)}")220 221 222def parse_args() -> argparse.Namespace:223 parser = argparse.ArgumentParser(description="Run baseline inference for sysadmin-env")224 parser.add_argument(225 "--env-url",226 default="",227 help="Environment base URL (default: ENV_BASE_URL or http://localhost:7860)",228 )229 parser.add_argument(230 "--tasks",231 default="",232 help="Comma-separated task IDs (default: TASK_IDS or built-in defaults)",233 )234 parser.add_argument(235 "--max-steps",236 type=int,237 default=8,238 help="Max steps per task before forcing done (default: 8)",239 )240 return parser.parse_args()241 242 243if __name__ == "__main__":244 asyncio.run(_amain(parse_args()))245 