Spaces:
Sleeping
Sleeping
| """ConfigDebugEnvironment - OpenEnv-compatible environment class. | |
| Inherits from openenv.core.env_server.Environment and implements | |
| the standard reset/step/state interface with multi-task logic. | |
| """ | |
| from typing import Optional, Any | |
| from uuid import uuid4 | |
| from openenv.core.env_server import Environment | |
| from server.models import ConfigDebugAction, ConfigDebugObservation, ConfigDebugState | |
| from server.tasks.task_registry import get_task, TASK_ORDER | |
| MAX_STEPS_PER_TASK = 5 | |
| class ConfigDebugEnvironment(Environment): | |
| """Multi-task config debugging environment. | |
| Manages 7 sequential tasks internally. Each WebSocket session | |
| (via create_fastapi_app) gets its own instance with independent state. | |
| """ | |
| SUPPORTS_CONCURRENT_SESSIONS = True | |
| def __init__(self): | |
| self._init_episode() | |
| def _init_episode(self): | |
| self.task_ids = list(TASK_ORDER) | |
| self.current_task_index = 0 | |
| self.current_step = 0 | |
| self.total_reward = 0.0 | |
| self._done = False | |
| self.tasks_completed: list = [] | |
| 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 | |
| self._episode_id = str(uuid4()) | |
| self._global_step = 0 | |
| # ---- OpenEnv interface methods ---- | |
| def reset(self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs: Any) -> ConfigDebugObservation: | |
| """Reset environment to initial state (task 1).""" | |
| print(f"[SESSION_DEBUG] reset() called on env object_id={id(self)}, episode_id={episode_id}") | |
| print(f"[PROGRESSION_DEBUG] RESET - initializing new episode") | |
| self._init_episode() | |
| if episode_id: | |
| self._episode_id = episode_id | |
| obs = self._build_observation() | |
| print(f"[SESSION_DEBUG] reset() complete, env_id={id(self)}, _episode_id={self._episode_id}") | |
| print(f"[PROGRESSION_DEBUG] RESET_DONE - task_id={self._current_task_id()}, current_task_index={self.current_task_index}") | |
| return obs | |
| def step(self, action: ConfigDebugAction, timeout_s: Optional[float] = None, **kwargs: Any) -> ConfigDebugObservation: | |
| """Process an action: run the grader, advance tasks if done.""" | |
| # DEBUG: Log environment instance ID and current state | |
| print(f"[SESSION_DEBUG] step() called on env object_id={id(self)}, _episode_id={self._episode_id}") | |
| print(f"[SESSION_DEBUG] step() state: current_step={self.current_step}, current_task_index={self.current_task_index}, _done={self._done}") | |
| if self._done: | |
| return self._build_observation() | |
| task_id = self._current_task_id() | |
| task = get_task(task_id) | |
| # DEBUG: Log before step | |
| print(f"[PROGRESSION_DEBUG] BEFORE_STEP - task_id={task_id}, current_step={self.current_step}, current_task_index={self.current_task_index}") | |
| # Run the grader - now returns FLOAT ONLY for validator compatibility | |
| grader_result = task.grader(action.fixed_config) | |
| # Handle result gracefully (should be float) | |
| if isinstance(grader_result, float): | |
| reward = grader_result | |
| elif isinstance(grader_result, tuple) and len(grader_result) > 0: | |
| # Fallback for legacy tuple format | |
| reward = grader_result[0] | |
| else: | |
| reward = 0.01 # Safe default | |
| reward = max(0.01, min(0.99, reward)) | |
| self.current_step += 1 | |
| self._global_step += 1 | |
| self.bugs_found_so_far = 0 # Default since grader no longer returns this | |
| self.previous_reward = round(reward, 4) | |
| self.current_error_message = "" # Default since grader no longer returns this | |
| # DEBUG: Log progression check | |
| print(f"[PROGRESSION_DEBUG] PROGRESSION_CHECK - reward={reward:.4f}, current_step={self.current_step}, max_steps={MAX_STEPS_PER_TASK}") | |
| # Check if task is complete | |
| task_done = reward >= 0.99 or self.current_step >= MAX_STEPS_PER_TASK | |
| print(f"[PROGRESSION_DEBUG] TASK_DONE={task_done} (reward >= 0.99: {reward >= 0.99}, step >= {MAX_STEPS_PER_TASK}: {self.current_step >= MAX_STEPS_PER_TASK})") | |
| if task_done: | |
| print(f"[PROGRESSION_DEBUG] ADVANCING - task_id={task_id} complete, advancing from index {self.current_task_index} to {self.current_task_index + 1}") | |
| self.total_reward += reward | |
| self.tasks_completed.append(task_id) | |
| self.current_task_index += 1 | |
| self.current_step = 0 | |
| self.bugs_found_so_far = 0 | |
| self.current_error_message = None | |
| self.current_broken_config = None | |
| print(f"[PROGRESSION_DEBUG] NEW_TASK - current_task_index={self.current_task_index}, new_task_id={self._current_task_id()}") | |
| if self.current_task_index >= len(self.task_ids): | |
| self._done = True | |
| print(f"[PROGRESSION_DEBUG] EPISODE_COMPLETE - all {len(self.task_ids)} tasks done") | |
| else: | |
| self.current_broken_config = action.fixed_config | |
| obs = self._build_observation() | |
| obs.done = self._done | |
| obs.reward = round(reward, 4) | |
| return obs | |
| def state(self) -> ConfigDebugState: | |
| """Return current environment state with enhanced RL signals.""" | |
| tasks_remaining = self.task_ids[self.current_task_index:] | |
| if self._done: | |
| tasks_remaining = [] | |
| total_tasks = len(self.task_ids) | |
| completed_tasks = len(self.tasks_completed) | |
| progress_ratio = completed_tasks / total_tasks if total_tasks > 0 else 0.0 | |
| current_task = get_task(self._current_task_id()) | |
| return ConfigDebugState( | |
| episode_id=self._episode_id, | |
| step_count=self._global_step, | |
| current_task_id=self._current_task_id(), | |
| current_step=self.current_step, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| total_reward=round(self.total_reward, 4), | |
| is_done=self._done, | |
| tasks_completed=list(self.tasks_completed), | |
| tasks_remaining=tasks_remaining, | |
| # Enhanced RL signals | |
| bugs_found_so_far=self.bugs_found_so_far, | |
| current_error_message=self.current_error_message, | |
| progress_ratio=round(progress_ratio, 2), | |
| current_difficulty=current_task.difficulty, | |
| ) | |
| # ---- Internal helpers ---- | |
| def _current_task_id(self) -> str: | |
| if self.current_task_index < len(self.task_ids): | |
| return self.task_ids[self.current_task_index] | |
| return self.task_ids[-1] | |
| def _build_observation(self) -> ConfigDebugObservation: | |
| task_id = self._current_task_id() | |
| task = get_task(task_id) | |
| broken = self.current_broken_config if self.current_broken_config is not None else task.broken_config | |
| error = self.current_error_message if self.current_error_message is not None else task.error_message | |
| return ConfigDebugObservation( | |
| broken_config=broken, | |
| ground_truth=task.ground_truth, | |
| file_type=task.file_type, | |
| error_message=error, | |
| task_id=task.task_id, | |
| task_description=task.description, | |
| difficulty=task.difficulty, | |
| num_bugs=task.num_bugs, | |
| bugs_found_so_far=self.bugs_found_so_far, | |
| previous_reward=self.previous_reward, | |
| done=self._done, | |
| reward=self.previous_reward, | |
| ) | |