Team Ai
Apppublic

Mahathi4554/sql-query-debugging

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
Baseline.py392 linesDownload Raw Back to root
1"""2Baseline inference script for the SQL Query Debugging OpenEnv environment.3Supports multiple model providers:4  - Groq (default): GROQ_API_KEY, fast inference5  - OpenAI: OPENAI_API_KEY6  - Any OpenAI-compatible endpoint via BASE_URL env var7 8v2 Enhancements:9  - Step-by-step agent trace tracking10  - Reflection prompt with error context11  - run_comparison_for_task for single-task multi-model comparison12  - Improved prompt with full conversation history13 14Usage:15    GROQ_API_KEY=gsk_... python baseline.py16    OPENAI_API_KEY=sk-... python baseline.py --provider openai17    GROQ_API_KEY=gsk_... python baseline.py --model mixtral-8x7b-3276818    GROQ_API_KEY=gsk_... python baseline.py --compare19"""20 21from __future__ import annotations22import os23import json24import time25import argparse26from typing import Optional27from environment import SQLEnv, Action28from tasks import TASKS29 30 31SYSTEM_PROMPT = """You are an expert SQL debugger and optimizer. You will be given:321. A database schema (CREATE TABLE statements)332. A broken or suboptimal SQL query343. Feedback from previous attempts (errors, partial results, hints)35 36Your job is to fix the SQL query so it returns the correct result.37Respond with ONLY valid SQL — no explanation, no markdown, no backticks.38Just the raw SQL query."""39 40USER_PROMPT_TEMPLATE = """Database schema:41{schema_sql}42 43Task: {task_description}44 45Current query (broken/suboptimal):46{broken_query}47 48{feedback}49Write the corrected SQL query:"""50 51REFLECTION_PROMPT = """Previous attempt #{attempt} failed.52 53Previous query:54{prev_query}55 56Result: score={prev_score:.4f}57Feedback: {prev_message}58{error_info}59{hint_info}60 61Reflect on why the previous query was wrong and fix it.62Write the corrected SQL query (raw SQL only, no explanation):"""63 64 65PROVIDER_CONFIGS = {66    "groq": {67        "base_url": "https://api.groq.com/openai/v1",68        "api_key_env": "GROQ_API_KEY",69        "default_model": "llama-3.3-70b-versatile",70        "models": [71            "llama-3.3-70b-versatile",72            "mixtral-8x7b-32768",73            "gemma2-9b-it",74        ],75    },76    "openai": {77        "base_url": "https://api.openai.com/v1",78        "api_key_env": "OPENAI_API_KEY",79        "default_model": "gpt-4o-mini",80        "models": [81            "gpt-4o-mini",82            "gpt-4o",83            "gpt-3.5-turbo",84        ],85    },86}87 88 89def build_feedback(obs, prev_score: Optional[float] = None, prev_message: Optional[str] = None,90                   prev_query: Optional[str] = None, attempt: int = 0) -> str:91    """Build rich feedback string for the next attempt."""92    parts = []93 94    if attempt > 0 and prev_query and prev_score is not None:95        parts.append(REFLECTION_PROMPT.format(96            attempt=attempt,97            prev_query=prev_query,98            prev_score=prev_score,99            prev_message=prev_message or "",100            error_info=f"Error: {obs.error_message}" if obs.error_message else "",101            hint_info=f"Hint: {obs.hint}" if obs.hint else "",102        ))103        if obs.hint_advanced:104            parts.append(f"Advanced hint: {obs.hint_advanced}")105        if obs.last_result_preview is not None:106            parts.append(f"Last result preview (first 5 rows): {json.dumps(obs.last_result_preview)}")107            if obs.expected_row_count is not None:108                parts.append(f"Expected row count: {obs.expected_row_count}")109    else:110        if obs.error_message:111            parts.append(f"Error from last attempt: {obs.error_message}")112        if obs.last_result_preview is not None:113            parts.append(f"Last result preview (first 5 rows): {json.dumps(obs.last_result_preview)}")114            if obs.expected_row_count is not None:115                parts.append(f"Expected row count: {obs.expected_row_count}")116        if obs.hint:117            parts.append(f"Hint: {obs.hint}")118        if obs.hint_advanced:119            parts.append(f"Advanced hint: {obs.hint_advanced}")120 121    return "\n".join(parts) if parts else ""122 123 124def make_client(provider: str = "groq", api_key: Optional[str] = None, base_url: Optional[str] = None):125    """Create an OpenAI-compatible client for the given provider."""126    try:127        from openai import OpenAI128    except ImportError:129        raise ImportError("openai package not installed. Run: pip install openai")130 131    config = PROVIDER_CONFIGS.get(provider, PROVIDER_CONFIGS["groq"])132    resolved_key = api_key or os.environ.get(config["api_key_env"])133    if not resolved_key:134        raise ValueError(135            f"API key not found. Set {config['api_key_env']} environment variable "136            f"or pass --api-key."137        )138 139    return OpenAI(140        api_key=resolved_key,141        base_url=base_url or config["base_url"],142    )143 144 145def call_llm(client, model: str, messages: list) -> str:146    """Call the LLM and return the cleaned SQL response."""147    try:148        response = client.chat.completions.create(149            model=model,150            messages=messages,151            temperature=0.0,152            max_tokens=512,153        )154        sql = response.choices[0].message.content.strip()155        # Strip accidental markdown fences156        if sql.startswith("```"):157            lines = sql.split("\n")158            sql = "\n".join(l for l in lines if not l.startswith("```")).strip()159        return sql160    except Exception as e:161        print(f"  API error: {e}")162        return "SELECT 1;"163 164 165def run_task_with_agent(env: SQLEnv, task_id: str, client, model: str, user_sql=None) -> dict:166    """Run one task episode with the agent. Returns result dict with step trace."""167    obs = env.reset(task_id=task_id)168 169    if user_sql:170        obs.broken_query = user_sql # 🔥 THIS IS THE KEY LINE171    task = TASKS[task_id]172    steps = []173    final_result = None174    prev_score = None175    prev_message = None176    prev_query = None177 178    for attempt_idx in range(task.max_attempts):179        feedback = build_feedback(180            obs,181            prev_score=prev_score,182            prev_message=prev_message,183            prev_query=prev_query,184            attempt=attempt_idx,185        )186 187        user_msg = USER_PROMPT_TEMPLATE.format(188            schema_sql=obs.schema_sql,189            task_description=obs.task_description,190            broken_query=obs.broken_query,191            feedback=feedback,192        )193 194        messages = [195            {"role": "system", "content": SYSTEM_PROMPT},196            {"role": "user", "content": user_msg},197        ]198 199        sql_attempt = call_llm(client, model, messages)200 201        result = env.step(Action(sql_query=sql_attempt))202        attempt_num = attempt_idx + 1203 204        step_record = {205            "attempt": attempt_num,206            "query": sql_attempt,207            "score": result.reward.value,208            "message": result.reward.message,209            "breakdown": result.reward.breakdown,210        }211        steps.append(step_record)212 213        prev_score = result.reward.value214        prev_message = result.reward.message215        prev_query = sql_attempt216 217        final_result = result218 219        if result.done and result.reward.breakdown.get("correctness", 0) >= 0.7:220            break221 222        obs = result.observation223        time.sleep(0.3)  # Rate limit courtesy224 225    scores = [s["score"] for s in steps]226    return {227        "task_id": task_id,228        "difficulty": task.difficulty,229        "attempts": len(steps),230        "steps": steps,231        "scores_per_attempt": scores,232        "final_score": scores[-1] if scores else 0.0,233        "best_score": max(scores) if scores else 0.0,234        "solved": (235            final_result.reward.breakdown.get("correctness", 0) >= 0.7236            if final_result else False237        ),238    }239 240 241def run_baseline(242    provider: str = "groq",243    model: Optional[str] = None,244    api_key: Optional[str] = None,245    base_url: Optional[str] = None,246    task_ids: Optional[list[str]] = None,247    user_sql=None248) -> dict:249    """Run baseline agent on all tasks (or a subset). Returns structured results with traces."""250    config = PROVIDER_CONFIGS.get(provider, PROVIDER_CONFIGS["groq"])251    resolved_model = model or config["default_model"]252    client = make_client(provider=provider, api_key=api_key, base_url=base_url)253    env = SQLEnv()254    target_tasks = task_ids or list(TASKS.keys())255    task_results = []256 257    print(f"Provider: {provider} | Model: {resolved_model}")258    print("=" * 55)259 260    for task_id in target_tasks:261        if task_id not in TASKS:262            print(f"  Skipping unknown task: {task_id}")263            continue264        print(f"\nTask: {task_id} ({TASKS[task_id].difficulty})")265        result = run_task_with_agent(env, task_id, client, resolved_model, user_sql)266        task_results.append(result)267        print(268            f"  Score: {result['final_score']:.4f} | "269            f"Solved: {result['solved']} | "270            f"Attempts: {result['attempts']}"271        )272 273    if not task_results:274        return {"error": "No tasks ran."}275 276    avg_score = sum(r["final_score"] for r in task_results) / len(task_results)277    avg_best = sum(r["best_score"] for r in task_results) / len(task_results)278    solve_rate = sum(1 for r in task_results if r["solved"]) / len(task_results)279 280    summary = {281        "provider": provider,282        "model": resolved_model,283        "environment": "sql-query-debugging",284        "tasks": task_results,285        "aggregate": {286            "average_final_score": round(avg_score, 4),287            "average_best_score": round(avg_best, 4),288            "solve_rate": round(solve_rate, 4),289            "total_tasks": len(task_results),290            "tasks_solved": sum(1 for r in task_results if r["solved"]),291        },292    }293 294    print("\n" + "=" * 55)295    print(f"Avg score: {avg_score:.4f} | Solve rate: {solve_rate:.0%} | Tasks: {len(task_results)}")296    return summary297 298 299def run_comparison(provider: str = "groq") -> dict:300    """Compare all available models for the given provider."""301    config = PROVIDER_CONFIGS.get(provider, PROVIDER_CONFIGS["groq"])302    models = config["models"]303    all_results = {}304 305    print(f"\nRunning model comparison for provider: {provider}")306    print(f"Models: {models}\n")307 308    for model in models:309        print(f"\n{'='*55}")310        print(f"MODEL: {model}")311        try:312            result = run_baseline(provider=provider, model=model)313            all_results[model] = result["aggregate"]314        except Exception as e:315            print(f"  Failed: {e}")316            all_results[model] = {"error": str(e)}317 318    return all_results319 320 321def run_comparison_for_task(provider: str = "groq", task_id: Optional[str] = None) -> list[dict]:322    """323    Compare all models for a given provider on a single task (or all tasks).324    Returns a list of {model, score, solved, attempts} dicts sorted by score.325    """326    config = PROVIDER_CONFIGS.get(provider, PROVIDER_CONFIGS["groq"])327    models = config["models"]328    results = []329 330    for model in models:331        try:332            client = make_client(provider=provider)333            env = SQLEnv()334            target_ids = [task_id] if task_id else list(TASKS.keys())335            task_results = []336            for tid in target_ids:337                if tid not in TASKS:338                    continue339                r = run_task_with_agent(env, tid, client, model)340                task_results.append(r)341 342            if not task_results:343                continue344 345            avg_score = sum(r["final_score"] for r in task_results) / len(task_results)346            total_attempts = sum(r["attempts"] for r in task_results)347            all_solved = all(r["solved"] for r in task_results)348 349            results.append({350                "model": model,351                "score": round(avg_score, 4),352                "solved": all_solved,353                "attempts": total_attempts,354                "tasks": task_results,355            })356        except Exception as e:357            results.append({358                "model": model,359                "score": 0.0,360                "solved": False,361                "attempts": 0,362                "error": str(e),363            })364 365    return sorted(results, key=lambda x: x["score"], reverse=True)366 367 368if __name__ == "__main__":369    parser = argparse.ArgumentParser(description="SQL Debugging Env Baseline Runner")370    parser.add_argument("--provider", default="groq", choices=list(PROVIDER_CONFIGS.keys()),371                        help="API provider (default: groq)")372    parser.add_argument("--model", default=None, help="Model name override")373    parser.add_argument("--api-key", default=None, help="API key override")374    parser.add_argument("--base-url", default=None, help="Base URL override for custom endpoints")375    parser.add_argument("--compare", action="store_true", help="Compare all models for the provider")376    parser.add_argument("--tasks", nargs="+", default=None, help="Specific task IDs to run")377    args = parser.parse_args()378 379    if args.compare:380        results = run_comparison(provider=args.provider)381        print("\nFull comparison results:")382        print(json.dumps(results, indent=2))383    else:384        results = run_baseline(385            provider=args.provider,386            model=args.model,387            api_key=args.api_key,388            base_url=args.base_url,389            task_ids=args.tasks,390        )391        print("\nFull results:")392        print(json.dumps(results, indent=2))