Team Ai
Apppublic

Codexzzz/sql-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
sql_environment.py1105 linesDownload Raw Back to server
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