razak123/code-migration-env
0
1import asyncio2import json3import os4import re5import sys6from pathlib import Path7from textwrap import dedent8from typing import Dict, List, Optional, Tuple9 10from dotenv import load_dotenv11 12try:13 from openai import OpenAI14except ImportError: # pragma: no cover15 OpenAI = None16 17 18# Ensure imports work whether run from repo root or package dir.19_here = Path(__file__).resolve().parent20sys.path.insert(0, str(_here))21sys.path.insert(0, str(_here / "code_migration_env"))22 23os.environ["OPENBLAS_NUM_THREADS"] = "1"24os.environ["OMP_NUM_THREADS"] = "1"25 26try:27 from code_migration_env.client import CodeMigrationEnv28 from code_migration_env.models import CodeMigrationAction29except ImportError:30 from client import CodeMigrationEnv31 from models import CodeMigrationAction32 33 34load_dotenv()35 36DEFAULT_MODEL_NAME = "mistralai/devstral-2-123b-instruct-2512"37 38API_BASE_URL = (39 os.getenv("API_BASE_URL")40 or os.getenv("OPENAI_BASE_URL")41 or "https://router.huggingface.co/v1"42)43MODEL_NAME = (44 os.getenv("MODEL_NAME")45 or os.getenv("OPENAI_MODEL")46 or os.getenv("LITELLM_MODEL")47 or DEFAULT_MODEL_NAME48).strip()49HF_TOKEN = (os.getenv("HF_TOKEN") or "").strip()50API_KEY = (51 os.getenv("API_KEY")52 or os.getenv("OPENAI_API_KEY")53 or HF_TOKEN54 or ""55).strip()56 57MAX_STEPS = int(os.environ.get("MAX_STEPS", "3"))58MAX_TOTAL_REWARD = float(os.environ.get("MAX_TOTAL_REWARD", "1.0"))59SUCCESS_SCORE_THRESHOLD = float(os.environ.get("SUCCESS_SCORE_THRESHOLD", "0.5"))60SCORE_EPSILON = 0.00161 62TASKS = [63 ("python_modernize", "easy"),64 ("python_to_node", "medium"),65 ("pandas_to_polars_advanced", "hard"),66]67 68DETERMINISTIC_SOLUTIONS: Dict[str, str] = {69 "python_modernize": dedent(70 """71 from pathlib import Path72 73 def get_config(name: str) -> str | int | None:74 match name:75 case "db_host":76 return "localhost"77 case "db_port":78 return 543279 case _:80 return None81 82 def read_file(path: str) -> str:83 return Path("/data", path).read_text()84 """85 ).strip(),86 "python_to_node": dedent(87 """88 async function getUser(userId, includeDetails = false) {89 try {90 const url = new URL(`https://api.example.com/users/${userId}`);91 url.searchParams.set("details", includeDetails ? "1" : "0");92 93 const response = await fetch(url.toString(), {94 headers: {95 Accept: "application/json",96 "X-Request-Source": "openenv"97 }98 });99 100 if (!response.ok) {101 throw new Error(`Failed: ${response.status}`);102 }103 104 const data = await response.json();105 data.source = "api";106 return data;107 } catch (error) {108 throw error;109 }110 }111 """112 ).strip(),113 "pandas_to_polars": dedent(114 """115 import polars as pl116 117 def process_sales(filepath: str, exclude_regions: list[str] | None = None) -> pl.DataFrame:118 exclude_regions = exclude_regions or []119 120 lf = pl.scan_csv(filepath, try_parse_dates=True)121 lf = lf.filter(pl.col("status") != "cancelled")122 123 if exclude_regions:124 lf = lf.filter(~pl.col("region").is_in(exclude_regions))125 126 lf = lf.with_columns(127 pl.col("region").str.strip_chars().str.to_titlecase()128 )129 130 lf = lf.with_columns(131 pl.col("cost").fill_null(132 pl.col("cost").median().over("region")133 )134 )135 136 lf = lf.sort(["region", "order_date"])137 138 lf = lf.with_columns(139 pl.col("revenue")140 .rolling_mean(window_size=3, min_periods=1)141 .over("region")142 .alias("region_rolling_rev")143 )144 145 lf = lf.with_columns(146 ((pl.col("revenue") - pl.col("cost")) / pl.col("revenue")).alias("profit_margin")147 )148 149 return (150 lf.group_by(["region", "status"])151 .agg(152 [153 pl.col("revenue").sum().alias("revenue_sum"),154 pl.col("revenue").count().alias("revenue_count"),155 pl.col("profit_margin").mean().alias("margin_mean"),156 pl.col("region_rolling_rev").mean().alias("rolling_rev_mean"),157 ]158 )159 .sort(["revenue_sum", "region"], descending=[True, False])160 .collect()161 )162 """163 ).strip(),164}165DETERMINISTIC_SOLUTIONS["pandas_to_polars_advanced"] = DETERMINISTIC_SOLUTIONS["pandas_to_polars"]166 167 168def log_start(task: str, env_target: str, model: str) -> None:169 print(f"[START] task={task} env={env_target} model={model}", flush=True)170 171 172def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str] = None) -> None:173 done_val = str(done).lower()174 error_val = error if error else "null"175 action_preview = action.replace("\r", " ").replace("\n", " ")176 print(177 f"[STEP] step={step} action={action_preview!r} reward={reward:.2f} done={done_val} error={error_val}",178 flush=True,179 )180 181 182def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:183 success_val = str(success).lower()184 rewards_str = ",".join(f"{reward:.2f}" for reward in rewards)185 print(186 f"[END] success={success_val} steps={steps} score={score:.3f} rewards={rewards_str}",187 flush=True,188 )189 190 191def clamp_open_score(score: float) -> float:192 return max(SCORE_EPSILON, min(1.0 - SCORE_EPSILON, score))193 194 195def get_model_candidates() -> List[Tuple[str, str]]:196 candidates = [197 (MODEL_NAME, API_KEY),198 ]199 return [(model, key) for model, key in candidates if model and key]200 201 202def has_proxy_llm_config() -> bool:203 return bool(API_BASE_URL and MODEL_NAME and API_KEY)204 205 206def clean_json_response(text: str) -> str:207 text = text.strip()208 text = re.sub(r"^```json\s*", "", text, flags=re.IGNORECASE)209 text = re.sub(r"^```\s*", "", text)210 text = re.sub(r"```$", "", text)211 return text.strip()212 213 214def extract_json_object(text: str) -> Optional[Dict[str, str]]:215 cleaned = clean_json_response(text)216 try:217 parsed = json.loads(cleaned)218 if isinstance(parsed, dict):219 return parsed220 except Exception as exc:221 print(f"[DEBUG] JSON parse failed: {exc}", flush=True)222 223 match = re.search(r"\{.*\}", cleaned, flags=re.DOTALL)224 if not match:225 return None226 227 try:228 parsed = json.loads(match.group(0))229 except Exception:230 return None231 232 return parsed if isinstance(parsed, dict) else None233 234 235def build_prompt(obs) -> str:236 history = "\n".join(obs.history or [])237 info = getattr(obs, "info", {}) or {}238 acceptance_checks = "\n".join(f"- {item}" for item in info.get("acceptance_checks", [])) or "(not provided)"239 pitfalls = "\n".join(f"- {item}" for item in info.get("pitfalls", [])) or "(not provided)"240 task_specific_hints = []241 242 if getattr(obs, "task_id", "") == "python_to_node":243 task_specific_hints.extend(244 [245 "Use the global fetch API directly.",246 "Construct the request with URL and url.searchParams rather than returning a mocked object.",247 "Preserve headers Accept=application/json and X-Request-Source=openenv.",248 ]249 )250 251 if getattr(obs, "task_id", "") == "pandas_to_polars_advanced":252 task_specific_hints.extend(253 [254 "Target Polars 1.x compatible code.",255 "Prefer pl.scan_csv(filepath, try_parse_dates=True) or pl.scan_csv(filepath).",256 "Do not use parse_dates=True with pl.scan_csv.",257 "Return a collected Polars DataFrame at the end.",258 ]259 )260 261 task_hint_block = "\n".join(f"- {item}" for item in task_specific_hints) or "(none)"262 return f"""263You are a code migration expert.264 265Migrate the following {obs.source_language} code to {obs.target_language}.266 267Task ID: {obs.task_id}268Difficulty: {obs.difficulty}269 270Requirements:271{obs.requirements}272 273Test description:274{obs.test_description}275 276Business context:277{info.get("business_context", "(not provided)")}278 279Stakeholder request:280{info.get("stakeholder_request", "(not provided)")}281 282Acceptance checks:283{acceptance_checks}284 285Known pitfalls:286{pitfalls}287 288Task-specific implementation hints:289{task_hint_block}290 291Runtime budget:292{info.get("runtime_budget", "(not provided)")}293 294Attempts:295Used {info.get("attempts_used", 0)} of {info.get("max_attempts", MAX_STEPS)}.296Remaining: {info.get("attempts_remaining", MAX_STEPS)}297 298Previous attempts:299{history if history else "(none)"}300 301Source code:302```{obs.source_language}303{obs.source_code}304```305 306Return ONLY valid JSON with exactly these keys:307 308translated_code309explanation310 311Do not include markdown fences.312When prior feedback exists, revise the earlier attempt instead of restarting blindly.313""".strip()314 315 316def get_deterministic_action(obs) -> Optional[Dict[str, str]]:317 task_id = (getattr(obs, "task_id", "") or "").strip()318 solution = DETERMINISTIC_SOLUTIONS.get(task_id)319 320 if not solution and getattr(obs, "difficulty", "") == "hard":321 solution = DETERMINISTIC_SOLUTIONS["pandas_to_polars_advanced"]322 323 if not solution and getattr(obs, "source_language", "") == "python" and getattr(obs, "target_language", "") == "javascript":324 solution = DETERMINISTIC_SOLUTIONS["python_to_node"]325 326 if not solution:327 return None328 329 return {330 "translated_code": solution,331 "explanation": "Deterministic baseline generated from the environment requirements.",332 }333 334 335def build_fallback_translation(obs) -> str:336 if getattr(obs, "target_language", "") == "javascript":337 return dedent(338 """339 async function solveTask() {340 throw new Error("No solver available for this task");341 }342 """343 ).strip()344 345 return dedent(346 """347 def solve_task() -> None:348 raise RuntimeError("No solver available for this task")349 """350 ).strip()351 352 353async def call_model_for_action(model: str, api_key: str, prompt: str) -> Dict[str, str]:354 if OpenAI is None:355 return {356 "translated_code": "",357 "explanation": "openai package is unavailable in this runtime.",358 }359 360 client = OpenAI(base_url=API_BASE_URL, api_key=api_key)361 362 try:363 completion = client.chat.completions.create(364 model=model,365 temperature=0.2,366 max_tokens=2048,367 messages=[368 {369 "role": "system",370 "content": "Return only valid JSON with keys translated_code and explanation.",371 },372 {"role": "user", "content": prompt},373 ],374 )375 text = (completion.choices[0].message.content or "").strip()376 print(f"[DEBUG] Raw model response: {text[:200]}", flush=True)377 parsed = extract_json_object(text)378 if parsed:379 translated_code = str(parsed.get("translated_code", "")).strip()380 explanation = str(parsed.get("explanation", "")).strip()381 if translated_code:382 return {383 "translated_code": translated_code,384 "explanation": explanation or "LLM-generated translation.",385 }386 387 return {388 "translated_code": text,389 "explanation": "Raw model output could not be parsed as JSON.",390 }391 except Exception as exc:392 print(f"[DEBUG] Model call failed for {model}: {exc}", flush=True)393 return {394 "translated_code": "",395 "explanation": f"API error: {exc}",396 }397 398 399async def choose_action(obs, models: List[Tuple[str, str]]) -> Tuple[Dict[str, str], str]:400 prompt = build_prompt(obs)401 for model_name, api_key in models:402 result = await call_model_for_action(model_name, api_key, prompt)403 if result.get("translated_code"):404 return result, model_name405 406 deterministic = get_deterministic_action(obs)407 if deterministic:408 deterministic["explanation"] = (409 "Model call failed during task execution, so a deterministic fallback was used."410 )411 return deterministic, "deterministic-fallback"412 413 return {414 "translated_code": build_fallback_translation(obs),415 "explanation": "No model-based solver was available.",416 }, "fallback-stub"417 418 419def get_env_candidates() -> List[str]:420 explicit_target = (os.environ.get("IMAGE_NAME") or os.environ.get("ENV_URL") or "").strip()421 candidates = []422 423 if explicit_target:424 candidates.append(explicit_target)425 426 candidates.extend(427 [428 "openenv-code_migration",429 "openenv-code_migration:latest",430 "openenv-code_migration_env",431 "openenv-code_migration_env:latest",432 "code-migration-env",433 "code-migration-env:latest",434 "code_migration_env-env:latest",435 "http://127.0.0.1:8000",436 "http://localhost:8000",437 ]438 )439 440 deduped = []441 seen = set()442 for candidate in candidates:443 if not candidate or candidate in seen:444 continue445 seen.add(candidate)446 deduped.append(candidate)447 return deduped448 449 450async def create_env_client() -> Tuple[CodeMigrationEnv, str]:451 errors: List[str] = []452 453 for target in get_env_candidates():454 env = None455 try:456 if target.startswith("http://") or target.startswith("https://"):457 env = CodeMigrationEnv(base_url=target)458 else:459 env = await CodeMigrationEnv.from_docker_image(target)460 461 await env.reset()462 return env, target463 except Exception as exc:464 errors.append(f"{target}: {exc}")465 if env is not None:466 try:467 await env.close()468 except Exception:469 pass470 471 joined_errors = " | ".join(errors[:5]) if errors else "no env targets were configured"472 raise RuntimeError(f"Could not connect to any environment target. {joined_errors}")473 474 475def default_runner_label(task_name: str, models: List[Tuple[str, str]]) -> str:476 if models:477 return models[0][0]478 return "fallback-stub"479 480 481async def run_single_task_with_env(482 env: CodeMigrationEnv,483 task_name: str,484 episode_id: str,485 models: List[Tuple[str, str]],486 env_target: str,487) -> Dict[str, object]:488 log_start(task=task_name, env_target=env_target, model=default_runner_label(task_name, models))489 490 rewards: List[float] = []491 steps_taken = 0492 score = 0.0493 success = False494 495 try:496 result = await env.reset(episode_id=episode_id)497 498 for step in range(1, MAX_STEPS + 1):499 if result.done:500 break501 502 obs = result.observation503 action_data, runner_label = await choose_action(obs, models)504 505 try:506 action = CodeMigrationAction(507 translated_code=action_data["translated_code"],508 explanation=action_data["explanation"][:1999],509 )510 except Exception as exc:511 action = CodeMigrationAction(512 translated_code=build_fallback_translation(obs),513 explanation=f"Validation error: {exc}",514 )515 runner_label = "fallback-stub"516 517 result = await env.step(action)518 reward = result.reward or 0.0519 rewards.append(reward)520 steps_taken = step521 522 action_preview = action.translated_code[:100]523 if len(action.translated_code) > 100:524 action_preview += "..."525 526 feedback = None527 if getattr(result, "observation", None) and getattr(result.observation, "history", None):528 feedback = result.observation.history[-1]529 print(f"[DEBUG] feedback={feedback}", flush=True)530 531 log_step(532 step=step,533 action=action_preview,534 reward=reward,535 done=result.done,536 error=feedback,537 )538 539 if step == 1 and runner_label != default_runner_label(task_name, models):540 print(f"[DEBUG] runner={runner_label}", flush=True)541 542 if result.done:543 break544 545 # In the iterative setting each reward is a quality snapshot, not an546 # additive return. The best achieved reward is the task score.547 score = max(rewards) if rewards else SCORE_EPSILON548 score = clamp_open_score(score / MAX_TOTAL_REWARD if MAX_TOTAL_REWARD > 0 else score)549 success = score >= SUCCESS_SCORE_THRESHOLD550 except Exception as exc:551 print(f"[ERROR] Task {task_name} failed: {exc}", flush=True)552 score = SCORE_EPSILON553 554 log_end(success=success, steps=steps_taken, score=score, rewards=rewards)555 return {556 "task": task_name,557 "success": success,558 "steps": steps_taken,559 "score": score,560 "rewards": rewards,561 }562 563 564async def main() -> None:565 models = get_model_candidates()566 task_results = []567 print(568 "[DEBUG] LLM config "569 f"base_url={bool(API_BASE_URL)} model={bool(MODEL_NAME)} api_key={bool(API_KEY)}",570 flush=True,571 )572 573 try:574 env, env_target = await create_env_client()575 except Exception as exc:576 env_target = (os.environ.get("IMAGE_NAME") or os.environ.get("ENV_URL") or "unavailable").strip() or "unavailable"577 print(f"[FATAL] Failed to initialize env: {exc}", flush=True)578 for task_name, _ in TASKS:579 log_start(task=task_name, env_target=env_target, model=default_runner_label(task_name, models))580 log_end(success=False, steps=0, score=SCORE_EPSILON, rewards=[])581 return582 583 try:584 for task_name, episode_id in TASKS:585 task_results.append(586 await run_single_task_with_env(587 env=env,588 task_name=task_name,589 episode_id=episode_id,590 models=models,591 env_target=env_target,592 )593 )594 finally:595 try:596 await env.close()597 except Exception as exc:598 print(f"[DEBUG] env.close() error: {exc}", flush=True)599 600 print("\n[SUMMARY] Task Results:", flush=True)601 for result in task_results:602 print(603 f" {result['task']}: success={result['success']} score={result['score']:.3f}",604 flush=True,605 )606 607 608if __name__ == "__main__":609 asyncio.run(main())610 