Team Ai
Apppublic

razak123/code-migration-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
code_migration_env_environment.py312 linesDownload Raw Back to server
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