Team Ai
Apppublic

Kalletlamadhav/sql-optimization-env

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