CoolFace
Apppublic

dkAmulet/sql-query-optimizer

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
test_env.py500 linesDownload Raw Back to root
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