AMD21/codefixerenv
0
1import os2import sys3sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))4 5from env.models import Observation, Action, Reward6from tasks.tasks import TASKS7 8class CodeFixerEnv:9 MAX_STEPS: int = 510 11 def __init__(self) -> None:12 self._task = None13 self._grader = None14 self._difficulty = None15 self._history = []16 self._step_number = 017 self._done = False18 self._best_score = 0.019 self._cumulative_reward = 0.020 21 def reset(self, difficulty: str = "easy") -> Observation:22 if difficulty not in TASKS:23 raise ValueError(24 f"Unknown difficulty '{difficulty}'. Valid choices: {sorted(TASKS.keys())}"25 )26 27 task_def, grader = TASKS[difficulty]28 self._task = task_def29 self._grader = grader30 self._difficulty = difficulty31 self._history = []32 self._step_number = 033 self._done = False34 self._best_score = 0.035 self._cumulative_reward = 0.036 37 return Observation(38 input=task_def["buggy_code"],39 context=task_def["context"],40 history=[],41 step_number=0,42 max_steps=self.MAX_STEPS,43 )44 45 def step(self, action: Action) -> tuple[Observation, Reward, bool, dict]:46 if self._done:47 raise RuntimeError("Episode is over. Call reset() to start a new episode.")48 if self._task is None:49 raise RuntimeError("Environment not initialised. Call reset() first.")50 51 self._step_number += 152 info = {53 "task_id": self._task["id"],54 "difficulty": self._difficulty,55 "step": self._step_number,56 }57 58 raw_value = 0.059 grader_score = None60 improvement = None61 efficiency_bonus = 0.062 penalty = 0.063 feedback = ""64 65 if action.type == "give_up":66 penalty = -0.1067 raw_value = penalty68 self._done = True69 feedback = "Agent chose to give up."70 info["outcome"] = "gave_up"71 72 elif action.type == "explain":73 penalty = -0.0574 raw_value = penalty75 feedback = (76 f"Hint: re-read the loop bounds carefully. "77 f"Task goal: {self._task['context']}"78 )79 info["outcome"] = "explained"80 81 elif action.type == "fix":82 grader_score, grade_detail = self._grader(action.content)83 info["grade_detail"] = grade_detail84 info["grader_score"] = grader_score85 86 improvement = grader_score - self._best_score87 88 if improvement > 0:89 raw_value = improvement90 self._best_score = grader_score91 else:92 penalty = -0.0593 raw_value = penalty94 95 if grader_score >= 1.0 and self._step_number <= 2:96 efficiency_bonus = 0.2097 raw_value += efficiency_bonus98 info["efficiency_bonus"] = True99 100 if grader_score >= 1.0:101 self._done = True102 feedback = "All test cases passed — perfect fix!"103 info["outcome"] = "solved"104 else:105 passed = grade_detail.get("tests_passed", "?")106 total = grade_detail.get("tests_total", "?")107 feedback = (108 f"Score {grader_score:.2f} — "109 f"{passed}/{total} tests passed. Keep refining."110 )111 info["outcome"] = "partial"112 113 raw_value = round(raw_value, 4)114 self._cumulative_reward = round(self._cumulative_reward + raw_value, 4)115 116 if self._step_number >= self.MAX_STEPS and not self._done:117 self._done = True118 info["outcome"] = info.get("outcome", "timeout")119 120 reward = Reward(121 value = raw_value,122 grader_score = grader_score,123 improvement = round(improvement, 4) if improvement is not None else None,124 efficiency_bonus = efficiency_bonus,125 penalty = penalty,126 reason = feedback,127 )128 129 self._history.append({130 "step": self._step_number,131 "action_type": action.type,132 "action_content": (133 action.content[:120] + "..."134 if len(action.content) > 120 else action.content135 ),136 "reward": raw_value,137 "feedback": feedback,138 })139 140 next_obs = Observation(141 input = self._task["buggy_code"],142 context = self._task["context"],143 history = list(self._history),144 step_number = self._step_number,145 max_steps = self.MAX_STEPS,146 )147 return next_obs, reward, self._done, info148 149 def state(self) -> dict:150 return {151 "task_id": self._task["id"] if self._task else None,152 "difficulty": self._difficulty,153 "step_number": self._step_number,154 "max_steps": self.MAX_STEPS,155 "done": self._done,156 "best_score": self._best_score,157 "cumulative_reward": self._cumulative_reward,158 "history_length": len(self._history),159 }160 