Team Ai
Apppublic

Mahathi4554/sql-query-debugging

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
Server.py398 linesDownload Raw Back to root
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)