PRANAV05092003/autonomous-code-refactoring-env
0
1"""2ACRE inference script for OpenEnv submission evaluation.3 4Environment variables:5 - API_BASE_URL: LLM API endpoint injected by evaluator6 - MODEL_NAME: model identifier (default allowed)7 - API_KEY: API token for the OpenAI-compatible proxy endpoint8 - ENV_URL: running ACRE server base URL (required)9 - LOCAL_IMAGE_NAME: present for evaluator compatibility (optional)10 - USE_LLM: set to "0" to disable LLM action selection11 12STRICT stdout format (do not change):13 [START] task=<task_id>14 [STEP] action=<action_int>15 [END] task=<task_id> score=<score_float>16"""17from __future__ import annotations18 19import json20import os21import re22import sys23import time24from typing import Dict, List, Optional, Tuple25 26import requests27from openai import OpenAI28 29MODEL_NAME = os.getenv("MODEL_NAME") or "gpt-4o-mini"30# Phase-2 validator expects API_KEY through provided proxy.31API_KEY = os.getenv("API_KEY")32ENV_URL: str = os.getenv("ENV_URL", "http://localhost:7860")33LOCAL_IMAGE_NAME: str | None = os.getenv("LOCAL_IMAGE_NAME")34 35TASKS: List[str] = ["rename_variables", "remove_dead_code", "full_refactor"]36 37ACTION_MEANINGS: Dict[int, str] = {38 0: "rename_variable",39 1: "remove_dead_code",40 2: "simplify_loop",41 3: "optimize_condition",42 4: "inline_function",43}44 45SYSTEM_PROMPT = """\46You are an RL agent that refactors Python code. Choose one action per step.47 48Actions:49 0 rename_variable - rename generic names (x, tmp, i) to descriptive ones50 1 remove_dead_code - remove unreachable stmts, if False blocks, unused vars51 2 simplify_loop - convert append-loops to list comprehensions52 3 optimize_condition- simplify 'not not x', 'if True/False', 'x==True'53 4 inline_function - inline simple single-return module-level functions54 55Respond ONLY with valid JSON (no markdown):56{"action": <0-4>, "reason": "<one sentence>"}"""57 58SAFE_FALLBACK_SCORES: Dict[str, float] = {59 "easy": 0.0,60 "medium": 0.0,61 "hard": 0.0,62 "final": 0.0,63}64 65 66def _safe_scores() -> Dict[str, float]:67 return dict(SAFE_FALLBACK_SCORES)68 69 70def _env_url() -> str:71 # Never crash due to missing env var.72 return str(ENV_URL or "http://localhost:7860").rstrip("/")73 74 75def _post(path: str, payload: dict | None = None) -> dict:76 try:77 response = requests.post(f"{_env_url()}{path}", json=payload or {}, timeout=5)78 response.raise_for_status()79 return response.json()80 except Exception:81 print("Warning: Could not reach environment", file=sys.stderr)82 return {}83 84 85def _get(path: str) -> dict:86 try:87 response = requests.get(f"{_env_url()}{path}", timeout=5)88 response.raise_for_status()89 return response.json()90 except Exception:91 print("Warning: Could not reach environment", file=sys.stderr)92 return {}93 94 95def reset_env(task_id: str) -> dict:96 return _post("/reset", {"task_id": task_id})97 98 99def step_env(action: int) -> dict:100 return _post("/step", {"action": action})101 102 103def get_state() -> dict:104 return _get("/state")105 106 107def grade(task_id: str, code: str) -> float:108 try:109 response = requests.post(110 f"{_env_url()}/tasks/{task_id}/grade",111 json={"code": code},112 timeout=5,113 )114 response.raise_for_status()115 return float(response.json().get("score", 0.0))116 except Exception:117 print("Warning: Could not reach environment", file=sys.stderr)118 return 0.0119 120 121def choose_action(client: Optional[OpenAI], state: dict, task_id: str) -> Tuple[int, str]:122 def heuristic_action() -> Tuple[int, str]:123 code = str(state.get("current_code", ""))124 step_i = int(state.get("episode_steps", 0))125 126 has_generic = re.search(r"\b(x|tmp|i)\b", code) is not None127 has_if_false = re.search(r"\bif\s+False\b", code) is not None128 has_if_true = re.search(r"\bif\s+True\b", code) is not None129 has_append_loop = ".append(" in code and "for " in code130 has_double_not = "not not" in code131 has_add_call = "add(" in code132 133 if task_id == "rename_variables":134 if has_generic:135 return 0, "heuristic: remove generic names first"136 if has_if_false or "unused" in code:137 return 1, "heuristic: remove dead code"138 if has_append_loop:139 return 2, "heuristic: simplify loop"140 if has_if_true or has_double_not:141 return 3, "heuristic: optimize conditions"142 return 4, "heuristic: inline simple function"143 144 if task_id == "remove_dead_code":145 if has_if_false or "unused" in code:146 return 1, "heuristic: remove dead code patterns"147 if has_append_loop:148 return 2, "heuristic: convert append-loop"149 if has_if_true or has_double_not:150 return 3, "heuristic: simplify conditions"151 if has_generic:152 return 0, "heuristic: clean generic names"153 return 4, "heuristic: inline helper"154 155 if has_generic:156 return 0, "heuristic: rename generic variables"157 if has_append_loop:158 return 2, "heuristic: simplify loop into listcomp"159 if has_if_false or has_if_true or has_double_not:160 return 3, "heuristic: optimize boolean branches"161 if has_add_call:162 return 4, "heuristic: inline add() call"163 if step_i >= 2:164 return 1, "heuristic: remove remaining dead code"165 return 3, "heuristic: condition optimization as safe default"166 167 # Enable LLM by default when credentials are present.168 use_llm = bool(API_KEY) and os.getenv("USE_LLM", "1") == "1"169 if (not use_llm) or client is None:170 return heuristic_action()171 172 messages = [173 {"role": "system", "content": SYSTEM_PROMPT},174 {175 "role": "user",176 "content": (177 f"Task: {task_id}\n"178 f"Steps remaining: {state.get('max_steps', 5) - state.get('episode_steps', 0)}\n"179 f"Complexity: {state.get('complexity', 0)}\n\n"180 f"Current code:\n```python\n{state.get('current_code', '')}\n```\n\n"181 "Choose the best action."182 ),183 },184 ]185 try:186 response = client.chat.completions.create(187 model=MODEL_NAME,188 messages=messages,189 temperature=0.0,190 max_tokens=120,191 )192 raw = (response.choices[0].message.content or "").strip()193 json_blob = raw194 195 if "{" not in json_blob or "}" not in json_blob:196 return heuristic_action()197 198 match = re.search(r"\{.*\}", json_blob, flags=re.DOTALL)199 if match:200 json_blob = match.group(0)201 202 parsed = json.loads(json_blob)203 action = int(parsed.get("action", -1))204 reason = str(parsed.get("reason", ""))205 if 0 <= action <= 4:206 return action, reason or "llm-selected action"207 return heuristic_action()208 except Exception:209 return heuristic_action()210 211 212def _build_openai_client() -> Optional[OpenAI]:213 """214 Build OpenAI-compatible client using hackathon-required proxy env vars.215 Falls back safely when vars are absent in local runs.216 """217 base_url = os.getenv("API_BASE_URL")218 api_key = os.getenv("API_KEY")219 220 if not base_url or not api_key:221 return None222 223 try:224 return OpenAI(base_url=base_url, api_key=api_key)225 except Exception:226 return None227 228 229def _touch_proxy(client: Optional[OpenAI]) -> None:230 """231 Ensure at least one request is sent through the provided proxy in Phase-2.232 """233 if client is None:234 return None235 try:236 client.chat.completions.create(237 model=MODEL_NAME,238 messages=[{"role": "user", "content": "Return exactly: ok"}],239 temperature=0.0,240 max_tokens=2,241 )242 except Exception:243 # Keep inference resilient even if proxy is temporarily unavailable.244 return None245 return None246 247 248def run_episode(client: Optional[OpenAI], task_id: str, episode_num: int) -> float:249 reset_env(task_id)250 state = get_state()251 252 # STRICT logging format required by evaluator.253 print(f"[START] task={task_id}", flush=True)254 255 cumulative_reward = 0.0256 257 for step_num in range(1, 6):258 action, reason = choose_action(client, state, task_id)259 result = step_env(action)260 state = get_state()261 262 reward_payload = result.get("reward", {})263 raw_reward = float(reward_payload.get("raw", 0.0))264 norm_reward = float(reward_payload.get("normalized", (raw_reward + 32) / 52))265 cumulative_reward += raw_reward266 267 # STRICT logging format required by evaluator.268 print(f"[STEP] action={int(action)}", flush=True)269 270 if result.get("done") or result.get("terminated") or result.get("truncated"):271 break272 273 final_state = get_state()274 task_score = grade(task_id, final_state.get("current_code", ""))275 276 # STRICT logging format required by evaluator.277 print(f"[END] task={task_id} score={task_score:.4f}", flush=True)278 279 return task_score280 281 282def run_all_tasks() -> Dict[str, float]:283 """284 Run all three tasks and return deterministic scores.285 286 This is used by the FastAPI server to show live demo results on the Space.287 """288 try:289 # Prefer local in-process execution when running inside the server (no ENV_URL needed).290 try:291 from acre.tasks.task_registry import TaskRegistry292 from openenv_interface import OpenEnvRefactorEnv293 except Exception:294 TaskRegistry = None # type: ignore[assignment]295 OpenEnvRefactorEnv = None # type: ignore[assignment]296 297 registry = TaskRegistry() if TaskRegistry is not None else None298 env = OpenEnvRefactorEnv(registry=registry) if OpenEnvRefactorEnv is not None else None299 300 client = _build_openai_client()301 _touch_proxy(client)302 303 task_plan = [304 "rename_variables",305 "remove_dead_code",306 "full_refactor",307 ]308 309 results: Dict[str, float] = _safe_scores()310 scores: List[float] = []311 312 # If we have a local env, use it. Otherwise fall back to HTTP.313 if env is None or registry is None:314 # Network safety: quick health probe before running.315 try:316 r = requests.get(f"{_env_url()}/health", timeout=5)317 r.raise_for_status()318 except Exception:319 print("Warning: Could not reach environment", file=sys.stderr)320 return _safe_scores()321 322 for task_id in task_plan:323 print(f"[START] task={task_id}", flush=True)324 reset_env(task_id)325 for _ in range(5):326 state = get_state()327 action, _reason = choose_action(client, state, task_id)328 print(f"[STEP] action={int(action)}", flush=True)329 step_env(action)330 final_state = get_state()331 score = float(grade(task_id, final_state.get("current_code", "")))332 print(f"[END] task={task_id} score={float(score):.4f}", flush=True)333 scores.append(score)334 if task_id == "rename_variables":335 results["easy"] = score336 elif task_id == "remove_dead_code":337 results["medium"] = score338 else:339 results["hard"] = score340 341 results["final"] = float(sum(scores) / len(scores)) if scores else 0.0342 return results343 344 else:345 # Local in-process execution (fast + no network recursion).346 for task_id in task_plan:347 print(f"[START] task={task_id}", flush=True)348 env.reset(seed=0, task_id=task_id)349 for _ in range(5):350 st = env.state()351 state_payload = {352 "current_code": str(st.current_code),353 "episode_steps": int(st.episode_steps),354 "max_steps": int(st.max_steps),355 "complexity": float(st.complexity),356 }357 action, _reason = choose_action(client, state_payload, task_id)358 action = int(action)359 print(f"[STEP] action={int(action)}", flush=True)360 env.step(action)361 st = env.state()362 task = registry.get_task(task_id)363 score = float(task.grade_against_expected(st.current_code)) if task is not None else 0.0364 print(f"[END] task={task_id} score={float(score):.4f}", flush=True)365 scores.append(score)366 if task_id == "rename_variables":367 results["easy"] = score368 elif task_id == "remove_dead_code":369 results["medium"] = score370 else:371 results["hard"] = score372 373 results["final"] = float(sum(scores) / len(scores)) if scores else 0.0374 return results375 except Exception as e:376 print(f"ERROR: {str(e)}", file=sys.stderr)377 return _safe_scores()378 379 380def main() -> None:381 # Never crash. Always produce output.382 result = run_all_tasks()383 print(f"Easy: {float(result.get('easy', 0.0)):.4f}", file=sys.stderr)384 print(f"Medium: {float(result.get('medium', 0.0)):.4f}", file=sys.stderr)385 print(f"Hard: {float(result.get('hard', 0.0)):.4f}", file=sys.stderr)386 print(f"Final: {float(result.get('final', 0.0)):.4f}", file=sys.stderr)387 return None388 389 390if __name__ == "__main__":391 try:392 run_all_tasks()393 except Exception as e:394 print(f"Fatal error: {e}", file=sys.stderr)395 