Team Ai
Apppublic

Kalletlamadhav/sql-optimization-env

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