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