ritvik360/nl2sql-bench
0
1"""2nl2sql-bench/server/environment.py3====================================4NL2SQL-Bench core environment — implements the OpenEnv Environment interface.5 6Episode flow7------------81. reset(task_name?) → picks a task + question, returns initial observation92. step(action) → executes the SQL, grades it, returns observation + reward103. state() → returns episode metadata114. Episode ends when: exact_match OR step count reaches max_steps12 13The environment manages its own SQLite connection (in-memory, seeded14deterministically). One connection per Environment instance; the FastAPI15server creates one Environment per WebSocket session.16"""17 18from __future__ import annotations19 20import os21import sqlite322import uuid23from pathlib import Path24from typing import Optional25 26from openenv.core.env_server import Environment27 28# Import after openenv so path is correct regardless of working directory29_HERE = Path(__file__).parent30 31# Lazy import of task registry (avoids circular imports)32from tasks import get_task, all_task_names, BaseTask33from tasks.base import TaskExample34from grader import (35 GradeResult,36 compute_ground_truth,37 execute_query,38 grade,39 has_order_by,40)41 42# We import our models from one level up (models.py at project root)43import sys44sys.path.insert(0, str(_HERE.parent))45from models import NL2SQLAction, NL2SQLObservation, NL2SQLState46 47# ── Constants ──────────────────────────────────────────────────────────────48DEFAULT_TASK = os.getenv("NL2SQL_DEFAULT_TASK", "simple-filter")49MAX_STEPS = int(os.getenv("NL2SQL_MAX_STEPS", "5"))50RESULT_LIMIT = 10 # Max rows shown to agent per step51 52 53class NL2SQLEnvironment(Environment):54 """55 OpenEnv-compliant environment for NL-to-SQL query generation.56 57 One instance per WebSocket session (created by create_fastapi_app).58 """59 60 def __init__(self) -> None:61 self._conn: Optional[sqlite3.Connection] = None62 self._task: Optional[BaseTask] = None63 self._example: Optional[TaskExample] = None64 self._ground_truth: list = []65 self._order_sensitive: bool = False66 self._state = NL2SQLState(67 episode_id=None,68 step_count=0,69 task_name="",70 task_difficulty="",71 question="",72 best_reward=0.0,73 cumulative_reward=0.0,74 solved=False75 )76 self._last_obs = NL2SQLObservation(77 question="",78 schema_context="",79 task_name="",80 last_query="",81 last_result=[],82 last_error=None,83 result_columns=[],84 step=0,85 max_steps=5,86 done=False,87 reward=None,88 score=0.089 )90 self._episode_rewards: list = []91 self._setup_db()92 93 # ── DB lifecycle ───────────────────────────────────────────────────────94 95 def _setup_db(self) -> None:96 """Create in-memory SQLite DB and seed it."""97 schema_path = _HERE / "db" / "schema.sql"98 from db.seed import seed_database # local import after sys.path setup99 conn = sqlite3.connect(":memory:", check_same_thread=False)100 conn.row_factory = sqlite3.Row101 conn.execute("PRAGMA foreign_keys = ON")102 conn.executescript(schema_path.read_text())103 seed_database(conn)104 self._conn = conn105 106 # ── OpenEnv interface ──────────────────────────────────────────────────107 108 def reset(self, task_name: Optional[str] = None) -> NL2SQLObservation:109 """110 Start a new episode.111 112 task_name: one of 'simple-filter', 'join-aggregation', 'analytics-window'.113 Defaults to NL2SQL_DEFAULT_TASK env-var or 'simple-filter'.114 """115 task_name = task_name or DEFAULT_TASK116 if task_name not in all_task_names():117 task_name = DEFAULT_TASK118 119 self._task = get_task(task_name)120 self._example = self._task.next_example()121 self._order_sensitive = has_order_by(self._example.sql)122 123 # Pre-compute ground truth once per episode124 self._ground_truth = compute_ground_truth(self._conn, self._example.sql)125 126 self._episode_rewards = []127 self._state = NL2SQLState(128 episode_id=str(uuid.uuid4()),129 step_count=0,130 task_name=self._task.name,131 task_difficulty=self._task.difficulty,132 question=self._example.question,133 best_reward=0.0,134 cumulative_reward=0.0,135 solved=False,136 )137 138 obs = NL2SQLObservation(139 question=self._example.question,140 schema_context=self._task.schema_context(),141 task_name=self._task.name,142 last_query="",143 last_result=[],144 last_error=None,145 result_columns=[],146 step=0,147 max_steps=MAX_STEPS,148 done=False,149 reward=None,150 score=0.0,151 )152 self._last_obs = obs153 return obs154 155 def step(self, action: NL2SQLAction) -> NL2SQLObservation:156 """Execute the agent's SQL and return graded observation."""157 if self._task is None or self._example is None:158 # Called before reset — auto-reset159 self.reset()160 161 self._state.step_count += 1162 current_step = self._state.step_count163 done = False164 165 # Execute the query166 rows, error = execute_query(self._conn, action.query)167 168 # Grade it169 result: GradeResult = grade(170 actual_rows=rows,171 ground_truth_rows=self._ground_truth,172 error=error,173 step=current_step,174 order_sensitive=self._order_sensitive,175 )176 177 reward = result.reward178 self._episode_rewards.append(reward)179 self._state.cumulative_reward += reward180 self._state.best_reward = max(self._state.best_reward, reward)181 182 if result.exact_match:183 self._state.solved = True184 done = True185 elif current_step >= MAX_STEPS:186 done = True187 188 # Prepare result rows for observation (truncated for agent readability)189 display_rows = (rows or [])[:RESULT_LIMIT]190 result_columns = list(display_rows[0].keys()) if display_rows else []191 # Convert sqlite3.Row objects if needed192 display_rows = [dict(r) for r in display_rows]193 194 # Normalised cumulative score195 n = len(self._episode_rewards)196 score = self._state.cumulative_reward / max(n, 1) if n else 0.0197 score = round(min(max(score, 0.0), 1.0), 4)198 199 obs = NL2SQLObservation(200 question=self._example.question,201 schema_context=self._task.schema_context(),202 task_name=self._task.name,203 last_query=action.query,204 last_result=display_rows,205 last_error=error,206 result_columns=result_columns,207 step=current_step,208 max_steps=MAX_STEPS,209 done=done,210 reward=reward,211 score=score,212 )213 self._last_obs = obs214 215 # openenv-core expects ONLY the observation returned from step().216 # The framework reads obs.reward and obs.done itself — do NOT return a tuple.217 return obs218 219 @property220 def state(self) -> NL2SQLState:221 return self._state222 223 # ── Helpers ────────────────────────────────────────────────────────────224 225 def available_tasks(self) -> list:226 return all_task_names()