CoolFace
Apppublic

akkki012/sysadmin_env

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes
inference.py245 linesDownload Raw Back to root
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