Team Ai
Apppublic

razak123/code-migration-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py610 linesDownload Raw Back to root
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