Team Ai
Apppublic

ritvik360/nl2sql-bench

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py226 linesDownload Raw Back to server
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()