Spaces:
Sleeping
Sleeping
| """ | |
| SQL Query Environment — Core logic. | |
| Implements the OpenEnv Environment interface: | |
| - reset() → creates fresh DB, picks a task, returns initial observation | |
| - step() → receives SQL, executes it, grades it, returns observation | |
| - state() → returns current episode state | |
| """ | |
| import sqlite3 | |
| import sys | |
| import os | |
| from uuid import uuid4 | |
| from openenv.core.env_server.interfaces import Environment | |
| from openenv.core.env_server.types import State | |
| sys.path.insert(0, os.path.dirname(__file__)) | |
| from database import create_database, execute_query, SCHEMA_DESCRIPTION | |
| from tasks import ALL_TASKS, TASK_LIST, TaskDefinition | |
| from grader import grade_query | |
| # Import models | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| from models import SQLAction, SQLObservation | |
| MAX_STEPS_PER_TASK = 3 | |
| class SQLQueryEnvironment(Environment): | |
| """ | |
| An environment where an AI agent writes SQL queries to answer | |
| natural language questions about a company database. | |
| """ | |
| def __init__(self): | |
| super().__init__() | |
| self._db: sqlite3.Connection | None = None | |
| self._state = State(episode_id=str(uuid4()), step_count=0) | |
| self._current_task: TaskDefinition | None = None | |
| self._step_count: int = 0 | |
| self._done: bool = False | |
| self._best_reward: float = 0.0 | |
| def reset(self, **kwargs) -> SQLObservation: | |
| """Reset environment: fresh DB, pick task, return initial observation.""" | |
| if self._db is not None: | |
| self._db.close() | |
| self._db = create_database() | |
| task_id = kwargs.get("task_id", None) | |
| if task_id and task_id in ALL_TASKS: | |
| self._current_task = ALL_TASKS[task_id] | |
| else: | |
| self._current_task = TASK_LIST[0] | |
| self._state = State(episode_id=str(uuid4()), step_count=0) | |
| self._step_count = 0 | |
| self._done = False | |
| self._best_reward = 0.0 | |
| return SQLObservation( | |
| task_id=self._current_task.task_id, | |
| task_description=self._current_task.description, | |
| difficulty=self._current_task.difficulty, | |
| schema_description=SCHEMA_DESCRIPTION, | |
| query_result="", | |
| query_error=False, | |
| feedback="Environment reset. Submit a SQL query to answer the question.", | |
| reward=0.0, | |
| done=False, | |
| step_count=0, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| ) | |
| def step(self, action: SQLAction) -> SQLObservation: | |
| """Execute agent's SQL, grade it, return observation.""" | |
| if self._done: | |
| return SQLObservation( | |
| task_id=self._current_task.task_id if self._current_task else "", | |
| task_description="", | |
| difficulty="", | |
| schema_description=SCHEMA_DESCRIPTION, | |
| query_result="", | |
| query_error=False, | |
| feedback="Episode is already complete. Call reset() to start a new one.", | |
| reward=0.0, | |
| done=True, | |
| step_count=self._step_count, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| ) | |
| self._step_count += 1 | |
| self._state.step_count = self._step_count | |
| # Validate task_id | |
| if action.task_id not in ALL_TASKS: | |
| return SQLObservation( | |
| task_id=action.task_id, | |
| task_description="", | |
| difficulty="", | |
| schema_description=SCHEMA_DESCRIPTION, | |
| query_result="", | |
| query_error=True, | |
| feedback=f"Unknown task_id: {action.task_id}. Valid: task_1, task_2, task_3", | |
| reward=0.0, | |
| done=False, | |
| step_count=self._step_count, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| ) | |
| task = ALL_TASKS[action.task_id] | |
| self._current_task = task | |
| # Execute agent's SQL | |
| agent_rows, agent_cols, error = execute_query(self._db, action.sql_query) | |
| # Grade it | |
| score, feedback = grade_query(task, agent_rows, agent_cols, error) | |
| if score > self._best_reward: | |
| self._best_reward = score | |
| # Format result for observation | |
| if error: | |
| result_str = f"ERROR: {error}" | |
| has_error = True | |
| elif agent_rows is not None and agent_cols is not None: | |
| header = " | ".join(agent_cols) | |
| separator = "-" * len(header) | |
| row_strs = [" | ".join(str(v) for v in row) for row in agent_rows[:20]] | |
| result_str = f"{header}\n{separator}\n" + "\n".join(row_strs) | |
| if len(agent_rows) > 20: | |
| result_str += f"\n... ({len(agent_rows) - 20} more rows)" | |
| has_error = False | |
| else: | |
| result_str = "(no result rows returned)" | |
| has_error = False | |
| # End episode if perfect score or max steps | |
| is_done = False | |
| if score >= 1.0: | |
| is_done = True | |
| feedback += " Task completed perfectly!" | |
| elif self._step_count >= MAX_STEPS_PER_TASK: | |
| is_done = True | |
| feedback += f" Maximum steps ({MAX_STEPS_PER_TASK}) reached." | |
| self._done = is_done | |
| return SQLObservation( | |
| task_id=task.task_id, | |
| task_description=task.description, | |
| difficulty=task.difficulty, | |
| schema_description=SCHEMA_DESCRIPTION, | |
| query_result=result_str, | |
| query_error=has_error, | |
| feedback=feedback, | |
| reward=score, | |
| done=is_done, | |
| step_count=self._step_count, | |
| max_steps=MAX_STEPS_PER_TASK, | |
| ) | |
| def state(self) -> State: | |
| return self._state | |