FixOS / server /my_env_environment.py
Dishamharitasa's picture
Update server/my_env_environment.py
601f240 verified
Raw
History Blame Contribute Delete
8.27 kB
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
"""FixOS environment implementation."""
from dataclasses import dataclass
from uuid import uuid4
from typing import Any, Callable, Dict, List
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State
try:
from ..models import FileInfo, FixOSAction, FixOSObservation, LogEntry, ProcessInfo, ServiceInfo
except ImportError:
from models import FileInfo, FixOSAction, FixOSObservation, LogEntry, ProcessInfo, ServiceInfo
# ---------------------------
# SAFE SCORE FUNCTION
# ---------------------------
def fix_score(score: float) -> float:
# strictly (0,1)
return max(0.01, min(0.99, float(score)))
# ---------------------------
# TASK DEFINITION
# ---------------------------
@dataclass(frozen=True)
class TaskDefinition:
task_id: str
difficulty: str
failures: Dict[str, Any]
optimal_steps: int
success_check: Callable[[Dict[str, Any]], bool]
# ---------------------------
# TASK ORDER
# ---------------------------
TASK_ORDER: List[str] = [
"easy_1", "easy_2", "medium_1", "medium_2", "hard_1", "hard_2", "hard_3",
]
# ---------------------------
# TASKS (REQUIRED — NO {...})
# ---------------------------
TASKS = {
"easy_1": TaskDefinition(
task_id="easy_1",
difficulty="easy",
failures={"nginx_status": "stopped"},
optimal_steps=3,
success_check=lambda s: s["nginx_status"] == "running",
),
"easy_2": TaskDefinition(
task_id="easy_2",
difficulty="easy",
failures={"mysql_status": "stopped"},
optimal_steps=3,
success_check=lambda s: s["mysql_status"] == "running",
),
"medium_1": TaskDefinition(
task_id="medium_1",
difficulty="medium",
failures={"nginx_config_valid": False},
optimal_steps=5,
success_check=lambda s: s["nginx_config_valid"] and s["nginx_status"] == "running",
),
"medium_2": TaskDefinition(
task_id="medium_2",
difficulty="medium",
failures={"mysql_config_valid": False},
optimal_steps=5,
success_check=lambda s: s["mysql_config_valid"] and s["mysql_status"] == "running",
),
"hard_1": TaskDefinition(
task_id="hard_1",
difficulty="hard",
failures={"disk_percent": 98.0},
optimal_steps=7,
success_check=lambda s: s["disk_percent"] < 90,
),
"hard_2": TaskDefinition(
task_id="hard_2",
difficulty="hard",
failures={"nginx_status": "stopped", "nginx_config_valid": False},
optimal_steps=8,
success_check=lambda s: s["nginx_status"] == "running" and s["nginx_config_valid"],
),
"hard_3": TaskDefinition(
task_id="hard_3",
difficulty="hard",
failures={"mysql_status": "stopped", "mysql_config_valid": False},
optimal_steps=8,
success_check=lambda s: s["mysql_status"] == "running" and s["mysql_config_valid"],
),
}
# ---------------------------
# BASE SCORE FUNCTION
# ---------------------------
def _task_score_base(state: Dict[str, Any]) -> float:
score = 0.0
if state["disk_percent"] <= 95:
score += 0.2
if state["nginx_status"] == "running":
score += 0.2
if state["mysql_status"] == "running":
score += 0.2
if state["nginx_config_valid"]:
score += 0.2
if state["mysql_config_valid"]:
score += 0.2
return fix_score(score)
# ---------------------------
# ENVIRONMENT CLASS
# ---------------------------
class FixOSEnvironment(Environment):
SUPPORTS_CONCURRENT_SESSIONS: bool = True
MAX_STEPS_PER_EPISODE: int = 50
def __init__(self):
self._state = State(episode_id=str(uuid4()), step_count=0)
self._current_task = None
self._task_pointer = -1
self._env_state: Dict[str, Any] = {}
self._previous_score = 0.0
# ---------------------------
# RESET
# ---------------------------
def reset(self):
self._state = State(episode_id=str(uuid4()), step_count=0)
self._task_pointer = (self._task_pointer + 1) % len(TASK_ORDER)
self._current_task = TASKS[TASK_ORDER[self._task_pointer]]
self._env_state = {
"task_id": self._current_task.task_id,
"difficulty": self._current_task.difficulty,
"nginx_status": "running",
"mysql_status": "running",
"nginx_config_valid": True,
"mysql_config_valid": True,
"disk_percent": 45.0,
"processes": [],
"files": {},
"logs": [],
"history": [],
"logs_checked": False,
"config_checked": {"nginx": False, "mysql": False},
}
self._apply_failures(self._current_task.failures)
self._refresh_resources()
# FIXED
self._previous_score = fix_score(self._compute_task_score())
return self._build_observation("[INFO] Environment reset")
# ---------------------------
# STEP
# ---------------------------
def step(self, action):
self._state.step_count += 1
reward = 0.0
current_score = fix_score(self._compute_task_score())
delta = max(0.0, current_score - self._previous_score)
if delta > 0:
reward += 0.4 * delta
self._previous_score = current_score
is_success = self._current_task.success_check(self._env_state)
if is_success:
reward += 0.5
done = is_success or self._state.step_count >= self.MAX_STEPS_PER_EPISODE
reward = fix_score(reward)
return self._build_observation(
command_output="step executed",
reward=reward,
done=done,
is_success_step=is_success,
task_score=current_score, # already fixed
)
# ---------------------------
# SCORE COMPUTATION
# ---------------------------
def _compute_task_score(self) -> float:
s = self._env_state
task_id = self._current_task.task_id
if task_id.startswith("easy"):
raw = (
(0.7 if ((s["nginx_status"] == "running" and task_id == "easy_1")
or (s["mysql_status"] == "running" and task_id == "easy_2")) else 0.0)
+ (0.3 if s["disk_percent"] < 95 else 0.0)
)
return fix_score(raw)
score = 0.0
if task_id == "medium_1":
score += 0.2 if s["logs_checked"] else 0.0
score += 0.3 if s["nginx_config_valid"] else 0.0
score += 0.5 if s["nginx_status"] == "running" else 0.0
elif task_id == "medium_2":
score += 0.2 if s["logs_checked"] else 0.0
score += 0.3 if s["mysql_config_valid"] else 0.0
score += 0.5 if s["mysql_status"] == "running" else 0.0
else:
return _task_score_base(s)
return fix_score(score)
# ---------------------------
# HELPERS
# ---------------------------
def _apply_failures(self, failures):
for k, v in failures.items():
self._env_state[k] = v
def _refresh_resources(self):
self._env_state["cpu_percent"] = 50.0
self._env_state["memory_mb"] = 500.0
def _build_observation(
self,
command_output,
reward=0.0,
done=False,
is_success_step=False,
task_score=None,
):
return FixOSObservation(
command_output=command_output,
processes=[],
services=[],
filesystem=[],
resources={},
logs=[],
history=[],
task_id=self._env_state["task_id"],
task_difficulty=self._env_state["difficulty"],
task_score=fix_score(
self._compute_task_score() if task_score is None else task_score
),
is_success_step=is_success_step,
remaining_steps=max(0, self.MAX_STEPS_PER_EPISODE - self._state.step_count),
done=done,
reward=fix_score(reward),
)
@property
def state(self):
return self._state