dkAmulet/sql-query-optimizer
0
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 )