Codexzzz/sql-env
0
1# import json2# import sqlite33# import tempfile4# import threading5# from pathlib import Path6# from uuid import uuid47 8# from openenv.core.env_server.interfaces import Environment9# from openenv.core.env_server.types import State10 11# try:12# from ..models import SqlAction, SqlObservation13# except ImportError:14# from models import SqlAction, SqlObservation15 16 17# # ─────────────────────────────────────────────────────────────────────────────18# # MODULE-LEVEL SESSION STORE19# #20# # The OpenEnv HTTP server creates a NEW SqlEnvironment instance on every21# # request, so self._episode_id would always be None in step().22# # Storing sessions at module level (shared across all instances in the same23# # process) fixes this. We also persist to disk so the DB survives a restart.24# # ─────────────────────────────────────────────────────────────────────────────25 26# _MEMORY_SESSIONS: dict = {} # { episode_id -> session_dict, "__latest__" -> episode_id }27# _SESSION_DIR = Path(tempfile.gettempdir()) / "openenv_sql_env"28# _SESSION_DIR.mkdir(parents=True, exist_ok=True)29 30 31# # ─────────────────────────────────────────────────────────────────────────────32# # TASK DEFINITIONS33# # ─────────────────────────────────────────────────────────────────────────────34 35# TASKS = {36# "select_basics": {37# "description": (38# "Find the full name and email address of all customers who live in 'New York'. "39# "Return results sorted alphabetically by name (A to Z)."40# ),41# "schema": (42# "CREATE TABLE customers (\n"43# " id INTEGER PRIMARY KEY,\n"44# " name TEXT NOT NULL,\n"45# " email TEXT NOT NULL,\n"46# " city TEXT NOT NULL,\n"47# " age INTEGER\n"48# ");"49# ),50# "seed_sql": """51# INSERT INTO customers VALUES (1, 'Alice Brown', 'alice@email.com', 'New York', 28);52# INSERT INTO customers VALUES (2, 'Bob Smith', 'bob@email.com', 'New York', 34);53# INSERT INTO customers VALUES (3, 'Carol Davis', 'carol@email.com', 'Chicago', 25);54# INSERT INTO customers VALUES (4, 'David Lee', 'david@email.com', 'New York', 41);55# INSERT INTO customers VALUES (5, 'Eve Wilson', 'eve@email.com', 'Boston', 30);56# """,57# "expected": [58# ("Alice Brown", "alice@email.com"),59# ("Bob Smith", "bob@email.com"),60# ("David Lee", "david@email.com"),61# ],62# "max_steps": 5,63# },64 65# "aggregate_filter": {66# "description": (67# "Find each customer who has placed MORE THAN 2 orders. "68# "Return their name and total amount spent (sum of all their order amounts). "69# "Sort by total amount spent, highest first."70# ),71# "schema": (72# "CREATE TABLE customers (\n"73# " id INTEGER PRIMARY KEY,\n"74# " name TEXT NOT NULL\n"75# ");\n"76# "CREATE TABLE orders (\n"77# " id INTEGER PRIMARY KEY,\n"78# " customer_id INTEGER NOT NULL,\n"79# " amount REAL NOT NULL,\n"80# " order_date TEXT NOT NULL\n"81# ");"82# ),83# "seed_sql": """84# INSERT INTO customers VALUES (1, 'Alice Brown');85# INSERT INTO customers VALUES (2, 'Bob Smith');86# INSERT INTO customers VALUES (3, 'Carol Davis');87# INSERT INTO orders VALUES (1, 1, 120.00, '2024-01-10');88# INSERT INTO orders VALUES (2, 1, 85.50, '2024-01-15');89# INSERT INTO orders VALUES (3, 1, 200.00, '2024-02-01');90# INSERT INTO orders VALUES (4, 2, 45.00, '2024-01-20');91# INSERT INTO orders VALUES (5, 2, 95.00, '2024-02-10');92# INSERT INTO orders VALUES (6, 2, 160.00, '2024-02-15');93# INSERT INTO orders VALUES (7, 3, 30.00, '2024-01-05');94# """,95# "expected": [96# ("Alice Brown", 405.50),97# ("Bob Smith", 300.00),98# ],99# "max_steps": 5,100# },101 102# "multi_join": {103# "description": (104# "Generate a monthly revenue report for the year 2024. "105# "For each month and product category return: "106# "month in 'YYYY-MM' format, category name, "107# "number of distinct orders that included products from that category, "108# "and total revenue (quantity x product price, summed across all items). "109# "Order by month ascending, then total revenue descending within each month. "110# "Only include data from 2024 - exclude records from other years."111# ),112# "schema": (113# "CREATE TABLE categories (\n"114# " id INTEGER PRIMARY KEY,\n"115# " name TEXT NOT NULL\n"116# ");\n"117# "CREATE TABLE products (\n"118# " id INTEGER PRIMARY KEY,\n"119# " name TEXT NOT NULL,\n"120# " category_id INTEGER NOT NULL,\n"121# " price REAL NOT NULL\n"122# ");\n"123# "CREATE TABLE orders (\n"124# " id INTEGER PRIMARY KEY,\n"125# " order_date TEXT NOT NULL\n"126# ");\n"127# "CREATE TABLE order_items (\n"128# " id INTEGER PRIMARY KEY,\n"129# " order_id INTEGER NOT NULL,\n"130# " product_id INTEGER NOT NULL,\n"131# " quantity INTEGER NOT NULL\n"132# ");"133# ),134# "seed_sql": """135# INSERT INTO categories VALUES (1, 'Electronics');136# INSERT INTO categories VALUES (2, 'Books');137# INSERT INTO products VALUES (1, 'Laptop', 1, 999.00);138# INSERT INTO products VALUES (2, 'Phone', 1, 599.00);139# INSERT INTO products VALUES (3, 'Python Book', 2, 49.00);140# INSERT INTO products VALUES (4, 'SQL Handbook', 2, 39.00);141# INSERT INTO orders VALUES (1, '2024-01-15');142# INSERT INTO orders VALUES (2, '2024-01-20');143# INSERT INTO orders VALUES (3, '2024-02-10');144# INSERT INTO orders VALUES (4, '2024-02-28');145# INSERT INTO orders VALUES (5, '2023-12-01');146# INSERT INTO order_items VALUES (1, 1, 1, 1);147# INSERT INTO order_items VALUES (2, 2, 3, 2);148# INSERT INTO order_items VALUES (3, 2, 4, 1);149# INSERT INTO order_items VALUES (4, 3, 2, 1);150# INSERT INTO order_items VALUES (5, 4, 3, 3);151# INSERT INTO order_items VALUES (6, 5, 1, 1);152# """,153# "expected": [154# ("2024-01", "Electronics", 1, 999.00),155# ("2024-01", "Books", 1, 137.00),156# ("2024-02", "Electronics", 1, 599.00),157# ("2024-02", "Books", 1, 147.00),158# ],159# "max_steps": 7,160# },161 162# "data_anomalies": {163# "description": (164# "Find data quality issues in the customers table. "165# "Return: the type of issue as a string and the count of affected rows. "166# "The three issue types to check are:\n"167# " 1. 'duplicate_email' - email addresses that appear more than once\n"168# " 2. 'invalid_age' - age values that are NULL, negative, or greater than 150\n"169# " 3. 'null_name' - rows where name is NULL\n"170# "Return all three rows ordered alphabetically by issue type. "171# "Use UNION ALL to combine the three checks into one result set."172# ),173# "schema": (174# "CREATE TABLE customers (\n"175# " id INTEGER PRIMARY KEY,\n"176# " name TEXT,\n"177# " email TEXT,\n"178# " age INTEGER\n"179# ");"180# ),181# "seed_sql": """182# INSERT INTO customers VALUES (1, 'Alice', 'a@test.com', 25);183# INSERT INTO customers VALUES (2, NULL, 'b@test.com', 30);184# INSERT INTO customers VALUES (3, 'Carol', 'a@test.com', 22);185# INSERT INTO customers VALUES (4, 'Dave', 'd@test.com', -5);186# INSERT INTO customers VALUES (5, 'Eve', 'e@test.com', 200);187# """,188# "expected": [189# ("duplicate_email", 2),190# ("invalid_age", 2),191# ("null_name", 1),192# ],193# "max_steps": 7,194# },195# }196 197 198# # ─────────────────────────────────────────────────────────────────────────────199# # THREAD-BASED QUERY TIMEOUT (cross-platform, no SIGALRM)200# # ─────────────────────────────────────────────────────────────────────────────201 202# def _run_query_with_timeout(db_path: Path, sql: str, timeout_seconds: int = 5):203# """Execute SQL in a daemon thread with a hard timeout.204# Returns (rows, error_or_None).205# """206# result_box: list = [None]207# error_box: list = [None]208 209# def _target():210# conn = None211# try:212# conn = sqlite3.connect(str(db_path))213# cursor = conn.execute(sql)214# result_box[0] = cursor.fetchall()215# except Exception as exc:216# error_box[0] = exc217# finally:218# if conn:219# try:220# conn.close()221# except Exception:222# pass223 224# t = threading.Thread(target=_target, daemon=True)225# t.start()226# t.join(timeout=timeout_seconds)227 228# if t.is_alive():229# return None, TimeoutError(f"Query exceeded {timeout_seconds}s — avoid full table scans.")230# if error_box[0] is not None:231# return None, error_box[0]232# return result_box[0], None233 234 235# # ─────────────────────────────────────────────────────────────────────────────236# # ENVIRONMENT CLASS237# # ─────────────────────────────────────────────────────────────────────────────238 239# class SqlEnvironment(Environment):240# SUPPORTS_CONCURRENT_SESSIONS = True241 242# def __init__(self):243# self._episode_id: str | None = None244# self._state = State(episode_id=str(uuid4()), step_count=0)245 246# # ── File helpers (best-effort persistence) ────────────────────────────────247 248# def _db_path(self, episode_id: str) -> Path:249# return _SESSION_DIR / f"{episode_id}.sqlite3"250 251# def _meta_path(self, episode_id: str) -> Path:252# return _SESSION_DIR / f"session_{episode_id}.json"253 254# def _save_session_file(self, session: dict) -> None:255# try:256# path = self._meta_path(session["episode_id"])257# tmp = path.with_suffix(".tmp")258# with tmp.open("w", encoding="utf-8") as f:259# json.dump(session, f)260# tmp.replace(path)261# except Exception:262# pass # file persistence is best-effort; memory is primary263 264# def _load_session_file(self, episode_id: str) -> dict:265# try:266# path = self._meta_path(episode_id)267# if not path.exists():268# return {}269# with path.open("r", encoding="utf-8") as f:270# return json.load(f)271# except Exception:272# return {}273 274# # ── DB initialisation ─────────────────────────────────────────────────────275 276# def _initialise_db(self, task_name: str, db_path: Path) -> None:277# task = TASKS[task_name]278# if db_path.exists():279# try:280# db_path.unlink()281# except OSError:282# pass283# conn = sqlite3.connect(str(db_path))284# try:285# conn.executescript(task["schema"])286# conn.executescript(task["seed_sql"])287# conn.commit()288# finally:289# conn.close()290 291# # ── Session resolution ────────────────────────────────────────────────────292 293# def _resolve_session(self) -> dict:294# """Find the session for the current episode.295 296# Priority order:297# 1. self._episode_id (WebSocket / persistent instance where reset was called)298# 2. _MEMORY_SESSIONS["__latest__"] (HTTP, same process, new instance per request)299# 3. Disk fallback for the resolved episode_id (container restart recovery)300# """301# episode_id = self._episode_id or _MEMORY_SESSIONS.get("__latest__")302# if not episode_id:303# return {}304 305# session = _MEMORY_SESSIONS.get(episode_id, {})306# if not session:307# session = self._load_session_file(episode_id)308# if session:309# _MEMORY_SESSIONS[episode_id] = session310 311# if session:312# self._episode_id = episode_id # pin for this request313 314# return session315 316# # ── reset() ──────────────────────────────────────────────────────────────317 318# def reset(self, seed=None, episode_id=None, **kwargs) -> SqlObservation:319# task_name = kwargs.get("task", "select_basics")320# if task_name not in TASKS:321# task_name = "select_basics"322 323# task = TASKS[task_name]324# new_id = episode_id or str(uuid4())325# db_path = self._db_path(new_id)326 327# self._initialise_db(task_name, db_path)328 329# session = {330# "episode_id": new_id,331# "task_name": task_name,332# "db_path": str(db_path),333# "step_count": 0,334# }335 336# # Primary store: module-level dict (survives across instances in same process)337# _MEMORY_SESSIONS[new_id] = session338# _MEMORY_SESSIONS["__latest__"] = new_id339# # Secondary store: disk (survives container restart)340# self._save_session_file(session)341 342# self._episode_id = new_id343# self._state = State(episode_id=new_id, step_count=0)344 345# return SqlObservation(346# task_description = task["description"],347# schema_info = task["schema"],348# query_result = [],349# error_message = "",350# feedback = "Episode started. Write a SQL query to solve the task above.",351# score_breakdown = {},352# attempts_remaining = task["max_steps"],353# done = False,354# reward = 0.0,355# )356 357# # ── step() ────────────────────────────────────────────────────────────────358 359# def step(self, action: SqlAction) -> SqlObservation:360# session = self._resolve_session()361 362# if not session:363# return SqlObservation(364# task_description = "",365# schema_info = "",366# query_result = [],367# error_message = "No active session. Call /reset first.",368# feedback = "No active session — call /reset before /step.",369# score_breakdown = {"execute": -0.05},370# attempts_remaining = 0,371# done = True,372# reward = -0.05,373# )374 375# task_name = session.get("task_name", "select_basics")376# if task_name not in TASKS:377# task_name = "select_basics"378 379# db_path = Path(session.get("db_path") or str(self._db_path(session["episode_id"])))380 381# # Re-seed if the DB file was lost (e.g. tmpfs wipe on restart)382# if not db_path.exists():383# self._initialise_db(task_name, db_path)384 385# # Advance step counter386# step_count = int(session.get("step_count", 0)) + 1387# session["step_count"] = step_count388# _MEMORY_SESSIONS[self._episode_id] = session389# self._save_session_file(session)390 391# self._state = State(392# episode_id = self._episode_id or str(uuid4()),393# step_count = step_count,394# )395 396# task = TASKS[task_name]397# attempts_remaining = task["max_steps"] - step_count398 399# rows, err = _run_query_with_timeout(db_path, action.sql_query, timeout_seconds=5)400 401# if isinstance(err, TimeoutError):402# reward = -0.10403# rows = []404# feedback = str(err)405# breakdown = {"execute": -0.10}406# error_msg = str(err)407# elif err is not None:408# reward = -0.05409# rows = []410# feedback = f"SQL Error: {err}. Fix your syntax and try again."411# breakdown = {"execute": -0.05}412# error_msg = str(err)413# else:414# reward, feedback, breakdown = self._grade(rows, task["expected"], action.sql_query)415# error_msg = ""416 417# done = reward >= 0.95 or attempts_remaining <= 0418 419# return SqlObservation(420# task_description = task["description"],421# schema_info = task["schema"],422# query_result = [list(r) for r in (rows or [])],423# error_message = error_msg,424# feedback = feedback,425# score_breakdown = breakdown,426# attempts_remaining = max(0, attempts_remaining),427# done = done,428# reward = float(max(-0.10, min(1.0, reward))),429# )430 431# # ── Grader ────────────────────────────────────────────────────────────────432 433# def _grade(self, result: list, expected: list, sql_query: str):434# breakdown: dict = {"execute": 0.10}435 436# if not result:437# return (438# 0.10,439# "Query ran but returned 0 rows. Check your WHERE clause or JOIN conditions.",440# breakdown,441# )442 443# result_set = set(tuple(r) for r in result)444# expected_set = set(tuple(e) for e in expected)445 446# result_cols = len(result[0]) if result else 0447# expected_cols = len(expected[0]) if expected else 0448# col_score = (449# 0.20 if result_cols == expected_cols450# else 0.20 * (min(result_cols, expected_cols) / max(result_cols, expected_cols, 1))451# )452# breakdown["columns"] = round(col_score, 3)453 454# row_score = 0.20 * min(1.0, len(result) / max(len(expected), 1))455# breakdown["rows"] = round(row_score, 3)456 457# f1 = self._f1(result_set, expected_set)458# val_score = 0.40 * f1459# breakdown["values"] = round(val_score, 3)460 461# uses_star = "select*" in sql_query.lower().replace(" ", "")462# eff_score = 0.0 if uses_star else 0.10463# breakdown["efficiency"] = eff_score464 465# total = breakdown["execute"] + col_score + row_score + val_score + eff_score466# pct = int(f1 * 100)467 468# if f1 >= 1.0 and col_score >= 0.20 and row_score >= 0.20:469# feedback = (470# "Perfect! Exact match."471# if not uses_star472# else "Correct result but avoid SELECT * — target only needed columns."473# )474# elif result_cols > expected_cols:475# feedback = (476# f"Too many columns ({result_cols} returned, {expected_cols} expected). "477# "Remove extra columns from SELECT."478# )479# elif result_cols < expected_cols:480# feedback = (481# f"Too few columns ({result_cols} returned, {expected_cols} expected). "482# "Add missing columns to SELECT."483# )484# elif len(result) > len(expected) * 1.5:485# feedback = (486# f"Too many rows ({len(result)} vs {len(expected)} expected). "487# "Check your WHERE or HAVING — a filter may be missing."488# )489# elif len(result) < len(expected):490# feedback = (491# f"Too few rows ({len(result)} vs {len(expected)} expected). "492# "Check your JOIN or WHERE — some matching rows are being excluded."493# )494# elif f1 >= 0.8:495# feedback = f"Very close! {pct}% of values match. Check column ordering or data type casting."496# elif f1 >= 0.5:497# feedback = f"Partial match: {pct}% correct. Re-read the task and check your filters."498# else:499# feedback = f"Mostly incorrect ({pct}% match). Start from the schema and re-read the task."500 501# return round(total, 3), feedback, breakdown502 503# @staticmethod504# def _f1(result_set: set, expected_set: set) -> float:505# if not result_set and not expected_set:506# return 1.0507# if not result_set or not expected_set:508# return 0.0509# intersection = result_set & expected_set510# precision = len(intersection) / len(result_set)511# recall = len(intersection) / len(expected_set)512# if precision + recall == 0:513# return 0.0514# return 2 * precision * recall / (precision + recall)515 516# @property517# def state(self) -> State:518# return self._state519 520 521import json522import sqlite3523import tempfile524import threading525from pathlib import Path526from uuid import uuid4527 528from openenv.core.env_server.interfaces import Environment529from openenv.core.env_server.types import State530 531try:532 from ..models import SqlAction, SqlObservation533except ImportError:534 from models import SqlAction, SqlObservation535 536 537# ─────────────────────────────────────────────────────────────────────────────538# MODULE-LEVEL SESSION STORE539#540# The OpenEnv HTTP server creates a NEW SqlEnvironment instance on every541# request, so self._episode_id would always be None in step().542# Storing sessions at module level (shared across all instances in the same543# process) fixes this. We also persist to disk so the DB survives a restart.544# ─────────────────────────────────────────────────────────────────────────────545 546_MEMORY_SESSIONS: dict = {} # { episode_id -> session_dict, "__latest__" -> episode_id }547_SESSION_DIR = Path(tempfile.gettempdir()) / "openenv_sql_env"548_SESSION_DIR.mkdir(parents=True, exist_ok=True)549 550 551# ─────────────────────────────────────────────────────────────────────────────552# TASK DEFINITIONS553# ─────────────────────────────────────────────────────────────────────────────554 555TASKS = {556 "select_basics": {557 "description": (558 "Find the full name and email address of all customers who live in 'New York'. "559 "Return results sorted alphabetically by name (A to Z)."560 ),561 "schema": (562 "CREATE TABLE customers (\n"563 " id INTEGER PRIMARY KEY,\n"564 " name TEXT NOT NULL,\n"565 " email TEXT NOT NULL,\n"566 " city TEXT NOT NULL,\n"567 " age INTEGER\n"568 ");"569 ),570 "seed_sql": """571INSERT INTO customers VALUES (1, 'Alice Brown', 'alice@email.com', 'New York', 28);572INSERT INTO customers VALUES (2, 'Bob Smith', 'bob@email.com', 'New York', 34);573INSERT INTO customers VALUES (3, 'Carol Davis', 'carol@email.com', 'Chicago', 25);574INSERT INTO customers VALUES (4, 'David Lee', 'david@email.com', 'New York', 41);575INSERT INTO customers VALUES (5, 'Eve Wilson', 'eve@email.com', 'Boston', 30);576""",577 "expected": [578 ("Alice Brown", "alice@email.com"),579 ("Bob Smith", "bob@email.com"),580 ("David Lee", "david@email.com"),581 ],582 "max_steps": 5,583 },584 585 "aggregate_filter": {586 "description": (587 "Find each customer who has placed MORE THAN 2 orders. "588 "Return their name and total amount spent (sum of all their order amounts). "589 "Sort by total amount spent, highest first."590 ),591 "schema": (592 "CREATE TABLE customers (\n"593 " id INTEGER PRIMARY KEY,\n"594 " name TEXT NOT NULL\n"595 ");\n"596 "CREATE TABLE orders (\n"597 " id INTEGER PRIMARY KEY,\n"598 " customer_id INTEGER NOT NULL,\n"599 " amount REAL NOT NULL,\n"600 " order_date TEXT NOT NULL\n"601 ");"602 ),603 "seed_sql": """604INSERT INTO customers VALUES (1, 'Alice Brown');605INSERT INTO customers VALUES (2, 'Bob Smith');606INSERT INTO customers VALUES (3, 'Carol Davis');607INSERT INTO orders VALUES (1, 1, 120.00, '2024-01-10');608INSERT INTO orders VALUES (2, 1, 85.50, '2024-01-15');609INSERT INTO orders VALUES (3, 1, 200.00, '2024-02-01');610INSERT INTO orders VALUES (4, 2, 45.00, '2024-01-20');611INSERT INTO orders VALUES (5, 2, 95.00, '2024-02-10');612INSERT INTO orders VALUES (6, 2, 160.00, '2024-02-15');613INSERT INTO orders VALUES (7, 3, 30.00, '2024-01-05');614""",615 "expected": [616 ("Alice Brown", 405.50),617 ("Bob Smith", 300.00),618 ],619 "max_steps": 5,620 },621 622 "multi_join": {623 "description": (624 "Generate a monthly revenue report for the year 2024. "625 "For each month and product category return: "626 "month in 'YYYY-MM' format, category name, "627 "number of distinct orders that included products from that category, "628 "and total revenue (quantity x product price, summed across all items). "629 "Order by month ascending, then total revenue descending within each month. "630 "Only include data from 2024 - exclude records from other years."631 ),632 "schema": (633 "CREATE TABLE categories (\n"634 " id INTEGER PRIMARY KEY,\n"635 " name TEXT NOT NULL\n"636 ");\n"637 "CREATE TABLE products (\n"638 " id INTEGER PRIMARY KEY,\n"639 " name TEXT NOT NULL,\n"640 " category_id INTEGER NOT NULL,\n"641 " price REAL NOT NULL\n"642 ");\n"643 "CREATE TABLE orders (\n"644 " id INTEGER PRIMARY KEY,\n"645 " order_date TEXT NOT NULL\n"646 ");\n"647 "CREATE TABLE order_items (\n"648 " id INTEGER PRIMARY KEY,\n"649 " order_id INTEGER NOT NULL,\n"650 " product_id INTEGER NOT NULL,\n"651 " quantity INTEGER NOT NULL\n"652 ");"653 ),654 "seed_sql": """655INSERT INTO categories VALUES (1, 'Electronics');656INSERT INTO categories VALUES (2, 'Books');657INSERT INTO products VALUES (1, 'Laptop', 1, 999.00);658INSERT INTO products VALUES (2, 'Phone', 1, 599.00);659INSERT INTO products VALUES (3, 'Python Book', 2, 49.00);660INSERT INTO products VALUES (4, 'SQL Handbook', 2, 39.00);661INSERT INTO orders VALUES (1, '2024-01-15');662INSERT INTO orders VALUES (2, '2024-01-20');663INSERT INTO orders VALUES (3, '2024-02-10');664INSERT INTO orders VALUES (4, '2024-02-28');665INSERT INTO orders VALUES (5, '2023-12-01');666INSERT INTO order_items VALUES (1, 1, 1, 1);667INSERT INTO order_items VALUES (2, 2, 3, 2);668INSERT INTO order_items VALUES (3, 2, 4, 1);669INSERT INTO order_items VALUES (4, 3, 2, 1);670INSERT INTO order_items VALUES (5, 4, 3, 3);671INSERT INTO order_items VALUES (6, 5, 1, 1);672""",673 "expected": [674 ("2024-01", "Electronics", 1, 999.00),675 ("2024-01", "Books", 1, 137.00),676 ("2024-02", "Electronics", 1, 599.00),677 ("2024-02", "Books", 1, 147.00),678 ],679 "max_steps": 7,680 },681 682 "data_anomalies": {683 "description": (684 "Find data quality issues in the customers table. "685 "Return: the type of issue as a string and the count of affected rows. "686 "The three issue types to check are:\n"687 " 1. 'duplicate_email' - email addresses that appear more than once\n"688 " 2. 'invalid_age' - age values that are NULL, negative, or greater than 150\n"689 " 3. 'null_name' - rows where name is NULL\n"690 "Return all three rows ordered alphabetically by issue type. "691 "Use UNION ALL to combine the three checks into one result set."692 ),693 "schema": (694 "CREATE TABLE customers (\n"695 " id INTEGER PRIMARY KEY,\n"696 " name TEXT,\n"697 " email TEXT,\n"698 " age INTEGER\n"699 ");"700 ),701 "seed_sql": """702INSERT INTO customers VALUES (1, 'Alice', 'a@test.com', 25);703INSERT INTO customers VALUES (2, NULL, 'b@test.com', 30);704INSERT INTO customers VALUES (3, 'Carol', 'a@test.com', 22);705INSERT INTO customers VALUES (4, 'Dave', 'd@test.com', -5);706INSERT INTO customers VALUES (5, 'Eve', 'e@test.com', 200);707""",708 "expected": [709 ("duplicate_email", 2),710 ("invalid_age", 2),711 ("null_name", 1),712 ],713 "max_steps": 7,714 },715 716 # ── NEW TASK ──────────────────────────────────────────────────────────────717 "window_functions": {718 "description": (719 "For each employee, calculate their salary rank within their department "720 "and the difference between their salary and their department's average salary. "721 "Return four columns in this exact order: "722 "employee name, department name, "723 "salary rank within the department (1 = highest paid, use RANK()), "724 "and the difference between their salary and the department average salary "725 "(rounded to 2 decimal places, positive means above average). "726 "Order results by department name ascending, then by rank ascending."727 ),728 "schema": (729 "CREATE TABLE departments (\n"730 " id INTEGER PRIMARY KEY,\n"731 " name TEXT NOT NULL\n"732 ");\n"733 "CREATE TABLE employees (\n"734 " id INTEGER PRIMARY KEY,\n"735 " name TEXT NOT NULL,\n"736 " department_id INTEGER NOT NULL,\n"737 " salary REAL NOT NULL\n"738 ");"739 ),740 "seed_sql": """741INSERT INTO departments VALUES (1, 'Engineering');742INSERT INTO departments VALUES (2, 'Marketing');743INSERT INTO employees VALUES (1, 'Alice', 1, 90000);744INSERT INTO employees VALUES (2, 'Bob', 1, 80000);745INSERT INTO employees VALUES (3, 'Carol', 1, 85000);746INSERT INTO employees VALUES (4, 'Dave', 2, 70000);747INSERT INTO employees VALUES (5, 'Eve', 2, 75000);748INSERT INTO employees VALUES (6, 'Frank', 2, 65000);749""",750 "expected": [751 ("Alice", "Engineering", 1, 5000.0),752 ("Carol", "Engineering", 2, 0.0),753 ("Bob", "Engineering", 3, -5000.0),754 ("Eve", "Marketing", 1, 5000.0),755 ("Dave", "Marketing", 2, 0.0),756 ("Frank", "Marketing", 3, -5000.0),757 ],758 "max_steps": 8,759 },760}761 762 763# ─────────────────────────────────────────────────────────────────────────────764# THREAD-BASED QUERY TIMEOUT (cross-platform, no SIGALRM)765# ─────────────────────────────────────────────────────────────────────────────766 767def _run_query_with_timeout(db_path: Path, sql: str, timeout_seconds: int = 5):768 """Execute SQL in a daemon thread with a hard timeout.769 Returns (rows, error_or_None).770 """771 result_box: list = [None]772 error_box: list = [None]773 774 def _target():775 conn = None776 try:777 conn = sqlite3.connect(str(db_path))778 cursor = conn.execute(sql)779 result_box[0] = cursor.fetchall()780 except Exception as exc:781 error_box[0] = exc782 finally:783 if conn:784 try:785 conn.close()786 except Exception:787 pass788 789 t = threading.Thread(target=_target, daemon=True)790 t.start()791 t.join(timeout=timeout_seconds)792 793 if t.is_alive():794 return None, TimeoutError(f"Query exceeded {timeout_seconds}s — avoid full table scans.")795 if error_box[0] is not None:796 return None, error_box[0]797 return result_box[0], None798 799 800# ─────────────────────────────────────────────────────────────────────────────801# ENVIRONMENT CLASS802# ─────────────────────────────────────────────────────────────────────────────803 804class SqlEnvironment(Environment):805 SUPPORTS_CONCURRENT_SESSIONS = True806 807 def __init__(self):808 self._episode_id: str | None = None809 self._state = State(episode_id=str(uuid4()), step_count=0)810 811 # ── File helpers (best-effort persistence) ────────────────────────────────812 813 def _db_path(self, episode_id: str) -> Path:814 return _SESSION_DIR / f"{episode_id}.sqlite3"815 816 def _meta_path(self, episode_id: str) -> Path:817 return _SESSION_DIR / f"session_{episode_id}.json"818 819 def _save_session_file(self, session: dict) -> None:820 try:821 path = self._meta_path(session["episode_id"])822 tmp = path.with_suffix(".tmp")823 with tmp.open("w", encoding="utf-8") as f:824 json.dump(session, f)825 tmp.replace(path)826 except Exception:827 pass # file persistence is best-effort; memory is primary828 829 def _load_session_file(self, episode_id: str) -> dict:830 try:831 path = self._meta_path(episode_id)832 if not path.exists():833 return {}834 with path.open("r", encoding="utf-8") as f:835 return json.load(f)836 except Exception:837 return {}838 839 # ── DB initialisation ─────────────────────────────────────────────────────840 841 def _initialise_db(self, task_name: str, db_path: Path) -> None:842 task = TASKS[task_name]843 if db_path.exists():844 try:845 db_path.unlink()846 except OSError:847 pass848 conn = sqlite3.connect(str(db_path))849 try:850 conn.executescript(task["schema"])851 conn.executescript(task["seed_sql"])852 conn.commit()853 finally:854 conn.close()855 856 # ── Session resolution ────────────────────────────────────────────────────857 858 def _resolve_session(self) -> dict:859 """Find the session for the current episode.860 861 Priority order:862 1. self._episode_id (WebSocket / persistent instance where reset was called)863 2. _MEMORY_SESSIONS["__latest__"] (HTTP, same process, new instance per request)864 3. Disk fallback for the resolved episode_id (container restart recovery)865 """866 episode_id = self._episode_id or _MEMORY_SESSIONS.get("__latest__")867 if not episode_id:868 return {}869 870 session = _MEMORY_SESSIONS.get(episode_id, {})871 if not session:872 session = self._load_session_file(episode_id)873 if session:874 _MEMORY_SESSIONS[episode_id] = session875 876 if session:877 self._episode_id = episode_id # pin for this request878 879 return session880 881 # ── reset() ──────────────────────────────────────────────────────────────882 883 def reset(self, seed=None, episode_id=None, **kwargs) -> SqlObservation:884 task_name = kwargs.get("task", "select_basics")885 if task_name not in TASKS:886 task_name = "select_basics"887 888 task = TASKS[task_name]889 new_id = episode_id or str(uuid4())890 db_path = self._db_path(new_id)891 892 self._initialise_db(task_name, db_path)893 894 session = {895 "episode_id": new_id,896 "task_name": task_name,897 "db_path": str(db_path),898 "step_count": 0,899 }900 901 # Primary store: module-level dict (survives across instances in same process)902 _MEMORY_SESSIONS[new_id] = session903 _MEMORY_SESSIONS["__latest__"] = new_id904 # Secondary store: disk (survives container restart)905 self._save_session_file(session)906 907 self._episode_id = new_id908 self._state = State(episode_id=new_id, step_count=0)909 910 return SqlObservation(911 task_description = task["description"],912 schema_info = task["schema"],913 query_result = [],914 error_message = "",915 feedback = "Episode started. Write a SQL query to solve the task above.",916 score_breakdown = {},917 attempts_remaining = task["max_steps"],918 done = False,919 reward = 0.0,920 )921 922 # ── step() ────────────────────────────────────────────────────────────────923 924 def step(self, action: SqlAction) -> SqlObservation:925 session = self._resolve_session()926 927 if not session:928 return SqlObservation(929 task_description = "",930 schema_info = "",931 query_result = [],932 error_message = "No active session. Call /reset first.",933 feedback = "No active session — call /reset before /step.",934 score_breakdown = {"execute": -0.05},935 attempts_remaining = 0,936 done = True,937 reward = -0.05,938 )939 940 task_name = session.get("task_name", "select_basics")941 if task_name not in TASKS:942 task_name = "select_basics"943 944 db_path = Path(session.get("db_path") or str(self._db_path(session["episode_id"])))945 946 # Re-seed if the DB file was lost (e.g. tmpfs wipe on restart)947 if not db_path.exists():948 self._initialise_db(task_name, db_path)949 950 # Advance step counter951 step_count = int(session.get("step_count", 0)) + 1952 session["step_count"] = step_count953 _MEMORY_SESSIONS[self._episode_id] = session954 self._save_session_file(session)955 956 self._state = State(957 episode_id = self._episode_id or str(uuid4()),958 step_count = step_count,959 )960 961 task = TASKS[task_name]962 attempts_remaining = task["max_steps"] - step_count963 964 rows, err = _run_query_with_timeout(db_path, action.sql_query, timeout_seconds=5)965 966 if isinstance(err, TimeoutError):967 reward = -0.10968 rows = []969 feedback = str(err)970 breakdown = {"execute": -0.10}971 error_msg = str(err)972 elif err is not None:973 reward = -0.05974 rows = []975 feedback = f"SQL Error: {err}. Fix your syntax and try again."976 breakdown = {"execute": -0.05}977 error_msg = str(err)978 else:979 reward, feedback, breakdown = self._grade(rows, task["expected"], action.sql_query)980 error_msg = ""981 982 done = reward >= 0.95 or attempts_remaining <= 0983 984 return SqlObservation(985 task_description = task["description"],986 schema_info = task["schema"],987 query_result = [list(r) for r in (rows or [])],988 error_message = error_msg,989 feedback = feedback,990 score_breakdown = breakdown,991 attempts_remaining = max(0, attempts_remaining),992 done = done,993 reward = float(max(-0.10, min(1.0, reward))),994 )995 996 # ── Grader ────────────────────────────────────────────────────────────────997 998 @staticmethod999 def _normalize_row(row: tuple) -> tuple:1000 """Normalize floats to 2 decimal places for robust set comparison.1001 1002 SQLite can return 405.4999999999 instead of 405.5 due to IEEE 7541003 floating-point arithmetic (SUM, AVG, ROUND operations).1004 Without this, F1 drops to 0 even when the answer is logically correct.1005 Strings and integers are returned unchanged.1006 1007 Examples:1008 (405.4999999999,) -> (405.5,)1009 ('Alice', 1) -> ('Alice', 1)1010 (5000.0,) -> (5000.0,)1011 """1012 def _norm(v):1013 if isinstance(v, float):1014 return round(v, 2)1015 return v1016 return tuple(_norm(v) for v in row)1017 1018 def _grade(self, result: list, expected: list, sql_query: str):1019 breakdown: dict = {"execute": 0.10}1020 1021 if not result:1022 return (1023 0.10,1024 "Query ran but returned 0 rows. Check your WHERE clause or JOIN conditions.",1025 breakdown,1026 )1027 1028 # Use _normalize_row so floating-point arithmetic differences don't cause1029 # false mismatches (e.g. 405.4999999999 vs 405.5)1030 result_set = set(self._normalize_row(tuple(r)) for r in result)1031 expected_set = set(self._normalize_row(tuple(e)) for e in expected)1032 1033 result_cols = len(result[0]) if result else 01034 expected_cols = len(expected[0]) if expected else 01035 col_score = (1036 0.20 if result_cols == expected_cols1037 else 0.20 * (min(result_cols, expected_cols) / max(result_cols, expected_cols, 1))1038 )1039 breakdown["columns"] = round(col_score, 3)1040 1041 row_score = 0.20 * min(1.0, len(result) / max(len(expected), 1))1042 breakdown["rows"] = round(row_score, 3)1043 1044 f1 = self._f1(result_set, expected_set)1045 val_score = 0.40 * f11046 breakdown["values"] = round(val_score, 3)1047 1048 uses_star = "select*" in sql_query.lower().replace(" ", "")1049 eff_score = 0.0 if uses_star else 0.101050 breakdown["efficiency"] = eff_score1051 1052 total = breakdown["execute"] + col_score + row_score + val_score + eff_score1053 pct = int(f1 * 100)1054 1055 if f1 >= 1.0 and col_score >= 0.20 and row_score >= 0.20:1056 feedback = (1057 "Perfect! Exact match."1058 if not uses_star1059 else "Correct result but avoid SELECT * — target only needed columns."1060 )1061 elif result_cols > expected_cols:1062 feedback = (1063 f"Too many columns ({result_cols} returned, {expected_cols} expected). "1064 "Remove extra columns from SELECT."1065 )1066 elif result_cols < expected_cols:1067 feedback = (1068 f"Too few columns ({result_cols} returned, {expected_cols} expected). "1069 "Add missing columns to SELECT."1070 )1071 elif len(result) > len(expected) * 1.5:1072 feedback = (1073 f"Too many rows ({len(result)} vs {len(expected)} expected). "1074 "Check your WHERE or HAVING — a filter may be missing."1075 )1076 elif len(result) < len(expected):1077 feedback = (1078 f"Too few rows ({len(result)} vs {len(expected)} expected). "1079 "Check your JOIN or WHERE — some matching rows are being excluded."1080 )1081 elif f1 >= 0.8:1082 feedback = f"Very close! {pct}% of values match. Check column ordering or data type casting."1083 elif f1 >= 0.5:1084 feedback = f"Partial match: {pct}% correct. Re-read the task and check your filters."1085 else:1086 feedback = f"Mostly incorrect ({pct}% match). Start from the schema and re-read the task."1087 1088 return round(total, 3), feedback, breakdown1089 1090 @staticmethod1091 def _f1(result_set: set, expected_set: set) -> float:1092 if not result_set and not expected_set:1093 return 1.01094 if not result_set or not expected_set:1095 return 0.01096 intersection = result_set & expected_set1097 precision = len(intersection) / len(result_set)1098 recall = len(intersection) / len(expected_set)1099 if precision + recall == 0:1100 return 0.01101 return 2 * precision * recall / (precision + recall)1102 1103 @property1104 def state(self) -> State:1105 return self._state