dkAmulet/sql-query-optimizer
0
1#!/usr/bin/env python32"""3test_env.py — Unit + integration tests for SQL Query Optimizer Environment.4 5Run with:6 python test_env.py7 # or: python -m pytest test_env.py -v8 9Tests cover:10 - Database seeding (determinism, row counts)11 - All three graders (perfect answer, bad answer, partial answer, invalid SQL)12 - Environment lifecycle (reset, step, state, done flag, step budget)13 - Reward range guarantees [0.0, 1.0]14 - Determinism (same input → same score every time)15"""16from __future__ import annotations17 18import sys19import time20 21# ─────────────────────────────────────────────── compatibility shim ──────────22try:23 import pydantic # noqa: F40124 HAS_PYDANTIC = True25except ImportError:26 HAS_PYDANTIC = False27 print("⚠ pydantic not installed — running grader-only tests via stdlib.\n")28 29# ─────────────────────────────────────────────────────────────────────────────30 31PASS = "✓"32FAIL = "✗"33results: list[tuple[str, bool, str]] = []34 35def check(name: str, condition: bool, detail: str = "") -> None:36 mark = PASS if condition else FAIL37 results.append((name, condition, detail))38 status = f" {mark} {name}"39 if detail:40 status += f" ({detail})"41 print(status)42 if not condition:43 # Don't raise immediately so all tests run44 pass45 46 47def section(title: str) -> None:48 print(f"\n{'─' * 60}")49 print(f" {title}")50 print("─" * 60)51 52 53# ══════════════════════════════════════════════════════════════════════════════54# STDLIB-ONLY TESTS (run regardless of pydantic)55# ══════════════════════════════════════════════════════════════════════════════56 57import sqlite3, random, re58 59def _make_db():60 """Inline version of create_database() for stdlib tests."""61 conn = sqlite3.connect(":memory:", check_same_thread=False)62 conn.executescript("""63 CREATE TABLE users (user_id INTEGER PRIMARY KEY, username TEXT NOT NULL,64 email TEXT NOT NULL, first_name TEXT, last_name TEXT, city TEXT,65 country TEXT, is_active INTEGER DEFAULT 1, created_at TEXT);66 CREATE TABLE categories (category_id INTEGER PRIMARY KEY, name TEXT NOT NULL,67 parent_category_id INTEGER);68 CREATE TABLE products (product_id INTEGER PRIMARY KEY, name TEXT NOT NULL,69 category_id INTEGER, price REAL, sku TEXT, is_available INTEGER DEFAULT 1);70 CREATE TABLE orders (order_id INTEGER PRIMARY KEY, user_id INTEGER,71 status TEXT, total_amount REAL, created_at TEXT);72 CREATE TABLE order_items (item_id INTEGER PRIMARY KEY, order_id INTEGER,73 product_id INTEGER, quantity INTEGER, unit_price REAL);74 CREATE INDEX idx_users_active ON users(is_active);75 CREATE INDEX idx_users_country ON users(country, is_active);76 CREATE INDEX idx_orders_user ON orders(user_id);77 CREATE INDEX idx_orders_status ON orders(status);78 CREATE INDEX idx_items_order ON order_items(order_id);79 CREATE INDEX idx_items_product ON order_items(product_id);80 CREATE INDEX idx_products_cat ON products(category_id);81 """)82 rng = random.Random(42)83 cats = [(1,"Electronics",None),(2,"Clothing",None),(3,"Books",None),(4,"Smartphones",1),84 (5,"Laptops",1),(6,"Tablets",1),(7,"T-Shirts",2),(8,"Jeans",2),(9,"Fiction",3),(10,"Non-Fiction",3)]85 conn.executemany("INSERT INTO categories VALUES (?,?,?)", cats)86 countries = ["USA","USA","USA","UK","Canada","Germany","France","Australia"]87 cities = ["New York","Los Angeles","Chicago","London","Toronto","Berlin","Paris","Sydney"]88 users = [(i,f"user_{i}",f"user{i}@example.com",f"First{i}",f"Last{i}",89 rng.choice(cities),rng.choice(countries),1 if rng.random()>0.25 else 0,"2022-01-01")90 for i in range(1,1501)]91 conn.executemany("INSERT INTO users VALUES (?,?,?,?,?,?,?,?,?)", users)92 leaf = [4,5,6,7,8,9,10]93 products = [(i,f"Product_{i:04d}",rng.choice(leaf),round(rng.uniform(9.99,1499.99),2),f"SKU-{i:05d}",1)94 for i in range(1,301)]95 conn.executemany("INSERT INTO products VALUES (?,?,?,?,?,?)", products)96 statuses=["pending","processing","shipped","delivered","cancelled"]; w=[.10,.10,.20,.50,.10]97 orders = [(i,rng.randint(1,1500),rng.choices(statuses,weights=w)[0],round(rng.uniform(20,3000),2),"2024-01-01")98 for i in range(1,3001)]99 conn.executemany("INSERT INTO orders VALUES (?,?,?,?,?)", orders)100 items = [(i,rng.randint(1,3000),rng.randint(1,300),rng.randint(1,5),round(rng.uniform(9.99,1499.99),2))101 for i in range(1,8001)]102 conn.executemany("INSERT INTO order_items VALUES (?,?,?,?,?)", items)103 conn.commit()104 return conn105 106 107def test_database():108 section("Database Seeding")109 conn = _make_db()110 111 n = conn.execute("SELECT COUNT(*) FROM users").fetchone()[0]112 check("users table has 1500 rows", n == 1500, f"got {n}")113 114 n = conn.execute("SELECT COUNT(*) FROM orders").fetchone()[0]115 check("orders table has 3000 rows", n == 3000, f"got {n}")116 117 n = conn.execute("SELECT COUNT(*) FROM order_items").fetchone()[0]118 check("order_items table has 8000 rows", n == 8000, f"got {n}")119 120 n = conn.execute("SELECT COUNT(*) FROM categories").fetchone()[0]121 check("categories table has 10 rows", n == 10, f"got {n}")122 123 n = conn.execute("SELECT COUNT(*) FROM products").fetchone()[0]124 check("products table has 300 rows", n == 300, f"got {n}")125 126 # Determinism: run twice and compare active user counts127 conn2 = _make_db()128 a1 = conn.execute("SELECT COUNT(*) FROM users WHERE is_active=1").fetchone()[0]129 a2 = conn2.execute("SELECT COUNT(*) FROM users WHERE is_active=1").fetchone()[0]130 check("database is deterministic (seed=42)", a1 == a2, f"both={a1}")131 132 # Indexes exist133 indexes = [r[1] for r in conn.execute("PRAGMA index_list(users)").fetchall()]134 check("idx_users_active index exists", "idx_users_active" in indexes)135 check("idx_users_country index exists", "idx_users_country" in indexes)136 137 conn.close(); conn2.close()138 139 140def _norm(q): return re.sub(r"\s+"," ",q.strip().upper())141 142def test_grader_helpers():143 section("Grader Helper Functions (regexp)")144 145 check("SELECT * detected", bool(re.search(r"SELECT\s+\*", _norm("SELECT * FROM users"))))146 check("SELECT col not flagged", not bool(re.search(r"SELECT\s+\*", _norm("SELECT id FROM users"))))147 148 check("IN (SELECT …) detected",149 bool(re.search(r"\bIN\s*\(\s*SELECT\b", _norm("WHERE id IN (SELECT id FROM t)"))))150 check("IN (1,2,3) not flagged",151 not bool(re.search(r"\bIN\s*\(\s*SELECT\b", _norm("WHERE id IN (1,2,3)"))))152 153 check("JOIN detected", bool(re.search(r"\bJOIN\b", _norm("FROM a JOIN b ON a.id=b.id"))))154 check("No JOIN detected", not bool(re.search(r"\bJOIN\b", _norm("FROM users WHERE id=1"))))155 156 check("SELECT (SELECT …) detected",157 bool(re.search(r"SELECT\s+\(\s*SELECT\b", _norm("SELECT (SELECT name FROM c WHERE c.id=p.id)"))))158 check("Normal SELECT not flagged",159 not bool(re.search(r"SELECT\s+\(\s*SELECT\b", _norm("SELECT name FROM cats"))))160 161 162def test_result_correctness():163 section("Result Correctness (SQL logic)")164 conn = _make_db()165 166 # Task 1: active user IDs167 ref = {r[0] for r in conn.execute("SELECT user_id FROM users WHERE is_active=1").fetchall()}168 opt = {r[0] for r in conn.execute(169 "SELECT user_id FROM users WHERE is_active=1").fetchall()}170 check("Task1: active user set reproducible", ref == opt, f"{len(ref)} users")171 172 # Task 2: delivered orders from USA active users173 q_slow = """SELECT order_id FROM orders WHERE user_id IN174 (SELECT user_id FROM users WHERE country='USA' AND is_active=1)175 AND status='delivered'"""176 q_join = """SELECT o.order_id FROM orders o JOIN users u ON u.user_id=o.user_id177 WHERE u.country='USA' AND u.is_active=1 AND o.status='delivered'"""178 ref2 = {r[0] for r in conn.execute(q_slow).fetchall()}179 opt2 = {r[0] for r in conn.execute(q_join).fetchall()}180 check("Task2: IN and JOIN produce identical order sets", ref2 == opt2, f"{len(ref2)} orders")181 182 # Task 3: category revenue183 q3 = """SELECT c.name, SUM(oi.quantity*oi.unit_price)184 FROM categories c JOIN products p ON p.category_id=c.category_id185 JOIN order_items oi ON oi.product_id=p.product_id186 GROUP BY c.category_id, c.name HAVING SUM(oi.quantity*oi.unit_price)>1000187 ORDER BY 2 DESC"""188 rows3 = conn.execute(q3).fetchall()189 check("Task3: categories with revenue > 1000 exist", len(rows3) > 0, f"{len(rows3)} categories")190 check("Task3: all revenues > 1000", all(r[1] > 1000 for r in rows3))191 192 conn.close()193 194 195# ══════════════════════════════════════════════════════════════════════════════196# PYDANTIC-DEPENDENT TESTS197# ══════════════════════════════════════════════════════════════════════════════198 199def test_env_lifecycle():200 section("Environment Lifecycle (pydantic required)")201 from env import SQLQueryOptimizerEnv202 from models import SQLAction203 204 env = SQLQueryOptimizerEnv()205 206 # reset() returns correct structure207 obs = env.reset("select_star_removal")208 check("reset() returns task_id", obs.task_id == "select_star_removal")209 check("reset() returns difficulty=easy", obs.difficulty == "easy")210 check("reset() step_number=0", obs.step_number == 0)211 check("reset() schema_ddl non-empty", len(obs.schema_ddl) > 100)212 check("reset() slow_query non-empty", len(obs.slow_query) > 10)213 check("reset() last_reward=0.0 on fresh episode", obs.last_reward == 0.0)214 215 # state() reflects reset216 st = env.state()217 check("state() task_id matches", st.task_id == "select_star_removal")218 check("state() step_number=0 after reset", st.step_number == 0)219 check("state() best_reward=0.0 after reset", st.best_reward == 0.0)220 check("state() done=False after reset", not st.done)221 222 # step() increments step counter223 result = env.step(SQLAction(optimized_query="SELECT user_id, username, email FROM users WHERE is_active=1"))224 check("step() increments step number", result.observation.step_number == 1)225 check("step() reward.value in [0,1]", 0.0 <= result.reward.value <= 1.0)226 check("step() returns info dict with best_reward", "best_reward" in result.info)227 228 # Episode ends after max_steps229 obs2 = env.reset("select_star_removal") # max_steps=5230 for _ in range(4):231 env.step(SQLAction(optimized_query="SELECT * FROM users WHERE is_active=1"))232 r_last = env.step(SQLAction(optimized_query="SELECT * FROM users WHERE is_active=1"))233 check("Episode done after max_steps exhausted", r_last.done)234 235 # RuntimeError when stepping into done episode236 try:237 env.step(SQLAction(optimized_query="SELECT 1"))238 check("RuntimeError raised after done episode", False)239 except RuntimeError:240 check("RuntimeError raised after done episode", True)241 242 # list_tasks returns all three243 tasks = env.list_tasks()244 check("list_tasks() returns 3 tasks", len(tasks) == 3)245 ids = [t["task_id"] for t in tasks]246 check("all task IDs present", "select_star_removal" in ids and "aggregation_optimization" in ids)247 248 env.close()249 250 251def test_grader_task1():252 section("Grader: Task 1 — SELECT * Elimination")253 from env import SQLQueryOptimizerEnv254 from models import SQLAction255 256 env = SQLQueryOptimizerEnv()257 258 # Perfect answer259 env.reset("select_star_removal")260 r = env.step(SQLAction(optimized_query="SELECT user_id, username, email FROM users WHERE is_active = 1"))261 check("Perfect T1: reward >= 0.90", r.reward.value >= 0.90, f"got {r.reward.value:.3f}")262 check("Perfect T1: style=0.30 (no SELECT *)", r.reward.breakdown.style == 0.30)263 check("Perfect T1: validity > 0", r.reward.breakdown.validity > 0)264 check("Perfect T1: correctness > 0", r.reward.breakdown.correctness > 0)265 266 # SELECT * still present267 env.reset("select_star_removal")268 r_bad = env.step(SQLAction(optimized_query="SELECT * FROM users WHERE is_active = 1"))269 check("SELECT * T1: style=0.0", r_bad.reward.breakdown.style == 0.0)270 check("SELECT * T1: reward < perfect", r_bad.reward.value < r.reward.value)271 272 # Invalid SQL273 env.reset("select_star_removal")274 r_inv = env.step(SQLAction(optimized_query="NOT VALID SQL AT ALL !!!"))275 check("Invalid SQL T1: reward=0.0", r_inv.reward.value == 0.0)276 check("Invalid SQL T1: validity=0.0", r_inv.reward.breakdown.validity == 0.0)277 278 # Wrong result set (missing WHERE clause)279 env.reset("select_star_removal")280 r_wrong = env.step(SQLAction(optimized_query="SELECT user_id, username, email FROM users"))281 check("Wrong result T1: correctness < perfect", r_wrong.reward.breakdown.correctness < 0.40)282 283 env.close()284 285 286def test_grader_task2():287 section("Grader: Task 2 — Correlated Subquery → JOIN")288 from env import SQLQueryOptimizerEnv289 from models import SQLAction290 291 env = SQLQueryOptimizerEnv()292 293 perfect_q2 = """294 SELECT o.order_id, o.user_id, o.total_amount295 FROM orders o296 JOIN users u ON u.user_id = o.user_id297 WHERE u.country = 'USA'298 AND u.is_active = 1299 AND o.status = 'delivered'300 """301 302 env.reset("subquery_to_join")303 r = env.step(SQLAction(optimized_query=perfect_q2))304 check("Perfect T2: reward >= 0.80", r.reward.value >= 0.80, f"got {r.reward.value:.3f}")305 check("Perfect T2: style=0.20 (no IN subquery)", r.reward.breakdown.style == 0.20)306 check("Perfect T2: correctness = 0.40", r.reward.breakdown.correctness == 0.40)307 308 # Still uses IN (SELECT …)309 env.reset("subquery_to_join")310 r_in = env.step(SQLAction(optimized_query=(311 "SELECT order_id, user_id, total_amount FROM orders "312 "WHERE user_id IN (SELECT user_id FROM users WHERE country='USA' AND is_active=1) "313 "AND status='delivered'"314 )))315 check("IN subquery T2: style=0.0", r_in.reward.breakdown.style == 0.0)316 check("IN subquery T2: reward < JOIN reward", r_in.reward.value < r.reward.value)317 318 # Invalid SQL319 env.reset("subquery_to_join")320 r_inv = env.step(SQLAction(optimized_query="BROKEN QUERY"))321 check("Invalid SQL T2: reward=0.0", r_inv.reward.value == 0.0)322 323 env.close()324 325 326def test_grader_task3():327 section("Grader: Task 3 — Aggregation Optimization")328 from env import SQLQueryOptimizerEnv329 from models import SQLAction330 331 env = SQLQueryOptimizerEnv()332 333 perfect_q3 = """334 SELECT c.name AS category_name,335 SUM(oi.quantity * oi.unit_price) AS total_revenue336 FROM categories c337 JOIN products p ON p.category_id = c.category_id338 JOIN order_items oi ON oi.product_id = p.product_id339 GROUP BY c.category_id, c.name340 HAVING SUM(oi.quantity * oi.unit_price) > 1000341 ORDER BY total_revenue DESC342 """343 344 env.reset("aggregation_optimization")345 r = env.step(SQLAction(optimized_query=perfect_q3))346 check("Perfect T3: reward >= 0.90", r.reward.value >= 0.90, f"got {r.reward.value:.3f}")347 check("Perfect T3: style=0.25 (no correlated subqueries)", r.reward.breakdown.style == 0.25)348 check("Perfect T3: correctness = 0.35", r.reward.breakdown.correctness == 0.35)349 check("Perfect T3: performance > 0", r.reward.breakdown.performance > 0)350 351 # Correlated subquery still present352 env.reset("aggregation_optimization")353 r_sub = env.step(SQLAction(optimized_query=(354 "SELECT (SELECT name FROM categories WHERE category_id=p.category_id) AS category_name, "355 "SUM(oi.quantity*oi.unit_price) AS total_revenue "356 "FROM products p JOIN order_items oi ON oi.product_id=p.product_id "357 "GROUP BY p.category_id HAVING SUM(oi.quantity*oi.unit_price)>1000 ORDER BY total_revenue DESC"358 )))359 check("Correlated T3: style < 0.25", r_sub.reward.breakdown.style < 0.25)360 check("Correlated T3: reward < perfect", r_sub.reward.value < r.reward.value)361 362 env.close()363 364 365def test_reward_range():366 section("Reward Range Guarantee [0.0, 1.0]")367 from env import SQLQueryOptimizerEnv368 from models import SQLAction369 from tasks import TASK_ORDER370 371 env = SQLQueryOptimizerEnv()372 queries = [373 "SELECT * FROM users",374 "SELECT user_id FROM users WHERE is_active=1",375 "COMPLETELY INVALID",376 "SELECT 1",377 "SELECT o.order_id, o.user_id, o.total_amount FROM orders o JOIN users u ON u.user_id=o.user_id WHERE u.country='USA' AND u.is_active=1 AND o.status='delivered'",378 "SELECT c.name AS category_name, SUM(oi.quantity*oi.unit_price) AS total_revenue FROM categories c JOIN products p ON p.category_id=c.category_id JOIN order_items oi ON oi.product_id=p.product_id GROUP BY c.category_id,c.name HAVING SUM(oi.quantity*oi.unit_price)>1000 ORDER BY total_revenue DESC",379 ]380 381 all_in_range = True382 for task_id in TASK_ORDER:383 env.reset(task_id)384 for q in queries:385 r = env.step(SQLAction(optimized_query=q))386 v = r.reward.value387 if not (0.0 <= v <= 1.0):388 all_in_range = False389 print(f" OUT OF RANGE: task={task_id} query={q[:40]} reward={v}")390 if r.done:391 env.reset(task_id)392 393 check("All rewards in [0.0, 1.0] across all tasks and queries", all_in_range)394 env.close()395 396 397def test_determinism():398 section("Determinism (same input → same output)")399 from env import SQLQueryOptimizerEnv400 from models import SQLAction401 402 q = "SELECT user_id, username, email FROM users WHERE is_active = 1"403 404 env1 = SQLQueryOptimizerEnv()405 env2 = SQLQueryOptimizerEnv()406 407 env1.reset("select_star_removal")408 env2.reset("select_star_removal")409 410 r1 = env1.step(SQLAction(optimized_query=q))411 r2 = env2.step(SQLAction(optimized_query=q))412 413 check("Same query → same reward across two env instances",414 r1.reward.value == r2.reward.value,415 f"{r1.reward.value} vs {r2.reward.value}")416 417 env1.close(); env2.close()418 419 420def test_episode_reset_cleans_state():421 section("Episode Reset Cleans State")422 from env import SQLQueryOptimizerEnv423 from models import SQLAction424 425 env = SQLQueryOptimizerEnv()426 427 # Run task1 to completion428 env.reset("select_star_removal")429 for _ in range(5):430 env.step(SQLAction(optimized_query="SELECT * FROM users WHERE is_active=1"))431 432 # Reset to task2 — state should be clean433 obs = env.reset("subquery_to_join")434 check("After reset: task_id switches", obs.task_id == "subquery_to_join")435 check("After reset: step_number=0", obs.step_number == 0)436 st = env.state()437 check("After reset: best_reward=0.0", st.best_reward == 0.0)438 check("After reset: done=False", not st.done)439 440 # Can step into the new episode normally441 r = env.step(SQLAction(optimized_query="SELECT order_id, user_id, total_amount FROM orders WHERE status='delivered'"))442 check("Can step after reset", r.reward.value >= 0.0)443 444 env.close()445 446 447# ══════════════════════════════════════════════════════════════════════════════448# RUNNER449# ══════════════════════════════════════════════════════════════════════════════450 451def main():452 print("╔══════════════════════════════════════════════════════════╗")453 print("║ SQL Query Optimizer — Test Suite ║")454 print("╚══════════════════════════════════════════════════════════╝")455 456 t0 = time.time()457 458 # Always run (no pydantic needed)459 test_database()460 test_grader_helpers()461 test_result_correctness()462 463 # Run pydantic-dependent tests only if available464 if HAS_PYDANTIC:465 import sys as _sys466 _sys.path.insert(0, ".")467 test_env_lifecycle()468 test_grader_task1()469 test_grader_task2()470 test_grader_task3()471 test_reward_range()472 test_determinism()473 test_episode_reset_cleans_state()474 else:475 print("\n ⚠ Skipping pydantic tests (install pydantic + dependencies)")476 477 elapsed = time.time() - t0478 479 # ── Summary ──────────────────────────────────────────────────────────────480 print(f"\n{'═' * 60}")481 passed = sum(1 for _, ok, _ in results if ok)482 failed = sum(1 for _, ok, _ in results if not ok)483 total = len(results)484 print(f" Results: {passed}/{total} passed ({failed} failed) [{elapsed:.2f}s]")485 486 if failed:487 print("\n FAILED TESTS:")488 for name, ok, detail in results:489 if not ok:490 print(f" {FAIL} {name} {detail}")491 print("═" * 60)492 sys.exit(1)493 else:494 print(" ALL TESTS PASSED ✓")495 print("═" * 60)496 497 498if __name__ == "__main__":499 main()500 