Mahathi4554/sql-query-debugging
0
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