Kalletlamadhav/sql-optimization-env
0
1# server/environment.py2 3import random4import sqlite35import time6import uuid7from pathlib import Path8from typing import Optional9 10from .models import *11from .reward import RewardComposer12from .hack_detector import HackDetector13from tasks.task_registry import TaskRegistry14from curriculum.curriculum_engine import CurriculumEngine15from data.seed_database import seed_database16 17DB_PATH = Path('data/fixtures/benchmark_seed42.db')18 19 20class SQLOptEnvironment:21 22 def __init__(self):23 self.task_registry = TaskRegistry()24 self.curriculum = CurriculumEngine()25 self.reward_composer = RewardComposer()26 self.hack_detector = HackDetector()27 self._state: Optional[EnvironmentState] = None28 self._current_task = None29 30 # โ
FIX: define DB path ONLY once31 self._db_path = Path(__file__).resolve().parent.parent / "data" / "fixtures" / "benchmark_seed42.db"32 33 # โ
FIX: seed ONLY if DB doesn't exist (prevents wiping)34 try:35 print("๐ Rebuilding database...", flush=True)36 # Always recreate DB (HF containers are not reliable)37 if self._db_path.exists():38 self._db_path.unlink()39 40 seed_database(str(self._db_path), 1000)41 42 print("โ
Database ready with tables", flush=True)43 44 except Exception as e:45 print(f"โ Seeding failed: {e}", flush=True)46 def reset(self, task_id: str = None) -> SQLOptObservation:47 if task_id:48 task = self.task_registry.get_task_by_id(task_id)49 if task is None:50 raise ValueError(f"Task '{task_id}' not found")51 else:52 pool = self.task_registry.get_all_tasks()53 task = random.choice(pool)54 55 print(f"Selected task: {task.task_id}", flush=True)56 self._current_task = task57 level = task.curriculum_level58 59 self._state = EnvironmentState(60 episode_id=str(uuid.uuid4())[:8],61 current_task_id=task.task_id,62 current_step=0,63 max_steps=task.max_steps,64 curriculum_level=level,65 episode_rewards=[],66 total_episodes=self.curriculum.total_episodes,67 is_running=True68 )69 70 orig_time, orig_rows, exec_plan = self._run_query(task.slow_query)71 72 task.original_time_ms = orig_time73 task.original_plan = exec_plan74 75 return SQLOptObservation(76 task_id=task.task_id,77 step_number=0,78 goal=task.goal,79 schema_ddl=task.schema_ddl,80 current_query=task.slow_query,81 execution_plan=exec_plan,82 execution_time_ms=orig_time,83 row_count=orig_rows,84 db_stats=self._get_db_stats(task.tables),85 curriculum_level=level,86 anti_pattern_hint=task.hint if level <= 2 else None87 )88 89 def step(self, action: SQLOptAction) -> StepResult:90 if not self._state or not self._state.is_running:91 raise ValueError('Call reset() first')92 93 self._state.current_step += 194 task = self._current_task95 96 hack = self.hack_detector.detect(action, task)97 98 try:99 opt_time, opt_rows, opt_plan = self._run_query(100 action.optimized_query, action.index_statements101 )102 query_error = None103 except Exception as e:104 opt_time = max(task.original_time_ms, 1.0) * 2105 opt_rows = -1106 opt_plan = None107 query_error = str(e)108 109 reward_detail = self.reward_composer.compute(110 task=task,111 action=action,112 orig_time=task.original_time_ms,113 opt_time=opt_time,114 opt_rows=opt_rows,115 opt_plan=opt_plan,116 hack=hack,117 query_error=query_error118 )119 120 done = (121 self._state.current_step >= self._state.max_steps or122 reward_detail.total >= 0.95 or123 query_error is not None124 )125 126 if done:127 self._state.is_running = False128 self.curriculum.record_episode(reward_detail.total)129 self._state.episode_rewards.append(reward_detail.total)130 131 next_obs = SQLOptObservation(132 task_id=task.task_id,133 step_number=self._state.current_step,134 goal=task.goal,135 schema_ddl=task.schema_ddl,136 current_query=action.optimized_query,137 execution_plan=opt_plan if opt_plan else task.original_plan,138 execution_time_ms=opt_time,139 row_count=opt_rows,140 db_stats=self._get_db_stats(task.tables),141 curriculum_level=self._state.curriculum_level,142 error_message=query_error143 )144 145 return StepResult(146 observation=next_obs,147 reward=reward_detail.total,148 reward_detail=reward_detail,149 done=done,150 info={'hack': hack}151 )152 153 def state(self) -> EnvironmentState:154 if not self._state:155 raise ValueError('Call reset() first')156 return self._state157 158 # โ
FIX: use in-memory DB + timeout + no mutation159 def _run_query(self, query: str, index_stmts: list = None):160 161 # disk DB162 disk_conn = sqlite3.connect(self._db_path)163 164 # memory DB (isolated)165 mem_conn = sqlite3.connect(":memory:")166 167 # copy DB โ memory (CRITICAL FIX)168 disk_conn.backup(mem_conn)169 disk_conn.close()170 171 mem_conn.execute('PRAGMA foreign_keys = ON')172 173 # โ
timeout protection174 mem_conn.set_progress_handler(lambda: 1/0, 1000000)175 176 if index_stmts:177 for stmt in index_stmts:178 try:179 mem_conn.execute(stmt)180 except:181 pass182 183 plan_rows = mem_conn.execute(f'EXPLAIN QUERY PLAN {query}').fetchall()184 plan_text = ' | '.join(str(r) for r in plan_rows)185 186 using_index = 'USING INDEX' in plan_text.upper()187 is_full_scan = 'SCAN' in plan_text.upper() and not using_index188 189 exec_plan = ExecutionPlan(190 operation='FULL TABLE SCAN' if is_full_scan else 'INDEX SCAN',191 rows_examined=0,192 rows_returned=0,193 cost_estimate=0.0,194 using_index=plan_text if using_index else None,195 missing_index_hint='Consider adding index' if is_full_scan else None,196 explain_raw=plan_text197 )198 199 start = time.perf_counter()200 rows = mem_conn.execute(query).fetchall()201 elapsed = (time.perf_counter() - start) * 1000202 203 mem_conn.close()204 return elapsed, len(rows), exec_plan205 206 def _get_db_stats(self, tables: list) -> dict:207 conn = sqlite3.connect(self._db_path)208 stats = {}209 210 for table in tables:211 try:212 count = conn.execute(f'SELECT COUNT(*) FROM {table}').fetchone()[0]213 stats[table] = {'row_count': count}214 except Exception as e:215 stats[table] = {'error': str(e)}216 217 conn.close()218 return stats