Mahathi4554/sql-query-debugging
0
1"""2FastAPI server for the SQL Query Debugging OpenEnv environment.3v2 Enhancements:4 - Session-based env isolation5 - /leaderboard and /stats endpoints6 - /baseline/compare endpoint7 - Improved error handling8 - SQL safety layer9"""10 11from __future__ import annotations12import os13import time14from collections import defaultdict15from typing import Optional16from fastapi import FastAPI, HTTPException, Request17from fastapi.middleware.cors import CORSMiddleware18from pydantic import BaseModel19from fastapi.responses import HTMLResponse, FileResponse20 21from environment import SQLEnv, Action, Observation, StepResult, EnvState22from tasks import TASK_LIST, TASKS23 24app = FastAPI(25 title="SQL Query Debugging — OpenEnv",26 description=(27 "An OpenEnv environment where AI agents learn to debug and optimize SQL queries. "28 "Simulates real-world data engineering tasks with deterministic graders. "29 "5 tasks ranging from easy syntax fixes to expert-level CTE rewrites."30 ),31 version="2.0.0",32)33 34app.add_middleware(35 CORSMiddleware,36 allow_origins=["*"],37 allow_methods=["*"],38 allow_headers=["*"],39)40 41# Session-based env management42_sessions: dict[str, SQLEnv] = {}43_DEFAULT_SESSION = "default"44 45# In-memory leaderboard46_leaderboard: dict[str, list[dict]] = defaultdict(list)47_total_episodes = 048_total_steps = 049_server_start = time.time()50 51UNSAFE_KEYWORDS = {"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE", "REPLACE", "CREATE", "ATTACH", "DETACH"}52 53 54# ─── Helpers ─────────────────────────────────────────────────────────────────55 56def get_env(session_id: str = _DEFAULT_SESSION) -> SQLEnv:57 if session_id not in _sessions:58 _sessions[session_id] = SQLEnv()59 return _sessions[session_id]60 61 62def check_sql_safety(sql: str) -> Optional[str]:63 """Return error message if SQL contains unsafe operations, else None."""64 upper = sql.upper()65 for kw in UNSAFE_KEYWORDS:66 if kw in upper:67 return f"Unsafe query detected: '{kw}' operation is not allowed."68 return None69 70def generate_hint(sql: str, attempt: int) -> str:71 if not sql:72 return "Write a SQL query to begin."73 74 sql_upper = sql.upper()75 76 if "LIKE '%" in sql_upper:77 return "Avoid leading wildcard LIKE — prevents index usage."78 79 if " OR " in sql_upper:80 return "Try replacing OR chains with IN or better filtering."81 82 if "SUM(" in sql_upper and "GROUP BY" not in sql_upper:83 return "Missing GROUP BY for aggregation."84 85 if "JOIN" in sql_upper and "ON" not in sql_upper:86 return "JOIN detected without ON condition."87 88 if attempt >= 3:89 return "Try restructuring query using a CTE (WITH clause)."90 91 return "Check filtering, joins, and aggregation carefully."92 93# ─── Request/Response models ──────────────────────────────────────────────────94 95class ResetRequest(BaseModel):96 task_id: Optional[str] = None97 session_id: Optional[str] = _DEFAULT_SESSION98 99 100class StepRequest(BaseModel):101 model_config = {"protected_namespaces": ()}102 sql_query: str103 model_name: Optional[str] = None104 session_id: Optional[str] = _DEFAULT_SESSION105 106 107class GraderRequest(BaseModel):108 model_config = {"protected_namespaces": ()}109 task_id: str110 sql_query: str111 model_name: Optional[str] = None112 113 114class BaselineRequest(BaseModel):115 model_config = {"protected_namespaces": ()}116 provider: str = "groq"117 model: Optional[str] = None118 task_id: Optional[str] = None119 user_sql: Optional[str] = None # 🔥 ADD THIS LINE120 121 122class CompareRequest(BaseModel):123 provider: str = "groq"124 task_id: Optional[str] = None125 126 127# ─── OpenEnv Core Endpoints ───────────────────────────────────────────────────128 129@app.get("/")130def root():131 try:132 return FileResponse("Dashboard.html")133 except Exception:134 return {"status": "ok", "env": "sql-query-debugging"}135 136 137@app.get("/health")138def health():139 try:140 uptime_seconds = int(time.time() - _server_start)141 return {142 "status": "ok",143 "env": "sql-query-debugging",144 "version": "2.0.0",145 "uptime_seconds": uptime_seconds,146 "total_episodes": _total_episodes,147 "total_steps": _total_steps,148 "task_count": len(TASKS),149 }150 except Exception as e:151 raise HTTPException(status_code=500, detail=str(e))152 153 154@app.post("/reset", response_model=Observation)155def reset(req: ResetRequest = ResetRequest()):156 """Start a new episode. Optionally specify task_id and session_id."""157 global _total_episodes158 try:159 env = get_env(req.session_id or _DEFAULT_SESSION)160 obs = env.reset(task_id=req.task_id)161 _total_episodes += 1162 return obs163 except ValueError as e:164 raise HTTPException(status_code=400, detail=str(e))165 except Exception as e:166 raise HTTPException(status_code=500, detail=str(e))167 168 169@app.post("/step", response_model=StepResult)170def step(req: StepRequest):171 """Submit a SQL query as an action."""172 global _total_steps173 try:174 safety_err = check_sql_safety(req.sql_query)175 if safety_err:176 raise HTTPException(status_code=400, detail=safety_err)177 178 env = get_env(req.session_id or _DEFAULT_SESSION)179 action = Action(sql_query=req.sql_query)180 result = env.step(action)181 182 hint = generate_hint(183 req.sql_query,184 result.observation.attempt185 )186 187 result.observation.hint = hint188 _total_steps += 1189 190 if req.model_name and result.reward.breakdown.get("correctness", 0) >= 0.7:191 _leaderboard[result.observation.task_id].append({192 "score": result.reward.value,193 "model": req.model_name,194 "timestamp": time.time(),195 "attempts": result.observation.attempt,196 })197 198 diff_data = result.info.get("diff_data", {})199 200 return {201 "observation": result.observation,202 "reward": result.reward,203 "done": result.done,204 "expected_output": diff_data.get("expected"),205 "actual_output": diff_data.get("actual")206 }207 except HTTPException:208 raise209 except RuntimeError as e:210 raise HTTPException(status_code=400, detail=str(e))211 except Exception as e:212 raise HTTPException(status_code=500, detail=str(e))213 214 215@app.get("/state", response_model=EnvState)216def state(session_id: str = _DEFAULT_SESSION):217 """Return current episode state without side effects."""218 try:219 env = get_env(session_id)220 return env.state()221 except RuntimeError as e:222 raise HTTPException(status_code=400, detail=str(e))223 except Exception as e:224 raise HTTPException(status_code=500, detail=str(e))225 226 227# ─── Required Judging Endpoints ───────────────────────────────────────────────228 229@app.get("/tasks")230def list_tasks():231 """Returns all tasks with descriptions, difficulty, and action schema."""232 try:233 return {"tasks": TASK_LIST}234 except Exception as e:235 raise HTTPException(status_code=500, detail=str(e))236 237 238@app.post("/grader")239def grader(req: GraderRequest):240 """Grade a SQL query against a specific task without running a full episode."""241 try:242 if req.task_id not in TASKS:243 raise HTTPException(status_code=404, detail=f"Unknown task_id: {req.task_id}")244 245 safety_err = check_sql_safety(req.sql_query)246 if safety_err:247 raise HTTPException(status_code=400, detail=safety_err)248 249 env = SQLEnv()250 score = env.grade(req.task_id, req.sql_query)251 252 if req.model_name and score >= 0.9:253 _leaderboard[req.task_id].append({254 "score": score,255 "model": req.model_name,256 "timestamp": time.time(),257 "attempts": 0,258 })259 260 return {261 "task_id": req.task_id,262 "score": score,263 "message": f"Score for task '{req.task_id}': {score:.4f}",264 }265 except HTTPException:266 raise267 except Exception as e:268 raise HTTPException(status_code=500, detail=str(e))269 270 271@app.get("/baseline")272def baseline_get():273 """Run baseline with Groq llama-3.3-70b (requires GROQ_API_KEY env var)."""274 try:275 return _run_baseline_internal(provider="groq", model=None, task_id=None)276 except HTTPException:277 raise278 except Exception as e:279 raise HTTPException(status_code=500, detail=str(e))280 281 282@app.post("/baseline")283def baseline_post(req: BaselineRequest = BaselineRequest()):284 """Run baseline with specified provider/model. Optionally target a single task and custom SQL."""285 try:286 from Baseline import run_baseline287 288 task_ids = [req.task_id] if req.task_id else None289 290 results = run_baseline(291 provider=req.provider,292 model=req.model,293 task_ids=task_ids,294 user_sql=req.user_sql # ✅ FIX: pass user input295 )296 297 return results298 299 except HTTPException:300 raise301 except Exception as e:302 raise HTTPException(status_code=500, detail=str(e))303 304 305@app.post("/baseline/compare")306def baseline_compare(req: CompareRequest = CompareRequest()):307 """Run all models for a provider and return comparison results."""308 try:309 from baseline import run_comparison_for_task310 results = run_comparison_for_task(provider=req.provider, task_id=req.task_id)311 return {"results": results}312 except ImportError:313 raise HTTPException(status_code=500, detail="baseline.py not found.")314 except ValueError as e:315 raise HTTPException(status_code=400, detail=str(e))316 except Exception as e:317 raise HTTPException(status_code=500, detail=f"Comparison failed: {e}")318 319 320def _run_baseline_internal(provider: str, model, task_id: Optional[str] = None, task_ids=None):321 try:322 from baseline import run_baseline323 results = run_baseline(provider=provider, model=model, task_ids=task_ids)324 325 # Update leaderboard326 for task_result in results.get("tasks", []):327 if task_result.get("solved"):328 tid = task_result["task_id"]329 _leaderboard[tid].append({330 "score": task_result["final_score"],331 "model": results.get("model", "unknown"),332 "timestamp": time.time(),333 "attempts": task_result.get("attempts", 0),334 })335 336 return results337 except ImportError:338 raise HTTPException(status_code=500, detail="baseline.py not found.")339 except ValueError as e:340 raise HTTPException(status_code=400, detail=str(e))341 except Exception as e:342 raise HTTPException(status_code=500, detail=f"Baseline run failed: {e}")343 344 345# ─── Analytics Endpoints ──────────────────────────────────────────────────────346 347@app.get("/leaderboard")348def leaderboard():349 """Returns best scores per task, sorted by score descending."""350 try:351 result = {}352 for task_id in TASKS:353 entries = _leaderboard.get(task_id, [])354 sorted_entries = sorted(entries, key=lambda x: x["score"], reverse=True)355 result[task_id] = {356 "task_id": task_id,357 "difficulty": TASKS[task_id].difficulty,358 "top_scores": sorted_entries[:10],359 "total_attempts": len(entries),360 }361 return {"leaderboard": result}362 except Exception as e:363 raise HTTPException(status_code=500, detail=str(e))364 365 366@app.get("/stats")367def stats():368 """Returns aggregate statistics about environment usage."""369 try:370 uptime = int(time.time() - _server_start)371 task_stats = {}372 for task_id, task in TASKS.items():373 entries = _leaderboard.get(task_id, [])374 scores = [e["score"] for e in entries]375 task_stats[task_id] = {376 "difficulty": task.difficulty,377 "total_solves": len(scores),378 "avg_score": round(sum(scores)/len(scores), 4) if scores else None,379 "best_score": max(scores) if scores else None,380 }381 return {382 "environment": "sql-query-debugging",383 "version": "2.0.0",384 "uptime_seconds": uptime,385 "total_episodes": _total_episodes,386 "total_steps": _total_steps,387 "tasks": task_stats,388 }389 except Exception as e:390 raise HTTPException(status_code=500, detail=str(e))391 392 393# ─── Dev entrypoint ──────────────────────────────────────────────────────────394 395if __name__ == "__main__":396 import uvicorn397 port = int(os.environ.get("PORT", 7860))398 uvicorn.run("server:app", host="0.0.0.0", port=port, reload=False)