Tsah00/sql-env
0
1"""2inference.py - Baseline inference script for the SQL Query Learning Environment.3 4MANDATORY STDOUT FORMAT:5 [START] task=<task_name> env=<benchmark> model=<model_name>6 [STEP] step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>7 [END] success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...,rn>8 9Environment variables:10 API_BASE_URL The API endpoint for the LLM (default: HF router)11 MODEL_NAME The model identifier (default: Qwen/Qwen2.5-72B-Instruct)12 HF_TOKEN Hugging Face / API key13 14Usage:15 python inference.py16"""17 18from __future__ import annotations19 20import os21import sys22import textwrap23from typing import List, Optional24 25from openai import OpenAI26 27# ── Config ──────────────────────────────────────────────────────────────────28 29API_KEY = (30 os.getenv("HF_TOKEN")31 or os.getenv("OPENAI_API_KEY")32 or os.getenv("API_KEY")33)34API_BASE_URL = os.getenv("API_BASE_URL") or "https://router.huggingface.co/v1"35MODEL_NAME = os.getenv("MODEL_NAME") or "Qwen/Qwen2.5-72B-Instruct"36 37BENCHMARK = "sql_env"38MAX_STEPS = 1039TEMPERATURE = 0.140MAX_TOKENS = 51241 42# All 9 tasks across 3 difficulty tiers43ALL_TASKS = [44 ("easy_1", "easy"),45 ("easy_2", "easy"),46 ("easy_3", "easy"),47 ("medium_1", "medium"),48 ("medium_2", "medium"),49 ("medium_3", "medium"),50 ("hard_1", "hard"),51 ("hard_2", "hard"),52 ("hard_3", "hard"),53]54 55SYSTEM_PROMPT = textwrap.dedent("""56You are a data analyst at an e-commerce company. Business stakeholders57(marketing, finance, CRM, merchandising) submit ad-hoc data requests and58you fulfil them by writing SQL queries against the company database.59 60RULES:61- Return ONLY the SQL query — no explanation, no markdown fences, no comments.62- Use standard SQLite syntax.63- Use column aliases to match the expected column names exactly as stated in the task.64- Do NOT use SELECT * — select only the required columns.65- Aim for the simplest correct query; avoid unnecessary subqueries or CROSS JOINs.66""").strip()67 68 69# ── Structured Logging (MANDATORY FORMAT) ───────────────────────────────────70 71def log_start(task: str, env: str, model: str) -> None:72 print(f"[START] task={task} env={env} model={model}", flush=True)73 74 75def log_step(step: int, action: str, reward: float, done: bool,76 error: Optional[str]) -> None:77 # Sanitize action string: remove newlines, limit length78 action_clean = action.replace("\n", " ").replace("\r", "").strip()79 error_val = error if error else "null"80 done_val = str(done).lower()81 print(82 f"[STEP] step={step} action={action_clean} "83 f"reward={reward:.2f} done={done_val} error={error_val}",84 flush=True,85 )86 87 88def log_end(success: bool, steps: int, score: float,89 rewards: List[float]) -> None:90 rewards_str = ",".join(f"{r:.2f}" for r in rewards)91 print(92 f"[END] success={str(success).lower()} steps={steps} "93 f"score={score:.2f} rewards={rewards_str}",94 flush=True,95 )96 97 98# ── LLM Query ──────────────────────────────────────────────────────────────99 100def get_sql_from_llm(101 client: OpenAI,102 schema_info: str,103 task_description: str,104 expected_columns: List[str],105 previous_attempts: List[dict],106) -> str:107 """Ask the LLM to produce a SQL query for the given task."""108 attempts_text = ""109 if previous_attempts:110 last = previous_attempts[-1]111 attempts_text = (112 f"\nPREVIOUS ATTEMPT:\n"113 f" Query: {last['query']}\n"114 f" Reward: {last['reward']}\n"115 f" Feedback: {last['message']}\n"116 f" Error: {last.get('error', 'none')}\n"117 f"\nImprove on this attempt.\n"118 )119 120 user_prompt = (121 f"DATABASE SCHEMA:\n{schema_info}\n\n"122 f"TASK:\n{task_description}\n\n"123 f"EXPECTED OUTPUT COLUMNS: {', '.join(expected_columns)}\n"124 f"{attempts_text}\n"125 f"SQL QUERY:"126 )127 128 try:129 completion = client.chat.completions.create(130 model=MODEL_NAME,131 messages=[132 {"role": "system", "content": SYSTEM_PROMPT},133 {"role": "user", "content": user_prompt},134 ],135 temperature=TEMPERATURE,136 max_tokens=MAX_TOKENS,137 stream=False,138 )139 text = (completion.choices[0].message.content or "").strip()140 # Strip markdown fences if present141 if text.startswith("```"):142 lines = text.split("\n")143 lines = [l for l in lines if not l.startswith("```")]144 text = "\n".join(lines).strip()145 return text if text else _fallback_query(task_description)146 except Exception as exc:147 print(f"[DEBUG] LLM request failed: {exc}", flush=True)148 return _fallback_query(task_description)149 150 151def _fallback_query(description: str) -> str:152 """153 Deterministic rule-based fallback when LLM API is unavailable.154 Pattern-matches task descriptions to known reference queries.155 """156 d = description.lower()157 158 # Easy159 if "usa" in d or "united states" in d:160 return "SELECT name, email FROM customers WHERE country = 'USA'"161 if "count" in d and "completed" in d:162 return "SELECT COUNT(*) AS total_completed FROM orders WHERE status = 'completed'"163 if "top 5" in d and "expensive" in d:164 return "SELECT name, category, price FROM products ORDER BY price DESC LIMIT 5"165 166 # Hard (checked before medium to avoid substring collisions)167 if "above" in d and "average" in d:168 return (169 "WITH customer_totals AS ("170 " SELECT c.name, SUM(o.total_amount) AS total_spent"171 " FROM customers c"172 " JOIN orders o ON c.id = o.customer_id"173 " WHERE o.status = 'completed'"174 " GROUP BY c.id, c.name"175 ") "176 "SELECT name, total_spent FROM customer_totals "177 "WHERE total_spent > (SELECT AVG(total_spent) FROM customer_totals) "178 "ORDER BY total_spent DESC"179 )180 if "best-selling" in d or "best selling" in d:181 return (182 "WITH product_sales AS ("183 " SELECT p.category, p.name AS product_name, SUM(oi.quantity) AS total_quantity"184 " FROM products p JOIN order_items oi ON p.id = oi.product_id"185 " GROUP BY p.id, p.category, p.name"186 "), ranked AS ("187 " SELECT category, product_name, total_quantity,"188 " RANK() OVER (PARTITION BY category ORDER BY total_quantity DESC) AS rnk"189 " FROM product_sales"190 ") "191 "SELECT category, product_name, total_quantity FROM ranked WHERE rnk = 1 "192 "ORDER BY category, product_name"193 )194 if "2022" in d and "2023" in d and "2024" in d:195 return (196 "SELECT c.name, c.email FROM customers c "197 "WHERE (SELECT COUNT(DISTINCT STRFTIME('%Y', o.order_date)) "198 "FROM orders o WHERE o.customer_id = c.id "199 "AND STRFTIME('%Y', o.order_date) IN ('2022','2023','2024')) = 3 "200 "ORDER BY c.name ASC"201 )202 203 # Medium204 if "total spending" in d or "total spent" in d:205 return (206 "SELECT c.name, SUM(o.total_amount) AS total_spent "207 "FROM customers c JOIN orders o ON c.id = o.customer_id "208 "WHERE o.status = 'completed' "209 "GROUP BY c.id, c.name ORDER BY total_spent DESC"210 )211 if "never" in d and ("order" in d or "appear" in d):212 return (213 "SELECT p.name, p.category, p.price FROM products p "214 "LEFT JOIN order_items oi ON p.id = oi.product_id WHERE oi.id IS NULL"215 )216 if "average" in d and "month" in d:217 return (218 "SELECT STRFTIME('%Y-%m', order_date) AS month, "219 "ROUND(AVG(total_amount), 2) AS avg_order_value "220 "FROM orders WHERE order_date LIKE '2023%' "221 "GROUP BY month ORDER BY month ASC"222 )223 224 return "SELECT name FROM customers LIMIT 5"225 226 227# ── Run One Task Episode ───────────────────────────────────────────────────228 229def run_task(230 task_id: str,231 difficulty: str,232 client: OpenAI,233) -> float:234 """235 Run a single task as one episode. Returns the best score in [0, 1].236 237 Emits [START], [STEP]..., [END] to stdout.238 """239 # Import env locally to keep module-level clean240 sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))241 from server.sql_environment import SQLEnvironment242 from models import SQLAction243 244 env = SQLEnvironment()245 rewards: List[float] = []246 steps_taken = 0247 best_score = 0.0248 success = False249 250 log_start(task=task_id, env=BENCHMARK, model=MODEL_NAME)251 252 try:253 obs = env.reset(difficulty=difficulty, task_id=task_id)254 previous_attempts: List[dict] = []255 256 for step in range(1, MAX_STEPS + 1):257 if obs.done:258 break259 260 # Get query from LLM or fallback261 if API_KEY and API_BASE_URL:262 query = get_sql_from_llm(263 client,264 schema_info=obs.schema_info,265 task_description=obs.task_description,266 expected_columns=obs.expected_columns,267 previous_attempts=previous_attempts,268 )269 else:270 query = _fallback_query(obs.task_description)271 272 action = SQLAction(query=query, difficulty=difficulty, task_id=task_id)273 obs = env.step(action)274 275 reward = obs.reward276 done = obs.done277 error = obs.error if obs.error else None278 279 rewards.append(reward)280 steps_taken = step281 best_score = max(best_score, reward)282 283 log_step(284 step=step,285 action=query,286 reward=reward,287 done=done,288 error=error,289 )290 291 previous_attempts.append({292 "query": query,293 "reward": reward,294 "message": obs.message,295 "error": obs.error,296 })297 if len(previous_attempts) > 2:298 previous_attempts.pop(0)299 300 # Near-perfect score — stop early301 if reward >= 0.99:302 break303 304 if done:305 break306 307 # Final score — clamped to open interval (0.01, 0.99) per Phase 2 spec308 score = min(0.99, max(0.01, best_score))309 success = score >= 0.5310 311 except Exception as exc:312 print(f"[DEBUG] Exception during episode: {exc}", flush=True)313 score = 0.0314 315 finally:316 try:317 env.close()318 except Exception:319 pass320 log_end(321 success=success,322 steps=steps_taken,323 score=score,324 rewards=rewards,325 )326 327 return score328 329 330# ── Main ────────────────────────────────────────────────────────────────────331 332def main() -> None:333 client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY or "sk-placeholder")334 335 total_score = 0.0336 task_scores = {}337 338 for task_id, difficulty in ALL_TASKS:339 score = run_task(task_id, difficulty, client)340 task_scores[task_id] = score341 total_score += score342 343 # Summary (not part of mandatory format — informational only)344 print("\n" + "=" * 60, flush=True)345 print("INFERENCE SUMMARY", flush=True)346 print("=" * 60, flush=True)347 for task_id, score in task_scores.items():348 status = "PASS" if score >= 0.5 else "FAIL"349 print(f" [{status}] {task_id}: score={score:.2f}", flush=True)350 avg = total_score / len(ALL_TASKS) if ALL_TASKS else 0.0351 print(f"\n Average score: {avg:.2f} ({total_score:.2f}/{len(ALL_TASKS)})", flush=True)352 353 354if __name__ == "__main__":355 main()356 