arnavzz
feat: initial OpenEnv code debugging environment
c14504c
Raw
History Blame
5.27 kB
"""
Core environment logic: task loading, reset, step, state.
"""
import json
import random
import uuid
from pathlib import Path
from .executor import run_code_safely
class CodeDebugEnvironment:
def __init__(self):
self.tasks: dict[str, dict] = {}
self.episodes: dict[str, dict] = {}
self._load_tasks()
def _load_tasks(self):
tasks_dir = Path(__file__).parent.parent / "tasks"
for json_file in sorted(tasks_dir.rglob("*.json")):
with open(json_file, encoding="utf-8") as f:
task = json.load(f)
self.tasks[task["task_id"]] = task
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def reset(self, task_id: str | None = None, seed: int | None = None) -> dict:
if seed is not None:
random.seed(seed)
if task_id is None:
task_id = random.choice(list(self.tasks.keys()))
if task_id not in self.tasks:
raise KeyError(f"Unknown task_id: {task_id!r}")
task = self.tasks[task_id]
episode_id = str(uuid.uuid4())
self.episodes[episode_id] = {
"episode_id": episode_id,
"task": task,
"step_count": 0,
"done": False,
"rewards": [],
"last_test_results": [],
}
observation = self._initial_observation(task)
return {"episode_id": episode_id, "observation": observation}
def step(self, episode_id: str, action: dict) -> dict:
if episode_id not in self.episodes:
raise KeyError(f"Unknown episode_id: {episode_id!r}")
ep = self.episodes[episode_id]
if ep["done"]:
raise ValueError("Episode is already finished. Call reset() to start a new episode.")
task = ep["task"]
submitted_code = action.get("code", "")
ep["step_count"] += 1
test_results_raw, stdout, stderr = run_code_safely(
submitted_code,
task["test_code"],
timeout=10,
)
tests_passed = sum(1 for t in test_results_raw if t.get("passed", False))
total_tests = len(test_results_raw)
reward = round(tests_passed / total_tests, 4) if total_tests > 0 else 0.0
max_steps = task.get("max_steps", 5)
done = reward == 1.0 or ep["step_count"] >= max_steps
ep["done"] = done
ep["rewards"].append(reward)
ep["last_test_results"] = test_results_raw
observation = {
"task_id": task["task_id"],
"difficulty": task["difficulty"],
"description": task["description"],
"buggy_code": task["buggy_code"],
"test_descriptions": task["test_descriptions"],
"test_results": test_results_raw,
"stdout": stdout,
"stderr": stderr,
"step_count": ep["step_count"],
"max_steps": max_steps,
"reward": reward,
"done": done,
"total_tests": total_tests,
"tests_passed": tests_passed,
}
return {"observation": observation, "reward": reward, "done": done, "info": {}}
def state(self, episode_id: str) -> dict:
if episode_id not in self.episodes:
raise KeyError(f"Unknown episode_id: {episode_id!r}")
ep = self.episodes[episode_id]
task = ep["task"]
last_results = ep.get("last_test_results", [])
return {
"episode_id": episode_id,
"task_id": task["task_id"],
"difficulty": task["difficulty"],
"step_count": ep["step_count"],
"max_steps": task.get("max_steps", 5),
"last_reward": ep["rewards"][-1] if ep["rewards"] else 0.0,
"cumulative_reward": round(sum(ep["rewards"]), 4),
"tests_passed": sum(1 for t in last_results if t.get("passed", False)),
"total_tests": len(last_results),
"done": ep["done"],
}
def list_tasks(self) -> list[dict]:
return [
{
"task_id": t["task_id"],
"difficulty": t["difficulty"],
"description": t["description"],
"max_steps": t.get("max_steps", 5),
"total_tests": len(t["test_descriptions"]),
}
for t in self.tasks.values()
]
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _initial_observation(self, task: dict) -> dict:
return {
"task_id": task["task_id"],
"difficulty": task["difficulty"],
"description": task["description"],
"buggy_code": task["buggy_code"],
"test_descriptions": task["test_descriptions"],
"test_results": [],
"stdout": "",
"stderr": "",
"step_count": 0,
"max_steps": task.get("max_steps", 5),
"reward": 0.0,
"done": False,
"total_tests": len(task["test_descriptions"]),
"tests_passed": 0,
}