PRANAV05092003/autonomous-code-refactoring-env
0
1from __future__ import annotations2 3import os4import sys5from typing import Any, Dict, Optional, Tuple6 7PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))8if PROJECT_ROOT not in sys.path:9 sys.path.insert(0, PROJECT_ROOT)10 11try:12 from openenv.env import Env as OpenEnvBase13except Exception: # pragma: no cover14 class OpenEnvBase:15 def __init__(self, *args: Any, **kwargs: Any) -> None:16 return None17 18from acre.datasets.code_samples import CodeSample, CodeSampleDataset19from acre.env.refactor_env import RefactorEnv20from acre.tasks.task_registry import TaskRegistry21from models import ActionModel, ObservationModel, RewardModel, StateResponse22 23 24class OpenEnvRefactorEnv(OpenEnvBase):25 """26 Canonical OpenEnv interface for ACRE.27 28 This wrapper keeps the strict hackathon contract:29 - reset() -> ObservationModel30 - step(action) -> (ObservationModel, RewardModel, done, info)31 - state() -> StateResponse32 """33 34 def __init__(35 self,36 *,37 env: Optional[RefactorEnv] = None,38 registry: Optional[TaskRegistry] = None,39 ) -> None:40 super().__init__(41 name="ACRE",42 state_space="ObservationModel",43 action_space="ActionModel",44 episode_max_length=RefactorEnv.MAX_STEPS,45 )46 self._env = env or RefactorEnv()47 self._registry = registry or TaskRegistry()48 self._task_id: Optional[str] = None49 self._last_reset_info: Dict[str, Any] = {}50 51 @property52 def action_meanings(self) -> Dict[int, str]:53 return self._env.ACTION_MEANINGS54 55 @property56 def last_reset_info(self) -> Dict[str, Any]:57 return dict(self._last_reset_info)58 59 def _load_episode_source(self, *, task_id: Optional[str], code: Optional[str]) -> None:60 initial_code = code61 if initial_code is None and task_id:62 task = self._registry.get_task(task_id)63 if task is None:64 raise ValueError(f"Task '{task_id}' not found")65 # Load a multi-sample dataset for this task. Sample selection is66 # deterministic given the `seed` passed to `reset()`.67 samples = list(getattr(task, "samples", []) or [])68 if not samples:69 initial_code = task.initial_code70 else:71 self._env.dataset = CodeSampleDataset(72 [73 CodeSample(74 id=f"{task_id}:{i}",75 language="python",76 code=str(src),77 )78 for i, src in enumerate(samples)79 ]80 )81 return None82 83 if initial_code is None:84 return None85 86 self._env.dataset = CodeSampleDataset(87 [88 CodeSample(89 id=task_id or "custom",90 language="python",91 code=initial_code,92 )93 ]94 )95 return None96 97 def reset(98 self,99 *,100 seed: Optional[int] = None,101 task_id: Optional[str] = None,102 code: Optional[str] = None,103 ) -> ObservationModel:104 self._task_id = task_id105 self._load_episode_source(task_id=task_id, code=code)106 observation, info = self._env.reset(seed=seed)107 self._last_reset_info = dict(info)108 return ObservationModel.from_vector(observation.tolist())109 110 def step(self, action: int | ActionModel) -> Tuple[ObservationModel, RewardModel, bool, Dict[str, Any]]:111 action_value = action.action if isinstance(action, ActionModel) else int(action)112 observation, raw_reward, terminated, truncated, info = self._env.step(action_value)113 reward = RewardModel(114 raw=float(raw_reward),115 normalized=float(info.get("normalized_reward", 0.0)),116 components=dict(info.get("reward_components", {})),117 )118 done = bool(terminated or truncated)119 return ObservationModel.from_vector(observation.tolist()), reward, done, dict(info)120 121 def state(self) -> StateResponse:122 raw_state = self._env.state()123 observation_vector = list(raw_state.get("observation", [0.0, 0.0, 0.0, 0.0]))124 observation = ObservationModel.from_vector(observation_vector)125 return StateResponse(126 current_code=str(raw_state.get("current_code", "")),127 episode_steps=int(raw_state.get("episode_steps", 0)),128 max_steps=int(raw_state.get("max_steps", RefactorEnv.MAX_STEPS)),129 complexity=float(raw_state.get("complexity", 0.0)),130 last_runtime=float(raw_state.get("last_runtime", 0.0)),131 last_error=bool(raw_state.get("last_error", False)),132 sample_id=raw_state.get("sample_id"),133 language=raw_state.get("language"),134 task_id=self._task_id,135 observation=observation,136 observation_vector=observation.to_vector(),137 action_meanings=dict(raw_state.get("action_meanings", {})),138 )139 