CoolFace
Apppublic

dkAmulet/sql-query-optimizer

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
env.py143 linesDownload Raw Back to root
1"""2SQLQueryOptimizerEnv — OpenEnv-compatible Python class.3All reward scores are strictly clamped to (0.001, 0.999).4"""5from __future__ import annotations6 7import sqlite38from typing import List, Optional9 10from db import create_database11from models import (12    EnvironmentState,13    ExecutionMetrics,14    RewardBreakdown,15    SQLAction,16    SQLObservation,17    SQLReward,18    StepResult,19)20from tasks import (21    SCHEMA_DDL,22    TASK_GRADERS,23    TASK_ORDER,24    TASKS,25    get_query_metrics,26)27 28 29def _clamp(v: float) -> float:30    """Force any score strictly into (0.001, 0.999) — never 0.0 or 1.0."""31    return max(0.001, min(0.999, float(v)))32 33 34class SQLQueryOptimizerEnv:35 36    def __init__(self) -> None:37        self._conn: Optional[sqlite3.Connection] = None38        self._task_id: Optional[str] = None39        self._step: int = 040        self._best_reward: float = 0.00141        self._best_query: str = ""42        self._done: bool = False43        self._last_feedback: str = ""44        self._last_reward: float = 0.00145        self._conn = create_database()46 47    def reset(self, task_id: Optional[str] = None) -> SQLObservation:48        if task_id is None:49            task_id = TASK_ORDER[0]50        if task_id not in TASKS:51            valid = list(TASKS.keys())52            raise ValueError(f"Unknown task_id '{task_id}'. Valid options: {valid}")53        self._task_id = task_id54        self._step = 055        self._best_reward = 0.00156        self._done = False57        self._last_feedback = ""58        self._last_reward = 0.00159        task = TASKS[task_id]60        slow_query = task["slow_query"]61        self._best_query = slow_query62        return self._make_obs(slow_query)63 64    def step(self, action: SQLAction) -> StepResult:65        if self._task_id is None:66            raise RuntimeError("No active episode — call reset() first.")67        if self._done:68            raise RuntimeError("Episode is over — call reset() to start a new one.")69        task = TASKS[self._task_id]70        self._step += 171        grader = TASK_GRADERS[self._task_id]72        raw_score, raw_bd, feedback = grader(action.optimized_query, self._conn)73        score = _clamp(raw_score)74        bd_dict = {k: _clamp(v) for k, v in raw_bd.items()}75        if score > self._best_reward:76            self._best_reward = score77            self._best_query = action.optimized_query78        self._last_feedback = feedback79        self._last_reward = score80        self._done = self._step >= task["max_steps"] or score >= 0.9581        reward = SQLReward(82            value=score,83            breakdown=RewardBreakdown(**bd_dict),84            feedback=feedback,85        )86        obs = self._make_obs(action.optimized_query)87        return StepResult(88            observation=obs,89            reward=reward,90            done=self._done,91            info={92                "best_reward": self._best_reward,93                "step": self._step,94                "task_id": self._task_id,95                "episode_done": self._done,96            },97        )98 99    def state(self) -> EnvironmentState:100        max_steps = TASKS[self._task_id]["max_steps"] if self._task_id else 0101        return EnvironmentState(102            task_id=self._task_id or "",103            step_number=self._step,104            max_steps=max_steps,105            best_reward=self._best_reward,106            done=self._done,107            current_query=self._best_query,108        )109 110    def list_tasks(self) -> List[dict]:111        return [112            {113                "task_id": tid,114                "name": t["name"],115                "difficulty": t["difficulty"],116                "max_steps": t["max_steps"],117                "description": t["description"],118            }119            for tid, t in TASKS.items()120        ]121 122    def close(self) -> None:123        if self._conn:124            self._conn.close()125            self._conn = None126 127    def _make_obs(self, current_query: str) -> SQLObservation:128        task = TASKS[self._task_id]129        slow_metrics = get_query_metrics(task["slow_query"], self._conn)130        return SQLObservation(131            task_id=self._task_id,132            task_name=task["name"],133            difficulty=task["difficulty"],134            description=task["description"],135            schema_ddl=SCHEMA_DDL,136            slow_query=task["slow_query"],137            current_query=current_query,138            step_number=self._step,139            max_steps=task["max_steps"],140            slow_metrics=slow_metrics,141            last_feedback=self._last_feedback,142            last_reward=self._last_reward,143        )