Team Ai
Apppublic

Mahathi4554/sql-query-debugging

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
Environment.py396 linesDownload Raw Back to root
1"""2SQL Query Debugging Environment — OpenEnv compliant3v2 Enhancements:4  - SQL safety layer (blocks DROP, DELETE, UPDATE, INSERT, ALTER, etc.)5  - Session isolation support6  - Progressive hint system7  - Richer reward breakdown8"""9 10from __future__ import annotations11import sqlite312import time13import uuid14from typing import Any, Optional15from pydantic import BaseModel, Field16from tasks import TASKS, Task17 18UNSAFE_KEYWORDS = {"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE", "REPLACE", "ATTACH", "DETACH"}19 20 21# ─── Pydantic Models (OpenEnv spec) ──────────────────────────────────────────22 23class Observation(BaseModel):24    task_id: str = Field(..., description="Unique task identifier")25    task_description: str = Field(..., description="Human-readable task goal")26    schema_sql: str = Field(..., description="CREATE TABLE statements for the database schema")27    broken_query: str = Field(..., description="The SQL query the agent must fix or improve")28    error_message: Optional[str] = Field(None, description="SQL error from last execution, if any")29    last_result_preview: Optional[list[dict]] = Field(None, description="First 5 rows of last query result")30    expected_row_count: Optional[int] = Field(None, description="Expected number of rows in correct result")31    attempt: int = Field(0, description="Number of attempts made so far in this episode")32    max_attempts: int = Field(5, description="Maximum allowed attempts per episode")33    hint: Optional[str] = Field(None, description="Hint revealed after 2 failed attempts")34    hint_advanced: Optional[str] = Field(None, description="Stronger hint revealed after 4 failed attempts")35 36 37class Action(BaseModel):38    sql_query: str = Field(..., description="The SQL query to execute as the agent's attempt")39 40 41class Reward(BaseModel):42    value: float = Field(..., ge=0.0, le=1.0, description="Reward score between 0.0 and 1.0")43    breakdown: dict[str, float] = Field(..., description="Component scores: syntax, execution, correctness, optimization")44    message: str = Field(..., description="Human-readable explanation of the reward")45 46 47class StepResult(BaseModel):48    observation: Observation49    reward: Reward50    done: bool51    info: dict[str, Any]52 53 54class EnvState(BaseModel):55    episode_id: str56    task_id: str57    attempt: int58    max_attempts: int59    done: bool60    best_score: float61    current_broken_query: str62    db_path: str63 64 65# ─── SQL Safety ───────────────────────────────────────────────────────────────66 67def check_sql_safety(sql: str) -> Optional[str]:68    """Return error message if SQL contains unsafe operations, else None."""69    upper = sql.upper()70    for kw in UNSAFE_KEYWORDS:71        # Check for keyword as word boundary to avoid false positives72        import re73        if re.search(r'\b' + kw + r'\b', upper):74            return f"Unsafe query detected: '{kw}' operation is not allowed."75    return None76 77 78# ─── Reward Engine ────────────────────────────────────────────────────────────79 80def compute_reward(81    task: Task,82    submitted_sql: str,83    exec_result: dict,84    attempt: int,85) -> Reward:86    """87    Partial reward function — provides signal at every step:88      0.1  — query is valid SQL (no syntax error)89      0.2  — query executes without runtime error90      0.3  — query returns the correct number of rows (partial correctness)91      0.7  — query returns exactly the correct result set92      +0.3 — optimization bonus (if task.check_optimization and efficient plan)93      -0.02/attempt — attempt penalty (max -0.10)94    """95    breakdown: dict[str, float] = {96        "syntax": 0.0,97        "execution": 0.0,98        "correctness": 0.0,99        "optimization": 0.0,100    }101 102    if exec_result.get("syntax_valid"):103        breakdown["syntax"] = 0.1104 105    if exec_result.get("executed"):106        breakdown["execution"] = 0.2107 108    rows = exec_result.get("rows", [])109    expected = task.expected_rows110 111    if exec_result.get("executed") and expected is not None:112        total_cells = 0113        correct_cells = 0114 115        for i in range(min(len(rows), len(expected))):116            for key in expected[i]:117                total_cells += 1118                if key in rows[i] and rows[i][key] == expected[i][key]:119                    correct_cells += 1120 121        if total_cells > 0:122            cell_score = correct_cells / total_cells123 124            if cell_score == 1.0 and len(rows) == len(expected):125                breakdown["correctness"] = 0.7126            else:127                breakdown["correctness"] = round(0.6 * cell_score, 3)128 129    # small bonus if row count matches130    if len(rows) == len(expected):131        breakdown["correctness"] += 0.1132 133    if task.check_optimization and exec_result.get("executed") and breakdown["correctness"] >= 0.7:134        cost = exec_result.get("query_cost", 999)135        if cost <= task.optimization_target_cost:136            breakdown["optimization"] = 0.3137        elif cost <= task.optimization_target_cost * 1.5:138            breakdown["optimization"] = 0.15139        elif cost <= task.optimization_target_cost * 2:140            breakdown["optimization"] = 0.05141 142    raw = sum(breakdown.values())143    attempt_penalty = min(attempt * 0.02, 0.10)144    value = max(0.0, min(1.0, raw - attempt_penalty))145 146    messages = []147    if breakdown["syntax"] == 0:148        messages.append("Query has a syntax error.")149    elif breakdown["execution"] == 0:150        messages.append("Query parsed but failed to execute.")151    elif breakdown["correctness"] < 0.3:152        got = len(rows)153        exp = len(expected) if expected else "?"154        messages.append(f"Query runs but returns wrong rows. Got {got}, expected {exp}.")155    elif breakdown["correctness"] < 0.7:156        messages.append(f"Row count matches ({len(rows)}) but values are incorrect.")157    else:158        messages.append("Correct result!")159 160    if breakdown["optimization"] == 0.3:161        messages.append("Optimization bonus: efficient query plan.")162    elif breakdown["optimization"] > 0:163        messages.append(f"Partial optimization bonus: {breakdown['optimization']:.2f} (plan could be more efficient).")164 165    if attempt > 0:166        messages.append(f"Attempt penalty: -{attempt_penalty:.2f}.")167 168    return Reward(value=round(value, 4), breakdown=breakdown, message=" ".join(messages))169 170 171def _rows_match(got: list[dict], expected: list[dict]) -> bool:172    if len(got) != len(expected):173        return False174    try:175        got_sorted = sorted([tuple(sorted(r.items())) for r in got])176        exp_sorted = sorted([tuple(sorted(r.items())) for r in expected])177        return got_sorted == exp_sorted178    except Exception:179        return False180 181 182def _count_matching_rows(got: list[dict], expected: list[dict]) -> int:183    try:184        exp_set = set(tuple(sorted(r.items())) for r in expected)185        return sum(1 for r in got if tuple(sorted(r.items())) in exp_set)186    except Exception:187        return 0188 189 190# ─── SQLite Executor ──────────────────────────────────────────────────────────191 192def execute_sql(db_path: str, sql: str) -> dict:193    """Safely execute SQL, return structured result."""194    result: dict[str, Any] = {195        "syntax_valid": False,196        "executed": False,197        "rows": [],198        "error": None,199        "query_cost": 999,200        "execution_time_ms": 0,201    }202 203    # Safety check204    safety_err = check_sql_safety(sql)205    if safety_err:206        result["error"] = safety_err207        return result208 209    try:210        compiled = sqlite3.complete_statement(sql)211        if not compiled:212            result["error"] = "Incomplete SQL statement."213            return result214        result["syntax_valid"] = True215    except Exception as e:216        result["error"] = f"Syntax error: {e}"217        return result218 219    try:220        conn = sqlite3.connect(db_path)221        conn.row_factory = sqlite3.Row222        cur = conn.cursor()223        t0 = time.perf_counter()224        cur.execute(sql)225        result["execution_time_ms"] = round((time.perf_counter() - t0) * 1000, 2)226        rows = cur.fetchall()227        result["rows"] = [dict(r) for r in rows]228        result["executed"] = True229 230        try:231            cur.execute(f"EXPLAIN QUERY PLAN {sql}")232            plan = cur.fetchall()233            scans = sum(1 for row in plan if "SCAN" in str(row).upper())234            result["query_cost"] = scans if scans > 0 else 1235        except Exception:236            result["query_cost"] = 1237 238        conn.close()239    except sqlite3.Error as e:240        result["error"] = str(e)241 242    return result243 244 245# ─── Main Environment Class ───────────────────────────────────────────────────246 247class SQLEnv:248    """OpenEnv-compliant SQL Query Debugging Environment."""249 250    def __init__(self):251        self._episode_id: Optional[str] = None252        self._task: Optional[Task] = None253        self._db_path: Optional[str] = None254        self._attempt: int = 0255        self._done: bool = False256        self._best_score: float = 0.0257        self._last_obs: Optional[Observation] = None258 259    def reset(self, task_id: Optional[str] = None) -> Observation:260        """Start a new episode. If task_id is None, picks randomly."""261        self._episode_id = str(uuid.uuid4())262        self._attempt = 0263        self._done = False264        self._best_score = 0.0265 266        if task_id is None:267            import random268            task_id = random.choice(list(TASKS.keys()))269 270        if task_id not in TASKS:271            raise ValueError(f"Unknown task_id '{task_id}'. Available: {list(TASKS.keys())}")272 273        self._task = TASKS[task_id]274        self._db_path = self._task.setup_db()275 276        obs = Observation(277            task_id=self._task.task_id,278            task_description=self._task.description,279            schema_sql=self._task.schema_sql,280            broken_query=self._task.broken_query,281            error_message=None,282            last_result_preview=None,283            expected_row_count=len(self._task.expected_rows) if self._task.expected_rows else None,284            attempt=0,285            max_attempts=self._task.max_attempts,286            hint=None,287            hint_advanced=None,288        )289        self._last_obs = obs290        return obs291 292    def step(self, action: Action) -> StepResult:293        """Execute one agent action (SQL submission)."""294        if self._done:295            raise RuntimeError("Episode is done. Call reset() to start a new one.")296        if self._task is None:297            raise RuntimeError("No active episode. Call reset() first.")298 299        # Safety check300        safety_err = check_sql_safety(action.sql_query)301        if safety_err:302            raise RuntimeError(safety_err)303 304        self._attempt += 1305        exec_result = execute_sql(self._db_path, action.sql_query)306        reward = compute_reward(self._task, action.sql_query, exec_result, self._attempt)307 308        if reward.value > self._best_score:309            self._best_score = reward.value310 311        done = (312            reward.value >= 0.95313            or self._attempt >= self._task.max_attempts314        )315        self._done = done316 317        hint = None318        hint_advanced = None319        if self._attempt >= 2 and reward.breakdown.get("correctness", 0) < 0.7:320            hint = self._task.hint321        if self._attempt >= 4 and reward.breakdown.get("correctness", 0) < 0.7:322            hint_advanced = self._task.hint_advanced323 324        obs = Observation(325            task_id=self._task.task_id,326            task_description=self._task.description,327            schema_sql=self._task.schema_sql,328            broken_query=self._task.broken_query,329            error_message=exec_result.get("error"),330            last_result_preview=exec_result.get("rows", [])[:5],331            expected_row_count=len(self._task.expected_rows) if self._task.expected_rows else None,332            attempt=self._attempt,333            max_attempts=self._task.max_attempts,334            hint=hint,335            hint_advanced=hint_advanced,336        )337        self._last_obs = obs338        339        expected_rows = self._task.expected_rows or []340        actual_rows = exec_result.get("rows", [])341 342        diff_data = {343            "expected": {344                "columns": list(expected_rows[0].keys()) if expected_rows else [],345                "rows": [list(r.values()) for r in expected_rows]346            },347            "actual": {348                "columns": list(actual_rows[0].keys()) if actual_rows else [],349                "rows": [list(r.values()) for r in actual_rows]350            }351        }352 353        return StepResult(354            observation=obs,355            reward=reward,356            done=done,357            info={358                "episode_id": self._episode_id,359                "attempt": self._attempt,360                "best_score": self._best_score,361                "execution_time_ms": exec_result.get("execution_time_ms", 0),362                "query_cost": exec_result.get("query_cost"),363                "exec_error": exec_result.get("error"),364                "diff_data": diff_data   # 🔥 ADD THIS365            },366        )367 368    def state(self) -> EnvState:369        """Return current episode state (no side effects)."""370        if self._task is None:371            raise RuntimeError("No active episode. Call reset() first.")372        return EnvState(373            episode_id=self._episode_id or "",374            task_id=self._task.task_id,375            attempt=self._attempt,376            max_attempts=self._task.max_attempts,377            done=self._done,378            best_score=self._best_score,379            current_broken_query=self._task.broken_query,380            db_path=self._db_path or "",381        )382 383    def grade(self, task_id: str, submitted_sql: str) -> float:384        """Standalone grader — run without a live episode. Returns 0.0–1.0."""385        if task_id not in TASKS:386            raise ValueError(f"Unknown task_id '{task_id}'")387 388        safety_err = check_sql_safety(submitted_sql)389        if safety_err:390            return 0.0391 392        task = TASKS[task_id]393        db_path = task.setup_db()394        exec_result = execute_sql(db_path, submitted_sql)395        reward = compute_reward(task, submitted_sql, exec_result, attempt=0)396        return reward.value