PRANAV05092003/autonomous-code-refactoring-env
0
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 