Team Ai
Apppublic

aady161103/reflection-debug-agent

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
inference.py373 linesDownload Raw Back to root
1"""2Inference Script — Reflection Debug Agent3===================================4MANDATORY5- Before submitting, ensure the following variables are defined in your environment configuration:6    API_BASE_URL   The API endpoint for the LLM.7    MODEL_NAME     The model identifier to use for inference.8    HF_TOKEN       Your Hugging Face / API key.9    LOCAL_IMAGE_NAME The name of the local image to use for the environment if you are using from_docker_image()10                     method11 12- Defaults are set only for API_BASE_URL and MODEL_NAME13    (and should reflect your active inference setup):14    API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")15    MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")16 17- The inference script must be named `inference.py` and placed in the root directory of the project18- Participants must use OpenAI Client for all LLM calls using above variables19 20STDOUT FORMAT21- The script must emit exactly three line types to stdout, in this order:22 23    [START] task=<task_name> env=<benchmark> model=<model_name>24    [STEP]  step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>25    [END]   success=<true|false> steps=<n> rewards=<r1,r2,...,rn>26 27  Rules:28    - One [START] line at episode begin.29    - One [STEP] line per step, immediately after env.step() returns.30    - One [END] line after env.close(), always emitted (even on exception).31    - reward and rewards are formatted to 2 decimal places.32    - done and success are lowercase booleans: true or false.33    - error is the raw last_action_error string, or null if none.34    - All fields on a single line with no newlines within a line.35"""36 37import asyncio38import json39import os40import textwrap41from typing import List, Optional42 43from openai import OpenAI44 45# Load .env for local development only (does NOT override existing env vars)46from dotenv import load_dotenv47load_dotenv(override=False)48 49# --- Environment Client Import ---50# When running with from_docker_image(), the OpenEnv framework generates51# a typed client. For direct HTTP mode, we use httpx.52import httpx53 54# --- Configuration ---55IMAGE_NAME = os.getenv("IMAGE_NAME")56API_KEY = os.getenv("API_KEY") or os.getenv("HF_TOKEN")57API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")58MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")59BENCHMARK = "reflection_debug_agent"60MAX_STEPS = 861TEMPERATURE = 0.762MAX_TOKENS = 204863 64# Task list to iterate through65TASKS = ["api_json_fix", "csv_processor_fix", "retry_decorator_fix"]66 67SYSTEM_PROMPT = textwrap.dedent("""\68    You are an expert software engineer debugging code.69    You will be given buggy code and test results.70    71    You must respond with a JSON object containing exactly these fields:72    {73        "hypothesis": "Why the bug exists — reference specific code constructs, variable names, or logic errors.",74        "action_description": "What you will change and why — describe the concrete modification.",75        "expected_result": "What you expect after the fix — which tests should now pass and why.",76        "edits": [{"search": "old code...", "replace": "new code..."}]77    }78    79    IMPORTANT:80    - "edits" must be a list of objects, each containing a "search" and "replace" string.81    - "search" MUST be an EXACT match of a block of code, including all spaces and indentation.82    - "replace" is the exact string to substitute it with.83    - INDENTATION MATTERS: Your "replace" string must include the exact leading spaces/tabs required for proper Python indentation!84    - To insert code, make "search" match the surrounding lines and include them in "replace".85    - Be specific in your hypothesis.86    - Respond ONLY with the JSON object, no other text.87""")88 89 90# --- Logging Functions (EXACTLY matching required format) ---91 92def log_start(task: str, env: str, model: str) -> None:93    print(f"[START] task={task} env={env} model={model}", flush=True)94 95 96def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:97    error_val = error if error else "null"98    done_val = str(done).lower()99    # Clamp reward to strictly (0, 1)100    clamped = min(max(reward, 0.01), 0.99)101    print(102        f"[STEP] step={step} action={action} reward={clamped:.2f} done={done_val} error={error_val}",103        flush=True,104    )105 106 107def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:108    # Clamp all values to strictly (0, 1)109    clamped_score = min(max(score, 0.01), 0.99)110    clamped_rewards = [min(max(r, 0.01), 0.99) for r in rewards]111    rewards_str = ",".join(f"{r:.2f}" for r in clamped_rewards)112    print(113        f"[END] success={str(success).lower()} steps={steps} score={clamped_score:.3f} rewards={rewards_str}",114        flush=True,115    )116 117 118# --- LLM Interaction ---119 120def build_user_prompt(buggy_code: str, test_output: str, step: int, history: List[str]) -> str:121    history_block = "\n".join(history[-4:]) if history else "None"122    return textwrap.dedent(f"""\123        Step: {step}124        125        Current Code:126        ```python127        {buggy_code}128        ```129        130        Test Results:131        {test_output}132        133        Previous attempts:134        {history_block}135        136        Analyze the bug, fix the code, and provide your structured reflection as JSON.137    """)138 139 140def get_model_response(client: OpenAI, buggy_code: str, test_output: str, step: int, history: List[str]) -> dict:141    """Call the LLM and parse its JSON response."""142    user_prompt = build_user_prompt(buggy_code, test_output, step, history)143    try:144        completion = client.chat.completions.create(145            model=MODEL_NAME,146            messages=[147                {"role": "system", "content": SYSTEM_PROMPT},148                {"role": "user", "content": user_prompt},149            ],150            temperature=TEMPERATURE,151            max_tokens=MAX_TOKENS,152            stream=False,153        )154        text = (completion.choices[0].message.content or "").strip()155 156        # Extract JSON from the response (handle markdown code blocks)157        if "```json" in text:158            text = text.split("```json")[1].split("```")[0].strip()159        elif "```" in text:160            text = text.split("```")[1].split("```")[0].strip()161 162        parsed = json.loads(text)163        return {164            "edits": parsed.get("edits", []),165            "hypothesis": parsed.get("hypothesis", "No hypothesis provided"),166            "action_description": parsed.get("action_description", "No action described"),167            "expected_result": parsed.get("expected_result", "No expected result"),168        }169    except json.JSONDecodeError as e:170        print(f"[DEBUG] JSON parse failed: {e}", flush=True)171        return {172            "edits": [],173            "hypothesis": "Failed to parse LLM response",174            "action_description": "No changes made due to parse error",175            "expected_result": "No improvement expected",176        }177    except Exception as exc:178        print(f"[DEBUG] Model request failed: {exc}", flush=True)179        return {180            "edits": [],181            "hypothesis": f"LLM call failed: {exc}",182            "action_description": "No changes possible",183            "expected_result": "No improvement expected",184        }185 186 187# --- Environment Interaction (HTTP-based) ---188 189class DebugEnvClient:190    """Simple HTTP client for the debug environment."""191 192    def __init__(self, base_url: str):193        self.base_url = base_url.rstrip("/")194        self.client = httpx.Client(timeout=120.0)195        self.session_id = None196 197    def reset(self, task_name: str) -> dict:198        """Reset the environment for a new task."""199        resp = self.client.post(200            f"{self.base_url}/reset",201            json={"task_name": task_name, "session_id": self.session_id},202        )203        resp.raise_for_status()204        data = resp.json()205        self.session_id = data.get("session_id")206        return data207 208    def step(self, action: dict) -> dict:209        """Take a step in the environment."""210        action["session_id"] = self.session_id211        resp = self.client.post(f"{self.base_url}/step", json=action)212        resp.raise_for_status()213        return resp.json()214 215    def state(self) -> dict:216        """Get current state."""217        resp = self.client.get(218            f"{self.base_url}/state",219            params={"session_id": self.session_id or ""},220        )221        resp.raise_for_status()222        return resp.json()223 224    def close(self):225        """Close the HTTP client."""226        self.client.close()227 228    def wait_until_ready(self, retries: int = 10, delay: float = 3.0) -> bool:229        """Wait for the env container to be reachable."""230        import time as _time231        for attempt in range(retries):232            try:233                resp = self.client.get(f"{self.base_url}/health")234                if resp.status_code == 200:235                    print(f"[DEBUG] Env ready after {attempt + 1} attempt(s)", flush=True)236                    return True237            except Exception:238                pass239            print(f"[DEBUG] Env not ready, retry {attempt + 1}/{retries}...", flush=True)240            _time.sleep(delay)241        return False242 243 244def run_task(client: OpenAI, env: DebugEnvClient, task_name: str) -> tuple:245    """Run a single task episode. Returns (success, steps, score, rewards)."""246    rewards: List[float] = []247    steps_taken = 0248    history: List[str] = []249 250    log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)251 252    try:253        # Reset environment254        reset_data = env.reset(task_name)255        obs = reset_data["observation"]256        buggy_code = obs["buggy_code"]257        test_output = obs["test_output"]258 259        for step in range(1, MAX_STEPS + 1):260            if obs.get("done", False):261                break262 263            # Get LLM's fix + reflection264            response = get_model_response(client, buggy_code, test_output, step, history)265 266            # Take a step267            step_data = env.step({268                "edits": response["edits"],269                "hypothesis": response["hypothesis"],270                "action_description": response["action_description"],271                "expected_result": response["expected_result"],272            })273 274            obs = step_data["observation"]275            reward = step_data.get("reward", 0.0) or 0.0276            reward = min(max(reward, 0.01), 0.99)277            done = step_data.get("done", False)278            error = obs.get("last_action_error")279 280            rewards.append(reward)281            steps_taken = step282            buggy_code = obs["buggy_code"]283            test_output = obs["test_output"]284 285            # Format action string for log (abbreviated)286            action_str = f"fix({response['hypothesis'][:50]})"287 288            log_step(step=step, action=action_str, reward=reward, done=done, error=error)289 290            error_info = ""291            if error:292                error_info = f" ⚠️ EDITS REJECTED: {error[:100]}"293            history.append(294                f"Step {step}: {response['hypothesis'][:80]} -> tests {obs.get('tests_passed', 0)}/{obs.get('tests_total', 0)}{error_info}"295            )296 297            if done:298                break299 300        # Compute final score301        score = sum(rewards) / len(rewards) if rewards else 0.01302        score = min(max(score, 0.01), 0.99)303        success = score >= 0.5304 305        return success, steps_taken, score, rewards306 307    except Exception as exc:308        print(f"[DEBUG] Task {task_name} failed: {exc}", flush=True)309        return False, steps_taken, 0.01, rewards310 311 312def main() -> None:313    """Run inference across all tasks."""314    print(f"[DEBUG] Initializing OpenAI client...", flush=True)315    print(f"[DEBUG] API_BASE_URL={API_BASE_URL}", flush=True)316    print(f"[DEBUG] API_KEY={'set (' + API_KEY[:8] + '...)' if API_KEY else 'NOT SET'}", flush=True)317    print(f"[DEBUG] MODEL_NAME={MODEL_NAME}", flush=True)318 319    client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY or "not-set")320 321    # Warmup: make a guaranteed LLM call BEFORE any environment interaction322    try:323        print("[DEBUG] Making warmup LLM call...", flush=True)324        warmup = client.chat.completions.create(325            model=MODEL_NAME,326            messages=[{"role": "user", "content": "Say OK"}],327            max_tokens=5,328        )329        print(f"[DEBUG] Warmup LLM call succeeded: {warmup.choices[0].message.content}", flush=True)330    except Exception as exc:331        print(f"[DEBUG] Warmup LLM call FAILED: {exc}", flush=True)332 333    # Connect to environment334    env_url = os.getenv("ENV_URL", "http://localhost:7860")335    print(f"[DEBUG] ENV_URL={env_url}", flush=True)336    env = DebugEnvClient(env_url)337 338    # Wait for the env container to be ready339    if not env.wait_until_ready(retries=15, delay=3.0):340        print("[DEBUG] FATAL: Env container never became ready", flush=True)341        # Emit required output so validator doesn't flag missing format342        for task_name in TASKS:343            log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)344            log_end(success=False, steps=0, score=0.01, rewards=[0.01])345        env.close()346        return347 348    all_success = True349 350    try:351        for task_name in TASKS:352            success, steps, score, rewards = run_task(client, env, task_name)353            log_end(success=success, steps=steps, score=score, rewards=rewards)354 355            if not success:356                all_success = False357 358            print(f"[DEBUG] Task {task_name}: score={score:.3f}, success={success}", flush=True)359 360    except Exception as exc:361        print(f"[DEBUG] Unhandled error: {exc}", flush=True)362    finally:363        env.close()364 365    print(f"[DEBUG] All tasks complete. Overall success: {all_success}", flush=True)366 367 368if __name__ == "__main__":369    try:370        main()371    except Exception as exc:372        print(f"[DEBUG] FATAL unhandled exception: {exc}", flush=True)373