Team Ai
Apppublic

PRANAV05092003/autonomous-code-refactoring-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
server.py642 linesDownload Raw Back to root
1"""2Legacy server runner.3 4OpenEnv validation expects the FastAPI app to be importable at:5  server.app:app6 7This file is kept as a thin runner for local execution and Docker CMD.8"""9 10from __future__ import annotations11 12import os13import uvicorn14 15from server.app import app  # noqa: F401  (re-export for convenience)16 17 18def get_env() -> OpenEnvRefactorEnv:19    global _env20    if _env is None:21        _env = OpenEnvRefactorEnv(registry=registry)22    return _env23 24 25def _state_response() -> StateResponse:26    return get_env().state()27 28 29def _choose_action_heuristic(code: str, task_id: Optional[str]) -> int:30    has_generic = re.search(r"\b(x|tmp|i)\b", code) is not None31    has_if_false = re.search(r"\bif\s+False\b", code) is not None32    has_if_true = re.search(r"\bif\s+True\b", code) is not None33    has_append_loop = ".append(" in code and "for " in code34    has_double_not = "not not" in code35    has_add_call = "add(" in code36 37    if task_id == "rename_variables":38        if has_generic:39            return 040        if has_if_false or "unused" in code:41            return 142        if has_append_loop:43            return 244        if has_if_true or has_double_not:45            return 346        return 447 48    if task_id == "remove_dead_code":49        if has_if_false or "unused" in code:50            return 151        if has_append_loop:52            return 253        if has_if_true or has_double_not:54            return 355        if has_generic:56            return 057        return 458 59    if has_generic:60        return 061    if has_append_loop:62        return 263    if has_if_false or has_if_true or has_double_not:64        return 365    if has_add_call:66        return 467    return 168 69 70def _choose_action_llm(71    *,72    code: str,73    task_id: Optional[str],74    step_index: int,75    max_steps: int,76    api_base_url: str,77    model_name: str,78    api_token: str,79) -> tuple[int, str, str]:80    if not api_token.strip():81        return _choose_action_heuristic(code, task_id), "empty token -> heuristic", "heuristic"82 83    client = OpenAI(base_url=api_base_url, api_key=api_token)84    messages = [85        {86            "role": "system",87            "content": (88                "You are a code-refactoring action selector. Return ONLY compact JSON: "89                '{"action": <0-4>, "reason": "..."}.\n'90                "Actions: 0=rename_variable,1=remove_dead_code,2=simplify_loop,3=optimize_condition,4=inline_function"91            ),92        },93        {94            "role": "user",95            "content": (96                f"task_id={task_id or 'auto'}\n"97                f"step={step_index}/{max_steps}\n"98                "Current code:\n"99                f"```python\n{code}\n```"100            ),101        },102    ]103    try:104        resp = client.chat.completions.create(105            model=model_name,106            messages=messages,107            temperature=0.0,108            max_tokens=120,109        )110        raw = (resp.choices[0].message.content or "").strip()111        m = re.search(r"\{.*\}", raw, flags=re.DOTALL)112        blob = m.group(0) if m else raw113        parsed = json.loads(blob)114        action = int(parsed.get("action", -1))115        reason = str(parsed.get("reason", "llm-selected action"))116        if 0 <= action <= 4:117            return action, reason, "llm"118    except Exception as exc:119        return _choose_action_heuristic(code, task_id), f"llm error -> heuristic: {exc}", "heuristic"120 121    return _choose_action_heuristic(code, task_id), "invalid llm output -> heuristic", "heuristic"122 123 124def _choose_action_rl(observation: list[float], model_path: str) -> tuple[Optional[int], str, str]:125    if PPO is None:126        return None, "stable-baselines3 unavailable", "rl"127    if not os.path.exists(model_path):128        return None, f"rl model not found: {model_path}", "rl"129 130    try:131        model = _rl_model_cache.get(model_path)132        if model is None:133            model = PPO.load(model_path)134            _rl_model_cache[model_path] = model135 136        obs = np.asarray(observation, dtype=np.float32)137        action, _ = model.predict(obs, deterministic=True)138        action_i = int(action)139        if 0 <= action_i <= 4:140            return action_i, "rl policy action", "rl"141        return None, f"invalid rl action: {action_i}", "rl"142    except Exception as exc:143        return None, f"rl failure: {exc}", "rl"144 145 146def _demo_html() -> str:147    return """<!doctype html>148<html lang=\"en\">149<head>150    <meta charset=\"utf-8\" />151    <meta name=\"viewport\" content=\"width=device-width, initial-scale=1\" />152    <title>ACRE Refactor Demo</title>153    <style>154        @import url('https://fonts.googleapis.com/css2?family=Space+Grotesk:wght@400;600;700&display=swap');155        :root {156            --bg0: #0b1f2a;157            --bg1: #14344a;158            --ink: #eaf7ff;159            --muted: #a7c8db;160            --brand: #1ec28b;161            --warn: #ffcb47;162            --panel: rgba(8, 24, 36, 0.72);163            --stroke: rgba(140, 197, 225, 0.35);164        }165        * { box-sizing: border-box; }166        body {167            margin: 0;168            color: var(--ink);169            font-family: 'Space Grotesk', sans-serif;170            background:171                radial-gradient(circle at 12% 18%, rgba(30, 194, 139, 0.28), transparent 35%),172                radial-gradient(circle at 88% 8%, rgba(255, 203, 71, 0.22), transparent 30%),173                linear-gradient(150deg, var(--bg0), var(--bg1));174            min-height: 100vh;175        }176        .wrap {177            max-width: 1200px;178            margin: 0 auto;179            padding: 28px 20px 40px;180        }181        h1 {182            margin: 0 0 6px;183            font-size: clamp(1.6rem, 2vw + 1rem, 2.6rem);184            letter-spacing: 0.2px;185        }186        .sub { margin: 0 0 20px; color: var(--muted); }187        .grid {188            display: grid;189            grid-template-columns: 1fr;190            gap: 16px;191        }192        .panel {193            border: 1px solid var(--stroke);194            border-radius: 14px;195            background: var(--panel);196            backdrop-filter: blur(4px);197            padding: 14px;198        }199        .controls {200            display: grid;201            grid-template-columns: 1fr 1fr;202            gap: 8px;203            margin-bottom: 10px;204        }205        textarea, pre {206            width: 100%;207            min-height: 260px;208            border: 1px solid var(--stroke);209            border-radius: 10px;210            padding: 12px;211            background: rgba(1, 13, 24, 0.82);212            color: #dcf4ff;213            font-family: Consolas, 'Courier New', monospace;214            font-size: 13px;215            line-height: 1.4;216            overflow: auto;217            white-space: pre;218        }219        button, select {220            border: 1px solid var(--stroke);221            border-radius: 10px;222            padding: 10px 12px;223            background: rgba(11, 36, 52, 0.9);224            color: var(--ink);225            font-weight: 600;226        }227        button.primary {228            background: linear-gradient(120deg, #19a7ff, #1ec28b);229            color: #032235;230            border: none;231        }232        .cols {233            display: grid;234            grid-template-columns: 1fr;235            gap: 14px;236        }237        .meta {238            color: var(--muted);239            font-size: 0.92rem;240            margin-top: 8px;241        }242        .badge {243            color: #082b22;244            background: var(--brand);245            border-radius: 999px;246            padding: 2px 9px;247            font-size: 12px;248            font-weight: 700;249        }250        .warn {251            color: #2a1c00;252            background: var(--warn);253        }254        @media (min-width: 900px) {255            .cols { grid-template-columns: 1fr 1fr; }256        }257    </style>258</head>259<body>260    <div class=\"wrap\">261        <h1>ACRE Live Refactor Arena</h1>262        <p class=\"sub\">Paste old code, run the agent, and compare before and after with a full diff and step-by-step rewards.</p>263 264        <div class=\"panel\">265            <div class=\"controls\">266                <button onclick=\"loadExample(1)\">Load Example 1</button>267                <button onclick=\"loadExample(2)\">Load Example 2</button>268                <select id=\"task\">269                    <option value=\"\">Auto strategy</option>270                    <option value=\"rename_variables\">rename_variables</option>271                    <option value=\"remove_dead_code\">remove_dead_code</option>272                    <option value=\"full_refactor\">full_refactor</option>273                </select>274                <button class=\"primary\" onclick=\"runOptimize()\">Run Optimization</button>275            </div>276            <div class=\"controls\" style=\"margin-bottom: 10px;\">277                <select id=\"mode\">278                    <option value=\"rl_then_llm\">RL First -> LLM Fallback</option>279                    <option value=\"heuristic\">Heuristic Agent (no API key)</option>280                    <option value=\"llm\">LLM Agent (OpenAI-compatible API)</option>281                </select>282                <input id=\"rlModelPath\" placeholder=\"RL model path\" value=\"acre_agent.zip\" style=\"border:1px solid var(--stroke);border-radius:10px;padding:10px 12px;background:rgba(1,13,24,0.82);color:#dcf4ff;\" />283                <input id=\"baseUrl\" placeholder=\"API base URL (optional)\" value=\"https://api.openai.com/v1\" style=\"border:1px solid var(--stroke);border-radius:10px;padding:10px 12px;background:rgba(1,13,24,0.82);color:#dcf4ff;\" />284                <input id=\"modelName\" placeholder=\"Model name (optional)\" value=\"gpt-4o-mini\" style=\"border:1px solid var(--stroke);border-radius:10px;padding:10px 12px;background:rgba(1,13,24,0.82);color:#dcf4ff;\" />285                <input id=\"apiToken\" type=\"password\" placeholder=\"Paste API token here for LLM mode\" style=\"border:1px solid var(--stroke);border-radius:10px;padding:10px 12px;background:rgba(1,13,24,0.82);color:#dcf4ff;\" />286            </div>287            <div class=\"controls\" style=\"margin-bottom: 10px;\">288                <label style=\"display:flex;align-items:center;gap:8px;padding:8px 10px;border:1px solid var(--stroke);border-radius:10px;\">289                    <input id=\"autoSuggest\" type=\"checkbox\" />290                    Auto suggest after typing pause291                </label>292            </div>293            <textarea id=\"input\" spellcheck=\"false\" placeholder=\"Paste your Python code here...\"></textarea>294            <p class=\"meta\" id=\"status\">Status: ready</p>295            <p class=\"meta\" id=\"liveResults\">Live results: loading...</p>296        </div>297 298        <div class=\"cols\" style=\"margin-top: 14px\">299            <div class=\"panel\">300                <h3>Original Code</h3>301                <pre id=\"original\"></pre>302            </div>303            <div class=\"panel\">304                <h3>Optimized Code</h3>305                <pre id=\"optimized\"></pre>306            </div>307        </div>308 309        <div class=\"panel\" style=\"margin-top: 14px\">310            <h3>Diff</h3>311            <pre id=\"diff\"></pre>312        </div>313 314        <div class=\"panel\" style=\"margin-top: 14px\">315            <h3>Step Logs</h3>316            <pre id=\"steps\"></pre>317        </div>318    </div>319 320    <script>321        const EX1 = `def compute(x, y, tmp):\n    tmp = x + y\n    x = tmp * 2\n    result = x\n    return result\n`;322        const EX2 = `def add(p, q):\n    return p + q\n\ndef compute(x, data, tmp):\n    result = []\n    for item in data:\n        result.append(item * 2)\n    if False:\n        y = 999\n    if True:\n        val = add(x, tmp)\n    unused = 0\n    flag = not not True\n    return val\n    print(\"dead\")\n`;323        let autoTimer = null;324 325        function loadExample(i) {326            document.getElementById('input').value = i === 1 ? EX1 : EX2;327            document.getElementById('status').textContent = `Status: loaded example ${i}`;328        }329 330        async function runOptimize() {331            const code = document.getElementById('input').value;332            const task = document.getElementById('task').value || null;333            const mode = document.getElementById('mode').value;334            const useRl = mode === 'rl_then_llm';335            const useLlm = mode === 'llm' || mode === 'rl_then_llm';336            const fallbackToLlm = mode === 'rl_then_llm';337            const rlModelPath = document.getElementById('rlModelPath').value || null;338            const apiToken = document.getElementById('apiToken').value || null;339            const apiBaseUrl = document.getElementById('baseUrl').value || null;340            const modelName = document.getElementById('modelName').value || null;341            if (!code.trim()) {342                document.getElementById('status').innerHTML = 'Status: <span class=\"badge warn\">please paste code first</span>';343                return;344            }345            if (mode === 'llm' && (!apiToken || !apiToken.trim())) {346                document.getElementById('status').innerHTML = 'Status: <span class=\"badge warn\">paste API token for LLM mode</span>';347                return;348            }349 350            document.getElementById('status').textContent = 'Status: running optimization...';351            try {352                const res = await fetch('/optimize', {353                    method: 'POST',354                    headers: {'Content-Type': 'application/json'},355                    body: JSON.stringify({356                        code,357                        task_id: task,358                        max_steps: 5,359                        use_rl: useRl,360                        use_llm: useLlm,361                        fallback_to_llm: fallbackToLlm,362                        rl_model_path: rlModelPath,363                        api_base_url: apiBaseUrl,364                        model_name: modelName,365                        api_token: apiToken,366                    })367                });368                const data = await res.json();369                if (!res.ok) {370                    throw new Error(data.detail || 'request failed');371                }372 373                document.getElementById('original').textContent = data.original_code;374                document.getElementById('optimized').textContent = data.optimized_code;375                document.getElementById('diff').textContent = data.diff || '(no diff)';376                document.getElementById('steps').textContent = JSON.stringify(data.steps, null, 2);377 378                const scoreText = data.task_score === null ? 'n/a' : data.task_score;379                document.getElementById('status').innerHTML = `Status: <span class=\"badge\">done</span> cumulative_reward=${data.cumulative_reward.toFixed(2)} task_score=${scoreText}`;380            } catch (err) {381                document.getElementById('status').innerHTML = `Status: <span class=\"badge warn\">error</span> ${err.message}`;382            }383        }384 385        async function loadLiveResults() {386            const el = document.getElementById('liveResults');387            try {388                const res = await fetch('/demo');389                const data = await res.json();390                const r = (data && data.results) ? data.results : null;391                if (!res.ok || !r) {392                    throw new Error('demo request failed');393                }394                const easy = (r.easy ?? 0).toFixed(4);395                const medium = (r.medium ?? 0).toFixed(4);396                const hard = (r.hard ?? 0).toFixed(4);397                const final = (r.final ?? 0).toFixed(4);398                el.textContent = `Live results: Easy=${easy}  Medium=${medium}  Hard=${hard}  Final=${final}`;399            } catch (err) {400                el.textContent = `Live results: error (${err.message || err})`;401            }402        }403 404        loadExample(1);405        loadLiveResults();406        document.getElementById('input').addEventListener('input', () => {407            if (!document.getElementById('autoSuggest').checked) {408                return;409            }410            if (autoTimer) {411                clearTimeout(autoTimer);412            }413            autoTimer = setTimeout(() => {414                runOptimize();415            }, 1200);416        });417    </script>418</body>419</html>"""420 421 422# ---------------------------------------------------------------------------423# Routes424# ---------------------------------------------------------------------------425 426@app.get("/", response_class=HTMLResponse)427def root() -> HTMLResponse:428    """429    Hugging Face Space homepage.430 431    Serve the interactive UI so opening the Space shows a real demo page.432    The live JSON execution results remain available at `GET /demo`.433    """434    return HTMLResponse(content=_demo_html())435 436 437@app.get("/health", response_model=CompatibilityHealthResponse)438def health_compat() -> CompatibilityHealthResponse:439    """Compatibility health route used by some OpenEnv reference environments."""440    return CompatibilityHealthResponse(status="healthy", service="acre-env")441 442 443@app.get("/demo")444def demo() -> JSONResponse:445    """Run all tasks and return JSON results."""446    from inference import run_all_tasks447 448    return JSONResponse(content={"results": run_all_tasks()})449 450 451@app.get("/ui", response_class=HTMLResponse)452def demo_ui() -> HTMLResponse:453    """Alias for the interactive UI (same as `/`)."""454    return HTMLResponse(content=_demo_html())455 456 457@app.post("/reset", response_model=ResetResponse)458def reset(req: ResetRequest = ResetRequest()) -> ResetResponse:459    """Reset the environment. Optionally load a task's initial code."""460    env = get_env()461    try:462        obs = env.reset(seed=req.seed, task_id=req.task_id, code=req.code)463    except ValueError as exc:464        raise HTTPException(status_code=404, detail=str(exc)) from exc465    return ResetResponse(466        observation=obs,467        observation_vector=obs.to_vector(),468        info=env.last_reset_info,469        task_id=req.task_id,470        state=_state_response(),471    )472 473 474@app.post("/step", response_model=StepResponse)475def step(req: StepRequest) -> StepResponse:476    """Take one refactoring step."""477    env = get_env()478    if not (0 <= req.action <= 4):479        raise HTTPException(status_code=400, detail="action must be 0–4")480 481    obs, reward, done, info = env.step(req.action)482    action_name = str(info.get("action_name", env.action_meanings.get(req.action, "unknown")))483 484    return StepResponse(485        action=ActionModel(action=req.action, action_name=action_name),486        observation=obs,487        observation_vector=obs.to_vector(),488        reward=reward,489        done=done,490        terminated=done,491        truncated=False,492        info=info,493        state=_state_response(),494    )495 496 497@app.get("/state", response_model=StateResponse)498def state() -> StateResponse:499    """Return full current environment state (OpenEnv spec requirement)."""500    return _state_response()501 502 503@app.get("/tasks", response_model=TasksResponse)504def list_tasks() -> TasksResponse:505    """Enumerate all tasks (easy → medium → hard)."""506    return TasksResponse(tasks=[TaskInfo.model_validate(t) for t in registry.list_tasks()])507 508 509@app.post("/tasks/{task_id}/grade", response_model=GradeResponse)510def grade(task_id: str, req: GradeRequest) -> GradeResponse:511    """Grade submitted code against a task's grader (returns score 0.0–1.0)."""512    task = registry.get_task(task_id)513    if task is None:514        raise HTTPException(status_code=404, detail=f"Task '{task_id}' not found")515    # Use the deterministic expected-output grader for the public grade endpoint.516    score = task.grade_against_expected(req.code)517    return GradeResponse(518        task_id=task_id,519        score=round(score, 4),520        passed=score >= 0.8,521    )522 523 524@app.post("/optimize", response_model=OptimizeResponse)525def optimize(req: OptimizeRequest) -> OptimizeResponse:526    """Run a full optimization episode and return code comparison artifacts."""527    code = req.code.strip("\n")528    if not code.strip():529        raise HTTPException(status_code=400, detail="code must be non-empty")530 531    env = get_env()532    try:533        env.reset(task_id=req.task_id, code=code)534    except ValueError as exc:535        raise HTTPException(status_code=404, detail=str(exc)) from exc536 537    steps: list[OptimizationStep] = []538    cumulative_reward = 0.0539 540    for step_idx in range(1, req.max_steps + 1):541        state_now = env.state()542        current_code = state_now.current_code543        obs_list = [float(x) for x in state_now.observation_vector]544 545        action: int546        reason: str547        source: str548 549        if req.use_rl:550            rl_action, rl_reason, rl_source = _choose_action_rl(551                observation=obs_list,552                model_path=req.rl_model_path or DEFAULT_RL_MODEL_PATH,553            )554            if rl_action is not None:555                action, reason, source = rl_action, rl_reason, rl_source556            elif req.fallback_to_llm and req.use_llm:557                action, reason, source = _choose_action_llm(558                    code=current_code,559                    task_id=req.task_id,560                    step_index=step_idx,561                    max_steps=req.max_steps,562                    api_base_url=req.api_base_url or DEFAULT_API_BASE_URL,563                    model_name=req.model_name or DEFAULT_MODEL_NAME,564                    api_token=req.api_token or "",565                )566                reason = f"{rl_reason}; {reason}"567            else:568                action = _choose_action_heuristic(current_code, req.task_id)569                reason = f"{rl_reason}; heuristic fallback"570                source = "heuristic"571        elif req.use_llm:572            action, reason, source = _choose_action_llm(573                code=current_code,574                task_id=req.task_id,575                step_index=step_idx,576                max_steps=req.max_steps,577                api_base_url=req.api_base_url or DEFAULT_API_BASE_URL,578                model_name=req.model_name or DEFAULT_MODEL_NAME,579                api_token=req.api_token or "",580            )581        else:582            action = _choose_action_heuristic(current_code, req.task_id)583            reason = "heuristic policy"584            source = "heuristic"585 586        _, reward, done, info = env.step(action)587        state_now = env.state()588 589        cumulative_reward += float(reward.raw)590        steps.append(591            OptimizationStep(592                step=step_idx,593                action=action,594                action_name=info.get("action_name", "unknown"),595                reason=reason,596                source=source,597                reward=float(reward.raw),598                normalized_reward=float(reward.normalized),599                changed=bool(info.get("changed", False)),600                complexity=float(state_now.complexity),601            )602        )603 604        if done:605            break606 607    final_code = str(env.state().current_code)608    diff_lines = difflib.unified_diff(609        code.splitlines(),610        final_code.splitlines(),611        fromfile="original.py",612        tofile="optimized.py",613        lineterm="",614    )615    diff_text = "\n".join(diff_lines)616 617    task_score: Optional[float] = None618    if req.task_id:619        task = registry.get_task(req.task_id)620        if task is None:621            raise HTTPException(status_code=404, detail=f"Task '{req.task_id}' not found")622        task_score = round(task.grade(final_code), 4)623 624    return OptimizeResponse(625        original_code=code,626        optimized_code=final_code,627        diff=diff_text,628        steps=steps,629        cumulative_reward=round(cumulative_reward, 4),630        task_id=req.task_id,631        task_score=task_score,632    )633 634 635# ---------------------------------------------------------------------------636# Entry point637# ---------------------------------------------------------------------------638 639if __name__ == "__main__":640    port = int(os.getenv("PORT", 7860))641    uvicorn.run("server.app:app", host="0.0.0.0", port=port)642