dkAmulet/sql-query-optimizer
0
1#!/usr/bin/env python32"""3Baseline inference script for SQL Query Optimizer Environment.4 5Required environment variables:6 API_BASE_URL - Base URL for the LLM API7 MODEL_NAME - Model identifier8 HF_TOKEN - API key (no default)9 10Usage:11 set API_BASE_URL=https://api.openai.com/v112 set MODEL_NAME=gpt-4o-mini13 set HF_TOKEN=sk-...14 python inference.py15"""16from __future__ import annotations17 18import json19import os20import sys21import time22import traceback23from typing import Dict, Optional24 25# ── Environment variables (required format) ───────────────────────────────────26API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1")27MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o-mini")28HF_TOKEN = os.getenv("HF_TOKEN")29 30TEMPERATURE = 0.131MAX_TOKENS = 102432 33SYSTEM_PROMPT = """\34You are an expert database engineer specialising in SQL query optimisation.35Rewrite slow SQL queries to be more efficient while returning the EXACT same result set.36 37Common optimisations:38- Replace SELECT * with only the required columns.39- Convert IN (SELECT ...) subqueries to explicit JOINs.40- Replace correlated subqueries in SELECT list with JOIN + GROUP BY.41- Exploit available indexes listed in the schema.42 43Return ONLY the optimised SQL, no markdown fences, no explanation.44"""45 46# ── Imports with error handling ───────────────────────────────────────────────47try:48 from openai import OpenAI49except ImportError as e:50 print(f"[ERROR] Failed to import openai: {e}")51 print("[ERROR] Run: pip install openai>=2.7.2")52 sys.exit(1)53 54try:55 from env import SQLQueryOptimizerEnv56 from models import SQLAction57 from tasks import TASK_ORDER, TASKS58except ImportError as e:59 print(f"[ERROR] Failed to import environment modules: {e}")60 print("[ERROR] Make sure pydantic and other deps are installed.")61 sys.exit(1)62 63 64# ── LLM helpers ───────────────────────────────────────────────────────────────65 66def build_user_message(obs_dict: dict, feedback: str) -> str:67 lines = [68 f"Task ({obs_dict['difficulty']}): {obs_dict['description']}",69 "",70 "=== Database Schema ===",71 obs_dict["schema_ddl"],72 "",73 "=== Slow Query to Optimise ===",74 obs_dict["slow_query"],75 ]76 if feedback:77 lines += [78 "",79 "=== Your Previous Attempt ===",80 obs_dict.get("current_query", ""),81 "",82 "=== Grader Feedback ===",83 feedback,84 "",85 "Fix the issues above and return an improved query.",86 ]87 else:88 lines += ["", "Write an optimised version of the slow query."]89 return "\n".join(lines)90 91 92def call_llm(client: OpenAI, user_message: str) -> str:93 try:94 response = client.chat.completions.create(95 model=MODEL_NAME,96 messages=[97 {"role": "system", "content": SYSTEM_PROMPT},98 {"role": "user", "content": user_message},99 ],100 temperature=TEMPERATURE,101 max_tokens=MAX_TOKENS,102 stream=False,103 )104 text = (response.choices[0].message.content or "").strip()105 # Strip markdown fences if present106 for fence in ("```sql", "```SQL", "```"):107 if text.startswith(fence):108 text = text[len(fence):]109 if text.endswith("```"):110 text = text[:-3]111 return text.strip()112 except Exception as e:113 print(f"[WARN] LLM call failed: {e}")114 return ""115 116 117# ── Fallback query per task (used if LLM fails) ───────────────────────────────118FALLBACK_QUERIES = {119 "select_star_removal": (120 "SELECT user_id, username, email FROM users WHERE is_active = 1"121 ),122 "subquery_to_join": (123 "SELECT o.order_id, o.user_id, o.total_amount "124 "FROM orders o JOIN users u ON u.user_id = o.user_id "125 "WHERE u.country = 'USA' AND u.is_active = 1 AND o.status = 'delivered'"126 ),127 "aggregation_optimization": (128 "SELECT c.name AS category_name, "129 "SUM(oi.quantity * oi.unit_price) AS total_revenue "130 "FROM categories c "131 "JOIN products p ON p.category_id = c.category_id "132 "JOIN order_items oi ON oi.product_id = p.product_id "133 "GROUP BY c.category_id, c.name "134 "HAVING SUM(oi.quantity * oi.unit_price) > 1000 "135 "ORDER BY total_revenue DESC"136 ),137}138 139 140# ── Per-task runner ───────────────────────────────────────────────────────────141 142def run_task(env: SQLQueryOptimizerEnv, client: Optional[OpenAI],143 task_id: str) -> float:144 task_meta = TASKS[task_id]145 146 print(f"[START] task_id={task_id} difficulty={task_meta['difficulty']} "147 f"max_steps={task_meta['max_steps']}")148 149 try:150 obs = env.reset(task_id)151 except Exception as e:152 print(f"[END] task_id={task_id} best_reward=0.0 error=RESET_FAILED")153 return 0.0154 155 obs_dict = obs.model_dump()156 feedback = ""157 best_reward: float = 0.0158 159 while True:160 step_num = obs.step_number + 1161 162 # Try LLM first, fall back to hardcoded optimal query163 optimised = ""164 if client is not None:165 optimised = call_llm(client, build_user_message(obs_dict, feedback))166 167 if not optimised:168 optimised = FALLBACK_QUERIES.get(task_id,169 "SELECT 1") # last resort170 print(f"[WARN] Using fallback query for {task_id}")171 172 try:173 result = env.step(SQLAction(optimized_query=optimised))174 reward = result.reward175 176 print(f"[STEP] task_id={task_id} step={step_num} "177 f"reward={reward.value:.4f} "178 f"validity={reward.breakdown.validity:.2f} "179 f"correctness={reward.breakdown.correctness:.2f} "180 f"performance={reward.breakdown.performance:.2f} "181 f"style={reward.breakdown.style:.2f} "182 f"done={result.done}")183 184 if reward.value > best_reward:185 best_reward = reward.value186 187 obs = result.observation188 obs_dict = obs.model_dump()189 feedback = reward.feedback190 191 if result.done:192 break193 194 except Exception as e:195 print(f"[STEP] task_id={task_id} step={step_num} "196 f"reward=0.0 error=STEP_FAILED detail={e}")197 break198 199 print(f"[END] task_id={task_id} best_reward={best_reward:.4f}")200 return best_reward201 202 203# ── Main ──────────────────────────────────────────────────────────────────────204 205def main() -> Dict[str, float]:206 print(f"[INFO] SQL Query Optimizer Baseline Inference")207 print(f"[INFO] API_BASE_URL={API_BASE_URL}")208 print(f"[INFO] MODEL_NAME={MODEL_NAME}")209 print(f"[INFO] HF_TOKEN={'set' if HF_TOKEN else 'NOT SET'}")210 211 # Build OpenAI client — handle missing token gracefully212 client: Optional[OpenAI] = None213 try:214 if HF_TOKEN:215 client = OpenAI(api_key=HF_TOKEN, base_url=API_BASE_URL)216 print("[INFO] OpenAI client initialised successfully")217 else:218 print("[WARN] HF_TOKEN not set — will use fallback queries only")219 except Exception as e:220 print(f"[WARN] Could not initialise OpenAI client: {e} — using fallbacks")221 222 # Initialise environment223 try:224 env = SQLQueryOptimizerEnv()225 except Exception as e:226 print(f"[ERROR] Failed to initialise environment: {e}")227 traceback.print_exc()228 sys.exit(1)229 230 results: Dict[str, float] = {}231 t0 = time.time()232 233 try:234 for task_id in TASK_ORDER:235 try:236 results[task_id] = run_task(env, client, task_id)237 except Exception as e:238 print(f"[ERROR] Task {task_id} failed unexpectedly: {e}")239 traceback.print_exc()240 results[task_id] = 0.0241 finally:242 try:243 env.close()244 except Exception:245 pass246 247 elapsed = time.time() - t0248 overall = sum(results.values()) / max(len(results), 1)249 250 print(f"\n[SUMMARY] overall_average={overall:.4f} elapsed_seconds={elapsed:.1f}")251 for task_id, score in results.items():252 print(f"[SUMMARY] task={task_id} score={score:.4f}")253 254 try:255 with open("baseline_results.json", "w") as fh:256 json.dump({257 "task_scores": results,258 "overall_average": overall,259 "model": MODEL_NAME,260 "elapsed_seconds": round(elapsed, 1),261 }, fh, indent=2)262 print("[INFO] Results saved to baseline_results.json")263 except Exception as e:264 print(f"[WARN] Could not save results file: {e}")265 266 return results267 268 269if __name__ == "__main__":270 try:271 results = main()272 # Exit 0 even if all scores are 0 — don't fail the pipeline273 sys.exit(0)274 except Exception as e:275 print(f"[ERROR] Unhandled exception: {e}")276 traceback.print_exc()277 sys.exit(1)278 