Kalletlamadhav/sql-optimization-env
0
1# inference.py2 3import asyncio, json, os, sys, time4from typing import List5from openai import OpenAI6from server.models import SQLOptAction, AntiPatternType7 8API_BASE_URL = os.environ.get('API_BASE_URL', 'https://api.openai.com/v1')9API_KEY = os.environ.get('OPENAI_API_KEY', os.environ.get('HF_TOKEN', ''))10MODEL_NAME = os.environ.get('MODEL_NAME', 'gpt-4o-mini')11ENV_BASE_URL = os.environ.get('ENV_URL', 'http://localhost:7860')12 13TEMPERATURE = 0.014MAX_TOKENS = 102415MAX_STEPS = 316SUCCESS_THRESHOLD = 0.717 18BENCHMARK_TASKS = [19 'gst_missing_index',20 'gst_n_plus_one',21 'gst_multi_join',22]23 24SYSTEM_PROMPT = '''You are an expert SQL database engineer.25You will be given a slow SQL query, its schema, and execution plan.26Your job is to:271. Identify the anti-pattern type causing the slowness282. Explain why the query is slow293. Rewrite the query to be faster304. Add appropriate CREATE INDEX statements31Respond ONLY in valid JSON with this exact structure:32{33 "optimized_query": "SELECT ...",34 "identified_pattern": "MISSING_INDEX",35 "explanation": "The query does a full table scan because...",36 "index_statements": ["CREATE INDEX idx_name ON table(col)"],37 "schema_analysis": "Table has 5 columns, no indexes on..."38}39identified_pattern must be one of: N_PLUS_ONE, CARTESIAN_PRODUCT, MISSING_INDEX,40SELECT_STAR, LEADING_WILDCARD, IMPLICIT_CAST, UNBOUNDED_AGGREGATION, NONE'''41 42 43def log_start(task: str, env: str, model: str):44 print(json.dumps({45 "event": "START",46 "task": task,47 "env": env,48 "model": model,49 "timestamp": time.time()50 }), flush=True)51 52 53def log_step(step: int, action: str, reward: float, done: bool, error=None):54 print(json.dumps({55 "event": "STEP",56 "step": step,57 "action": action[:200] if action else "",58 "reward": round(reward, 4),59 "done": done,60 "error": str(error) if error else None61 }), flush=True)62 63 64def log_end(success: bool, steps: int, score: float, rewards: List[float]):65 print(json.dumps({66 "event": "END",67 "success": success,68 "steps": steps,69 "score": round(score, 4),70 "rewards": [round(r,4) for r in rewards]71 }), flush=True)72 73 74def build_user_prompt(obs: dict) -> str:75 return f'''76TASK: {obs['goal']}77SCHEMA:78{obs['schema_ddl'][:500]}79CURRENT QUERY (SLOW):80{obs['current_query']}81EXECUTION PLAN:82{json.dumps(obs['execution_plan'], indent=2)}83EXECUTION TIME: {obs['execution_time_ms']:.0f}ms84DB STATS: {json.dumps(obs['db_stats'])}85Optimize this query. Respond in JSON only.'''86 87 88def get_model_action(client: OpenAI, obs: dict) -> dict:89 prompt = build_user_prompt(obs)90 91 try:92 completion = client.chat.completions.create(93 model=MODEL_NAME,94 messages=[95 {'role': 'system', 'content': SYSTEM_PROMPT},96 {'role': 'user', 'content': prompt}97 ],98 temperature=TEMPERATURE,99 max_tokens=MAX_TOKENS,100 stream=False101 )102 103 raw = (completion.choices[0].message.content or '').strip()104 105 if raw.startswith('```'):106 raw = raw.split('```')[1]107 if raw.startswith('json'):108 raw = raw[4:]109 110 return json.loads(raw)111 112 except Exception as e:113 print(f'[DEBUG] Model error: {e}', flush=True)114 115 return {116 'optimized_query': obs['current_query'],117 'identified_pattern': 'NONE',118 'explanation': 'Could not parse model response',119 'index_statements': [],120 'schema_analysis': ''121 }122 123 124def run_task(client: OpenAI, task_id: str) -> float:125 import requests126 127 log_start(task=task_id, env='sql-optimization-env', model=MODEL_NAME)128 129 rewards = []130 steps_taken = 0131 132 resp = requests.post(f'{ENV_BASE_URL}/reset', json={'task_id': task_id})133 obs = resp.json()134 135 last_reward = 0.0136 done = False137 138 for step in range(1, MAX_STEPS + 1):139 if done:140 break141 142 action_dict = get_model_action(client, obs)143 144 step_resp = requests.post(f'{ENV_BASE_URL}/step', json=action_dict)145 result = step_resp.json()146 147 last_reward = result.get('reward', 0.0)148 done = result.get('done', False)149 obs = result.get('observation', obs)150 151 rewards.append(last_reward)152 steps_taken = step153 154 log_step(step=step,155 action=action_dict.get('optimized_query',''),156 reward=last_reward,157 done=done)158 159 score = rewards[-1] if rewards else 0.0160 score = max(0.0, min(1.0, score))161 162 success = score >= SUCCESS_THRESHOLD163 164 log_end(success=success,165 steps=steps_taken,166 score=score,167 rewards=rewards)168 169 return score170 171 172def main():173 client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)174 175 print('[DEBUG] Starting SQL Optimization Benchmark', flush=True)176 print(f'[DEBUG] Model: {MODEL_NAME}', flush=True)177 print(f'[DEBUG] Tasks: {BENCHMARK_TASKS}', flush=True)178 179 all_scores = {}180 181 for task_id in BENCHMARK_TASKS:182 print(f'[DEBUG] Running task: {task_id}', flush=True)183 184 score = run_task(client, task_id)185 all_scores[task_id] = score186 187 print(f'[DEBUG] Task {task_id} score: {score:.4f}', flush=True)188 189 avg = sum(all_scores.values()) / len(all_scores)190 191 # ✅ FIXED LINE + INDENTATION192 print(json.dumps({193 'event': 'BENCHMARK_COMPLETE',194 'scores': all_scores,195 'average': round(avg, 4)196 }), flush=True)197 198 199if __name__ == '__main__':200 main()