razak123/code-migration-env
0
1# code_migration_env\server\code_migration_env_environment.py2 3# Copyright (c) Meta Platforms, Inc. and affiliates.4# All rights reserved.5#6# This source code is licensed under the BSD-style license found in the7# LICENSE file in the root directory of this source tree.8 9"""10Code Migration Env Environment Implementation.11 12This environment evaluates production-style code migration tasks. It supports13iterative refinement by returning actionable feedback and partial credit across14multiple attempts.15"""16 17import json18import os19import re20import subprocess21import sys22from pathlib import Path23from typing import Any, Dict, List, Tuple24 25from openenv.core.env_server.interfaces import Environment26 27try:28 from ..models import CodeMigrationAction, CodeMigrationObservation29except ImportError:30 from models import CodeMigrationAction, CodeMigrationObservation31 32 33class CodeMigrationEnvironment(Environment):34 SCORE_EPSILON = 0.0135 TARGET_SUCCESS_IDIOM_SCORE = 0.836 37 SUPPORTS_CONCURRENT_SESSIONS: bool = True38 39 def __init__(self, scenario_dir: str = None):40 repo_root = Path(__file__).resolve().parent.parent41 default_scenario = repo_root / "scenarios" / "easy"42 43 if scenario_dir is None:44 scenario_dir = os.environ.get("SCENARIO_DIR", str(default_scenario))45 46 self.scenario_dir = Path(scenario_dir)47 self.task_meta = json.loads((self.scenario_dir / "meta.json").read_text())48 self.attempts = 049 self.max_attempts = 350 self.history: List[str] = []51 self.source_code = (self.scenario_dir / "source.py").read_text()52 53 def reset(self, seed=None, episode_id=None, **kwargs):54 if not episode_id:55 episode_id = "easy"56 57 repo_root = Path(__file__).resolve().parent.parent58 base_dir = Path(os.environ.get("SCENARIOS_BASE", str(repo_root / "scenarios")))59 60 self.scenario_dir = base_dir / episode_id61 self.task_meta = json.loads((self.scenario_dir / "meta.json").read_text())62 self.source_code = (self.scenario_dir / "source.py").read_text()63 self.attempts = 064 self.history = []65 66 return CodeMigrationObservation(67 task_id=self.task_meta["id"],68 difficulty=self.task_meta["difficulty"],69 source_code=self.source_code,70 source_language=self.task_meta["source_lang"],71 target_language=self.task_meta["target_lang"],72 requirements=self.task_meta["requirements"],73 test_description=self.task_meta["test_desc"],74 history=self.history,75 info=self._build_task_info(),76 )77 78 def _validate_syntax(self, code: str, lang: str) -> Tuple[bool, str]:79 if lang == "python":80 try:81 import ast82 83 ast.parse(code)84 return True, "Syntax OK"85 except SyntaxError as exc:86 return False, f"SyntaxError: {exc}"87 88 if lang == "javascript":89 try:90 import tempfile91 92 with tempfile.NamedTemporaryFile(suffix=".js", mode="w", delete=False) as handle:93 handle.write(code)94 tmp = handle.name95 result = subprocess.run(96 ["node", "--check", tmp],97 capture_output=True,98 text=True,99 timeout=5,100 )101 os.unlink(tmp)102 return result.returncode == 0, result.stderr or "Syntax OK"103 except (subprocess.TimeoutExpired, FileNotFoundError):104 return True, "Syntax check skipped (node unavailable)"105 106 return True, "Syntax check skipped"107 108 def step(self, action: CodeMigrationAction, timeout_s=None, **kwargs) -> CodeMigrationObservation:109 self.attempts += 1110 code = action.translated_code111 112 syntax_ok, syntax_msg = self._validate_syntax(code, self.task_meta["target_lang"])113 idiom_score = self._grade_idioms(code, self.task_meta["required_idioms"])114 missing_idioms = self._get_missing_idioms(code, self.task_meta["required_idioms"])115 116 if not syntax_ok:117 reward = self._clamp_open_score(self.SCORE_EPSILON + (0.05 * idiom_score))118 done = self.attempts >= self.max_attempts119 feedback = self._build_feedback(120 syntax_ok=False,121 syntax_msg=syntax_msg,122 test_ok=False,123 test_msg="Tests were skipped because the submission did not parse.",124 idiom_score=idiom_score,125 missing_idioms=missing_idioms,126 )127 obs = self._get_observation(feedback)128 obs.reward = reward129 obs.done = done130 return obs131 132 reward = 0.3133 test_ok, test_msg = self._run_tests(code, self.task_meta["target_lang"])134 test_progress = self._extract_test_progress(test_msg, test_ok)135 reward += 0.2 * test_progress136 137 reward += idiom_score * 0.4138 reward -= 0.02 * max(0, self.attempts - 1)139 reward = self._clamp_open_score(reward)140 141 done = (142 (test_ok and idiom_score >= self.TARGET_SUCCESS_IDIOM_SCORE)143 or self.attempts >= self.max_attempts144 )145 146 feedback = self._build_feedback(147 syntax_ok=True,148 syntax_msg=syntax_msg,149 test_ok=test_ok,150 test_msg=test_msg,151 idiom_score=idiom_score,152 missing_idioms=missing_idioms,153 )154 obs = self._get_observation(feedback)155 obs.reward = reward156 obs.done = done157 158 info = {159 "attempt": self.attempts,160 "final_reward": reward,161 "tests_passed": test_ok,162 "test_progress": test_progress,163 "idiom_score": idiom_score,164 "missing_idioms": missing_idioms,165 "attempts_remaining": max(0, self.max_attempts - self.attempts),166 "submitted_code_preview": code[:200] + "..." if len(code) > 200 else code,167 }168 print("info: ", info)169 170 return obs171 172 def _clamp_open_score(self, score: float) -> float:173 return max(self.SCORE_EPSILON, min(1.0 - self.SCORE_EPSILON, score))174 175 def _build_task_info(self) -> Dict[str, Any]:176 return {177 "business_context": self.task_meta.get("business_context", ""),178 "stakeholder_request": self.task_meta.get("stakeholder_request", ""),179 "acceptance_checks": self.task_meta.get("acceptance_checks", []),180 "pitfalls": self.task_meta.get("pitfalls", []),181 "runtime_budget": self.task_meta.get("runtime_budget", ""),182 "max_attempts": self.max_attempts,183 "attempts_used": self.attempts,184 "attempts_remaining": max(0, self.max_attempts - self.attempts),185 }186 187 @property188 def state(self) -> Dict[str, Any]:189 return {190 "task_id": self.task_meta["id"],191 "difficulty": self.task_meta["difficulty"],192 "source_language": self.task_meta["source_lang"],193 "target_language": self.task_meta["target_lang"],194 "attempts_used": self.attempts,195 "attempts_remaining": max(0, self.max_attempts - self.attempts),196 "history": self.history[-5:],197 "required_idioms": self.task_meta["required_idioms"],198 }199 200 def _run_tests(self, code: str, lang: str) -> Tuple[bool, str]:201 tests = json.loads((self.scenario_dir / "tests.json").read_text())202 timeout_s = float(self.task_meta.get("test_timeout_s", 2.0))203 204 if lang == "python":205 try:206 wrapper = f"CODE_SOURCE = {code!r}\n{tests.get('runner', '')}"207 result = subprocess.run(208 [sys.executable, "-c", wrapper],209 capture_output=True,210 text=True,211 timeout=timeout_s,212 )213 return result.returncode == 0, result.stdout or result.stderr214 except subprocess.TimeoutExpired:215 return False, "Test timeout"216 except Exception as exc:217 return False, f"Test runner error: {exc}"218 219 if lang == "javascript":220 try:221 import tempfile222 223 with tempfile.NamedTemporaryFile(suffix=".js", mode="w", delete=False) as handle:224 handle.write(code + "\n\n" + tests.get("runner", ""))225 tmp = handle.name226 227 result = subprocess.run(228 ["node", tmp],229 capture_output=True,230 text=True,231 timeout=timeout_s,232 )233 os.unlink(tmp)234 return result.returncode == 0, result.stdout or result.stderr235 except subprocess.TimeoutExpired:236 return False, "Test timeout"237 except FileNotFoundError:238 return False, "Node not available"239 except Exception as exc:240 return False, f"Test runner error: {exc}"241 242 return True, "Tests skipped"243 244 def _grade_idioms(self, code: str, idioms: List[str]) -> float:245 matches = sum(1 for pattern in idioms if pattern in code)246 return matches / len(idioms) if idioms else 1.0247 248 def _get_missing_idioms(self, code: str, idioms: List[str]) -> List[str]:249 return [pattern for pattern in idioms if pattern not in code]250 251 def _extract_test_progress(self, test_msg: str, test_ok: bool) -> float:252 if test_ok:253 return 1.0254 255 match = re.search(r"SCORE:([0-9]*\.?[0-9]+)", str(test_msg))256 if match:257 try:258 score = float(match.group(1))259 except ValueError:260 return 0.0261 return max(0.0, min(1.0, score))262 263 return 0.0264 265 def _normalize_feedback(self, text: str) -> str:266 compact = " ".join(str(text).split())267 return compact[:240]268 269 def _build_feedback(270 self,271 syntax_ok: bool,272 syntax_msg: str,273 test_ok: bool,274 test_msg: str,275 idiom_score: float,276 missing_idioms: List[str],277 ) -> str:278 messages: List[str] = []279 280 if not syntax_ok:281 messages.append(f"Syntax issue: {self._normalize_feedback(syntax_msg)}")282 elif not test_ok:283 messages.append(f"Functional check failed: {self._normalize_feedback(test_msg)}")284 else:285 messages.append("Functional tests passed.")286 287 if missing_idioms:288 messages.append("Still missing target idioms: " + ", ".join(missing_idioms[:4]))289 else:290 messages.append("Target-language idioms look strong.")291 292 messages.append(f"Current idiom score: {idiom_score:.2f}")293 294 if self.attempts < self.max_attempts and (not test_ok or missing_idioms):295 messages.append("Revise the submission using this feedback and try again.")296 297 return " ".join(messages)298 299 def _get_observation(self, feedback: str) -> CodeMigrationObservation:300 self.history.append(f"Attempt {self.attempts}: {feedback}")301 return CodeMigrationObservation(302 task_id=self.task_meta["id"],303 difficulty=self.task_meta["difficulty"],304 source_code=self.source_code,305 source_language=self.task_meta["source_lang"],306 target_language=self.task_meta["target_lang"],307 requirements=self.task_meta["requirements"],308 test_description=self.task_meta["test_desc"],309 history=self.history[-5:],310 info=self._build_task_info(),311 )312 