Spaces:
Sleeping
Sleeping
| from fastapi import FastAPI, HTTPException | |
| from typing import Optional | |
| import json | |
| import gradio as gr | |
| from server.models import ConfigDebugAction, ConfigDebugObservation, ConfigDebugState | |
| from server.tasks.task_registry import get_task, get_all_task_ids, TASK_ORDER | |
| app = FastAPI(title="ConfigDebugEnv", version="1.0.0") | |
| # --- Environment State --- | |
| MAX_STEPS_PER_TASK = 5 | |
| class EnvironmentState: | |
| """Mutable environment state that persists across requests.""" | |
| def __init__(self): | |
| self.reset_state() | |
| def reset_state(self): | |
| self.task_ids = list(TASK_ORDER) | |
| self.current_task_index = 0 | |
| self.current_step = 0 | |
| self.total_reward = 0.0 | |
| self.is_done = False | |
| self.tasks_completed = [] | |
| self.bugs_found_so_far = 0 | |
| self.previous_reward = 0.0 | |
| self.current_error_message: Optional[str] = None | |
| self.current_broken_config: Optional[str] = None | |
| env_state = EnvironmentState() | |
| def _get_current_task_id() -> str: | |
| if env_state.current_task_index < len(env_state.task_ids): | |
| return env_state.task_ids[env_state.current_task_index] | |
| return env_state.task_ids[-1] | |
| def _build_observation() -> ConfigDebugObservation: | |
| task_id = _get_current_task_id() | |
| task = get_task(task_id) | |
| broken_config = ( | |
| env_state.current_broken_config | |
| if env_state.current_broken_config is not None | |
| else task.broken_config | |
| ) | |
| error_message = ( | |
| env_state.current_error_message | |
| if env_state.current_error_message is not None | |
| else task.error_message | |
| ) | |
| return ConfigDebugObservation( | |
| broken_config=broken_config, | |
| file_type=task.file_type, | |
| error_message=error_message, | |
| task_id=task.task_id, | |
| task_description=task.description, | |
| difficulty=task.difficulty, | |
| num_bugs=task.num_bugs, | |
| bugs_found_so_far=env_state.bugs_found_so_far, | |
| previous_reward=env_state.previous_reward, | |
| ) | |
| def _build_state() -> ConfigDebugState: | |
| task_id = _get_current_task_id() | |
| tasks_remaining = env_state.task_ids[env_state.current_task_index:] | |
| if env_state.is_done: | |
| tasks_remaining = [] | |
| return ConfigDebugState( | |
| current_task_id=task_id, | |
| current_step=env_state.current_step, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| total_reward=round(env_state.total_reward, 4), | |
| is_done=env_state.is_done, | |
| tasks_completed=list(env_state.tasks_completed), | |
| tasks_remaining=tasks_remaining, | |
| ) | |
| # --- API Endpoints --- | |
| def info(): | |
| return {"name": "ConfigDebugEnv", "version": "1.0.0", "status": "running"} | |
| def health(): | |
| return {"status": "ok"} | |
| def reset(task_id: str = None): | |
| """Reset the environment to initial state, return first observation.""" | |
| env_state.reset_state() | |
| if task_id and task_id in env_state.task_ids: | |
| env_state.current_task_index = env_state.task_ids.index(task_id) | |
| observation = _build_observation() | |
| state = _build_state() | |
| return { | |
| "observation": observation.model_dump(), | |
| "reward": 0.0, | |
| "done": False, | |
| "state": state.model_dump(), | |
| "info": { | |
| "task_id": _get_current_task_id(), | |
| "tasks_total": len(env_state.task_ids), | |
| }, | |
| } | |
| def step(action: ConfigDebugAction): | |
| """ | |
| Take a step in the environment. | |
| The agent submits a fixed config, the grader evaluates it. | |
| """ | |
| if env_state.is_done: | |
| raise HTTPException( | |
| status_code=400, | |
| detail="Environment is done. Call /reset to start a new episode.", | |
| ) | |
| task_id = _get_current_task_id() | |
| task = get_task(task_id) | |
| # Run the grader | |
| reward, error_message, bugs_fixed = task.grader(action.fixed_config) | |
| reward = max(0.01, min(0.90, reward)) | |
| env_state.current_step += 1 | |
| env_state.bugs_found_so_far = len(bugs_fixed) | |
| env_state.previous_reward = round(reward, 4) | |
| # Update error message for feedback | |
| env_state.current_error_message = error_message | |
| # Check if task is complete (perfect score or max steps reached) | |
| task_done = reward >= 0.85 or env_state.current_step >= MAX_STEPS_PER_TASK | |
| task_reward = reward # Reward for this step | |
| if task_done: | |
| # Record the best reward for this task | |
| env_state.total_reward += reward | |
| env_state.tasks_completed.append(task_id) | |
| env_state.current_task_index += 1 | |
| # Reset per-task state | |
| env_state.current_step = 0 | |
| env_state.bugs_found_so_far = 0 | |
| env_state.current_error_message = None | |
| env_state.current_broken_config = None | |
| # Check if all tasks are done | |
| if env_state.current_task_index >= len(env_state.task_ids): | |
| env_state.is_done = True | |
| else: | |
| # If the agent submitted something, use it as the new "broken" config | |
| # so the agent can iterate | |
| env_state.current_broken_config = action.fixed_config | |
| observation = _build_observation() | |
| state = _build_state() | |
| return { | |
| "observation": observation.model_dump(), | |
| "reward": round(task_reward, 4), | |
| "done": env_state.is_done, | |
| "state": state.model_dump(), | |
| "info": { | |
| "task_id": task_id, | |
| "bugs_fixed": bugs_fixed, | |
| "error_message": error_message, | |
| "task_done": task_done, | |
| }, | |
| } | |
| def state(): | |
| """Return current environment state.""" | |
| return _build_state().model_dump() | |
| def observation(): | |
| """Return current observation.""" | |
| if env_state.is_done: | |
| raise HTTPException( | |
| status_code=400, | |
| detail="Environment is done. Call /reset to start a new episode.", | |
| ) | |
| return _build_observation().model_dump() | |
| def metadata(): | |
| return { | |
| "env_name": "config_debug_env", | |
| "version": "1.0.0", | |
| "description": "Config file debugging environment", | |
| "tasks": [ | |
| {"id": "task1_json", "difficulty": "easy", "num_bugs": 2, "has_grader": True, "grader": "server.graders.grader_api:grade_task1"}, | |
| {"id": "task2_yaml", "difficulty": "easy", "num_bugs": 2, "has_grader": True, "grader": "server.graders.grader_api:grade_task2"}, | |
| {"id": "task3_dockerfile", "difficulty": "medium", "num_bugs": 3, "has_grader": True, "grader": "server.graders.grader_api:grade_task3"}, | |
| {"id": "task4_compose", "difficulty": "medium", "num_bugs": 4, "has_grader": True, "grader": "server.graders.grader_api:grade_task4"}, | |
| {"id": "task5_k8s", "difficulty": "hard", "num_bugs": 5, "has_grader": True, "grader": "server.graders.grader_api:grade_task5"}, | |
| {"id": "task6_github_actions", "difficulty": "hard", "num_bugs": 5, "has_grader": True, "grader": "server.graders.grader_api:grade_task6"}, | |
| {"id": "task7_nginx", "difficulty": "very_hard", "num_bugs": 6, "has_grader": True, "grader": "server.graders.grader_api:grade_task7"}, | |
| ], | |
| "action_model": "ConfigDebugAction", | |
| "observation_model": "ConfigDebugObservation", | |
| } | |
| def tasks(): | |
| return { | |
| "tasks": [ | |
| { | |
| "id": tid, | |
| "name": get_task(tid).description, | |
| "difficulty": get_task(tid).difficulty, | |
| "num_bugs": get_task(tid).num_bugs, | |
| "has_grader": True, | |
| "grader": f"server.graders.grader_api:grade_{tid}", | |
| } | |
| for tid in TASK_ORDER | |
| ], | |
| "total_tasks": len(TASK_ORDER), | |
| "tasks_with_graders": len(TASK_ORDER), | |
| } | |
| def schema(): | |
| return { | |
| "action": ConfigDebugAction.model_json_schema(), | |
| "observation": ConfigDebugObservation.model_json_schema(), | |
| "state": ConfigDebugState.model_json_schema(), | |
| } | |
| # --- Gradio Web UI --- | |
| def ui_reset(): | |
| env_state.reset_state() | |
| obs = _build_observation() | |
| st = _build_state() | |
| return ( | |
| f"Task: {obs.task_id} | Difficulty: {obs.difficulty} | Bugs: {obs.num_bugs}", | |
| obs.task_description, | |
| obs.broken_config, | |
| obs.error_message, | |
| json.dumps(st.model_dump(), indent=2), | |
| "Environment reset. Submit a fixed config to begin.", | |
| ) | |
| def ui_step(fixed_config): | |
| if env_state.is_done: | |
| st = _build_state() | |
| return ( | |
| "All tasks completed!", | |
| "", | |
| "", | |
| "Episode done. Click Reset to start again.", | |
| json.dumps(st.model_dump(), indent=2), | |
| f"Final score: {env_state.total_reward:.1f} / {len(TASK_ORDER)}.0", | |
| ) | |
| task_id = _get_current_task_id() | |
| task = get_task(task_id) | |
| reward, error_message, bugs_fixed = task.grader(fixed_config) | |
| reward = max(0.01, min(0.90, reward)) | |
| env_state.current_step += 1 | |
| env_state.bugs_found_so_far = len(bugs_fixed) | |
| env_state.previous_reward = round(reward, 4) | |
| env_state.current_error_message = error_message | |
| task_done = reward >= 0.85 or env_state.current_step >= MAX_STEPS_PER_TASK | |
| if task_done: | |
| env_state.total_reward += reward | |
| env_state.tasks_completed.append(task_id) | |
| env_state.current_task_index += 1 | |
| env_state.current_step = 0 | |
| env_state.bugs_found_so_far = 0 | |
| env_state.current_error_message = None | |
| env_state.current_broken_config = None | |
| if env_state.current_task_index >= len(env_state.task_ids): | |
| env_state.is_done = True | |
| else: | |
| env_state.current_broken_config = fixed_config | |
| obs = _build_observation() | |
| st = _build_state() | |
| history = f"Reward: {reward:.2f} | Bugs fixed: {bugs_fixed} | Task done: {task_done}\nFeedback: {error_message}" | |
| return ( | |
| f"Task: {obs.task_id} | Difficulty: {obs.difficulty} | Bugs: {obs.num_bugs}", | |
| obs.task_description, | |
| obs.broken_config, | |
| obs.error_message, | |
| json.dumps(st.model_dump(), indent=2), | |
| history, | |
| ) | |
| def ui_get_state(): | |
| st = _build_state() | |
| return json.dumps(st.model_dump(), indent=2) | |
| with gr.Blocks(title="ConfigDebugEnv", theme=gr.themes.Soft()) as demo: | |
| gr.Markdown("# ConfigDebugEnv") | |
| gr.Markdown("An RL environment for debugging broken config files across 7 real-world formats: JSON, YAML, Dockerfile, docker-compose, Kubernetes, GitHub Actions, nginx.") | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| gr.Markdown("### HumanAgent Interface") | |
| task_info = gr.Textbox(label="Current Task", interactive=False) | |
| task_desc = gr.Textbox(label="Task Description", interactive=False, lines=2) | |
| broken_config = gr.Textbox(label="Broken Config", interactive=False, lines=10) | |
| error_msg = gr.Textbox(label="Error Message", interactive=False, lines=2) | |
| gr.Markdown("### Take Action") | |
| fixed_config_input = gr.Textbox(label="Your Fixed Config", placeholder="Paste your fixed configuration here...", lines=10) | |
| with gr.Row(): | |
| reset_btn = gr.Button("Reset Environment", variant="secondary") | |
| step_btn = gr.Button("Step", variant="primary") | |
| state_btn = gr.Button("Get State", variant="secondary") | |
| with gr.Column(scale=1): | |
| gr.Markdown("### State Observer") | |
| state_display = gr.Textbox(label="Current State", interactive=False, lines=12) | |
| history_display = gr.Textbox(label="Action History / Reward", interactive=False, lines=4) | |
| reset_btn.click( | |
| fn=ui_reset, | |
| outputs=[task_info, task_desc, broken_config, error_msg, state_display, history_display], | |
| ) | |
| step_btn.click( | |
| fn=ui_step, | |
| inputs=[fixed_config_input], | |
| outputs=[task_info, task_desc, broken_config, error_msg, state_display, history_display], | |
| ) | |
| state_btn.click( | |
| fn=ui_get_state, | |
| outputs=[state_display], | |
| ) | |
| app = gr.mount_gradio_app(app, demo, path="/") |