vinayaknandi05/sql-optimization-openenv
0
1"""2Baseline Inference Script — SQL Optimization OpenEnv3=====================================================4Reads API credentials from environment variables:5 API_BASE_URL - LLM endpoint (OpenAI-compatible)6 API_KEY - API key7 MODEL_NAME - model identifier8 HF_TOKEN - Hugging Face token (for Space access)9 10Produces structured [START] / [STEP] / [END] stdout logs as required.11"""12 13import asyncio14import json15import os16import sys17import time18from typing import Optional19 20from openai import AsyncOpenAI21 22# ── Config ────────────────────────────────────────────────────────────────────23 24API_BASE_URL: str = os.environ.get("API_BASE_URL", "https://api.openai.com/v1")25API_KEY: str = os.environ.get("API_KEY", "sk-placeholder")26MODEL_NAME: str = os.environ.get("MODEL_NAME", "gpt-4o-mini")27HF_TOKEN: str = os.environ.get("HF_TOKEN", "")28IMAGE_NAME: str = os.environ.get("IMAGE_NAME", "sql-optimization-openenv:latest")29 30TASK_NAME: str = "sql_optimization"31BENCHMARK: str = "sql-openenv"32MAX_STEPS: int = 1033TEMPERATURE: float = 0.034MAX_TOKENS: int = 51235 36# ── Logging helpers (required format) ────────────────────────────────────────37 38def log_start(task: str, env: str, model: str):39 print(f"[START] task={task}", flush=True)40 41def log_step(step: int, action: dict, observation: dict, reward: float, done: bool):42 print(f"[STEP] step={step} reward={reward}", flush=True)43 44def log_end(task: str, score: float, steps: int, success: bool):45 print(f"[END] task={task} score={score} steps={steps}", flush=True)46 47# ── Local env shim (avoids needing a running Docker container for baseline) ───48 49class MyEnvV4Env:50 """Thin wrapper around the local environment for baseline testing."""51 def __init__(self):52 sys.path.insert(0, os.path.dirname(__file__))53 from env.environment import SQLOptimizationEnv, SQLAction as _SQLAction54 self._env = SQLOptimizationEnv()55 self._SQLAction = _SQLAction56 57 async def reset(self):58 obs = self._env.reset()59 return _ObsWrapper(obs)60 61 async def step(self, action):62 obs, reward, done, info = self._env.step(action)63 return _StepResult(obs, reward, done, info)64 65 66class _ObsWrapper:67 def __init__(self, obs):68 self.observation = obs69 70 71class _StepResult:72 def __init__(self, obs, reward, done, info):73 self.observation = obs74 self.reward = reward75 self.done = done76 self.info = info77 78 79class MyEnvV4Action:80 """Compatibility shim matching competition interface."""81 def __init__(self, message: str):82 sys.path.insert(0, os.path.dirname(__file__))83 from env.environment import SQLAction84 # Parse the message as JSON or plain SQL85 try:86 data = json.loads(message)87 self._action = SQLAction(**data)88 except Exception:89 self._action = SQLAction(query=message, message="")90 91 def __new__(cls, message: str):92 obj = object.__new__(cls)93 obj.__init__(message)94 return obj._action95 96 97# ── LLM call ─────────────────────────────────────────────────────────────────98 99async def get_model_message(100 client: AsyncOpenAI,101 step: int,102 last_echoed: str,103 last_reward: float,104 history: list[str],105) -> str:106 system_prompt = (107 "You are an expert SQL optimization agent. "108 "You will be given a poorly written SQL query and a database schema. "109 "Your job is to rewrite the query to be correct, efficient, and follow best practices. "110 "Respond ONLY with a JSON object: {\"query\": \"<optimized SQL>\", \"message\": \"<brief explanation>\"}. "111 "No markdown, no explanation outside the JSON."112 )113 messages = [114 {"role": "system", "content": system_prompt},115 ]116 for h in history:117 messages.append({"role": "user", "content": h})118 messages.append({"role": "user", "content": f"Step {step}. Last observation: {last_echoed}\nLast reward: {last_reward}"})119 120 try:121 completion = await client.chat.completions.create(122 model=MODEL_NAME,123 messages=messages,124 temperature=TEMPERATURE,125 max_tokens=MAX_TOKENS,126 stream=False,127 )128 text = (completion.choices[0].message.content or "").strip()129 return text if text else '{"query": "SELECT 1", "message": "fallback"}'130 except Exception as exc:131 print(f"[DEBUG] Model request failed: {exc}", flush=True)132 return '{"query": "SELECT 1", "message": "error fallback"}'133 134 135# ── Main loop ─────────────────────────────────────────────────────────────────136 137async def main() -> None:138 client = AsyncOpenAI(base_url=API_BASE_URL, api_key=API_KEY)139 140 # Run all 3 tasks141 task_ids = ["task_easy", "task_medium", "task_hard"]142 all_scores = []143 144 for task_id in task_ids:145 env = MyEnvV4Env()146 147 history: list[str] = []148 rewards: list[float] = []149 steps_taken = 0150 score = 0.0151 success = False152 153 log_start(task=f"{TASK_NAME}/{task_id}", env=BENCHMARK, model=MODEL_NAME)154 155 try:156 result = await env.reset()157 obs = result.observation158 # Include task context in first history entry159 history.append(160 f"Task: {obs.task_description}\n\nOriginal Query:\n{obs.original_query}\n\nSchema:\n{obs.schema_info}"161 )162 last_echoed = obs.echoed_message163 last_reward = 0.0164 165 for step in range(1, MAX_STEPS + 1):166 if hasattr(obs, "done") and obs.done:167 break168 169 message = await get_model_message(client, step, last_echoed, last_reward, history)170 171 from env.environment import SQLAction172 try:173 action_data = json.loads(message)174 action = SQLAction(**action_data)175 except Exception:176 action = SQLAction(query=message, message="")177 178 result = await env.step(action)179 obs = result.observation180 reward = result.reward181 done = result.done182 183 steps_taken = step184 rewards.append(reward)185 score = obs.score186 last_echoed = obs.echoed_message187 last_reward = reward188 189 history.append(f"Observation: {obs.echoed_message}")190 if obs.last_query_result:191 history.append(f"Query result (first 5 rows): {obs.last_query_result}")192 if obs.last_query_error:193 history.append(f"Query error: {obs.last_query_error}")194 195 log_step(196 step=step,197 action=action.model_dump(),198 observation=obs.model_dump(),199 reward=reward,200 done=done,201 )202 203 if done:204 break205 206 success = score >= 0.90207 all_scores.append(score)208 209 except Exception as exc:210 print(f"[DEBUG] Episode error: {exc}", flush=True)211 all_scores.append(0.0)212 213 log_end(task=f"{TASK_NAME}/{task_id}", score=score, steps=steps_taken, success=success)214 215 print(json.dumps({"type": "SUMMARY", "scores": dict(zip(task_ids, all_scores)), "mean": round(sum(all_scores)/len(all_scores), 4)}), flush=True)216 217 218def entry_point():219 """Entry point for package installation that handles async main."""220 asyncio.run(main())221 222 223if __name__ == "__main__":224 asyncio.run(main())225 