CoolFace
Apppublic

vinayaknandi05/sql-optimization-openenv

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