""" 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, ) @property def state(self) -> State: return self._state