Team Ai
Apppublic

aananda/sql-agent-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
test_local.py266 linesDownload Raw Back to root
1"""2test_local.py — local pre-submission validation3 4Validates:5  1. openenv.yaml structure6  2. All task schemas load and seed data inserts cleanly7  3. All graders return scores in [0.0, 1.0]8  4. Perfect answers score 1.0 on every task9  5. SQL errors score 0.010  6. Partial answers score between 0 and 111  7. Environment class: reset / step / state lifecycle12  8. Hint injection after 3 failed steps13  9. Episode terminates correctly on submit and on max_steps14 15Run with:16  python test_local.py17"""18 19import sqlite320import sys21import traceback22 23# ── Colour helpers ─────────────────────────────────────────────────────────────24 25def ok(msg):  print(f"  \033[32m✓\033[0m {msg}")26def fail(msg):print(f"  \033[31m✗\033[0m {msg}"); _FAILURES.append(msg)27 28_FAILURES = []29 30# ──────────────────────────────────────────────────────────────────────────────31# 1. openenv.yaml32# ──────────────────────────────────────────────────────────────────────────────33 34print("\n── 1. openenv.yaml ──")35try:36    import yaml37    with open("openenv.yaml") as f:38        meta = yaml.safe_load(f)39    for field in ("name", "version", "tasks", "observation", "action", "reward"):40        assert field in meta, f"Missing field: {field}"41    assert len(meta["tasks"]) >= 3, "Need at least 3 tasks"42    for t in meta["tasks"]:43        assert "id" in t and "difficulty" in t44    ok("openenv.yaml is valid and has 3+ tasks")45except Exception as e:46    fail(f"openenv.yaml: {e}")47 48# ──────────────────────────────────────────────────────────────────────────────49# 2 + 3. Tasks load and graders are range-safe50# ──────────────────────────────────────────────────────────────────────────────51 52print("\n── 2+3. Task schemas + grader range safety ──")53try:54    from app.tasks import TASKS55    for tid, task in TASKS.items():56        conn = sqlite3.connect(":memory:")57        conn.executescript(task["schema_sql"])58        # empty result59        s, _ = task["grader"](conn, "SELECT 1 WHERE 1=0")60        assert 0.0 <= s <= 1.0, f"Out of range: {s}"61        # syntax error62        s, _ = task["grader"](conn, "INVALID SQL !!!")63        assert s == 0.0, f"Syntax error should give 0.0, got {s}"64        ok(f"{tid}: schema loads, grader range-safe")65except Exception as e:66    fail(f"Tasks: {e}\n{traceback.format_exc()}")67 68# ──────────────────────────────────────────────────────────────────────────────69# 4. Perfect answers score 1.070# ──────────────────────────────────────────────────────────────────────────────71 72print("\n── 4. Perfect answers → 1.0 ──")73 74PERFECT_ANSWERS = {75    "task_1_easy": """76        SELECT DISTINCT c.name, c.email77        FROM customers c78        JOIN orders o ON c.id = o.customer_id79    """,80    "task_2_medium": """81        WITH recent AS (82            SELECT u.plan, u.id AS uid, COUNT(e.id) AS evts83            FROM users u84            LEFT JOIN events e ON e.user_id = u.id AND e.ts >= '2024-01-01'85            GROUP BY u.plan, u.id86        )87        SELECT plan,88               SUM(evts) AS total_events,89               ROUND(CAST(SUM(evts) AS REAL) / COUNT(uid), 2) AS avg_events_per_user90        FROM recent91        GROUP BY plan92        ORDER BY plan93    """,94    "task_3_hard": """95        WITH running AS (96            SELECT a.holder, t.tx_date,97                   SUM(t.amount) OVER (98                       PARTITION BY t.account_id ORDER BY t.tx_date99                       ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW100                   ) AS running_balance, t.amount101            FROM transactions t102            JOIN accounts a ON a.id = t.account_id103        ),104        milestones AS (105            SELECT holder, MIN(tx_date) AS first_milestone_date106            FROM running WHERE running_balance > 5000 GROUP BY holder107        ),108        credits AS (109            SELECT a.holder, SUM(t.amount) AS total_credits110            FROM transactions t111            JOIN accounts a ON a.id = t.account_id112            WHERE t.amount > 0 GROUP BY a.holder113        )114        SELECT c.holder, m.first_milestone_date, c.total_credits115        FROM credits c LEFT JOIN milestones m ON c.holder = m.holder116        ORDER BY c.holder117    """,118}119 120try:121    from app.tasks import TASKS122    for tid, sql in PERFECT_ANSWERS.items():123        conn = sqlite3.connect(":memory:")124        conn.executescript(TASKS[tid]["schema_sql"])125        s, fb = TASKS[tid]["grader"](conn, sql)126        assert s == 1.0, f"Expected 1.0, got {s}. Feedback: {fb}"127        ok(f"{tid}: perfect answer → 1.0")128except Exception as e:129    fail(f"Perfect answers: {e}")130 131# ──────────────────────────────────────────────────────────────────────────────132# 5. Partial answers give 0 < score < 1133# ──────────────────────────────────────────────────────────────────────────────134 135print("\n── 5. Partial answers → (0, 1) ──")136try:137    from app.tasks import TASKS138    conn = sqlite3.connect(":memory:")139    conn.executescript(TASKS["task_1_easy"]["schema_sql"])140    # Returns all customers — includes 2 extras141    s, _ = TASKS["task_1_easy"]["grader"](conn, "SELECT name, email FROM customers")142    assert 0.0 < s < 1.0, f"Expected partial, got {s}"143    ok(f"task_1_easy: partial answer → {s:.3f}")144 145    conn2 = sqlite3.connect(":memory:")146    conn2.executescript(TASKS["task_2_medium"]["schema_sql"])147    # Correct totals but avg wrong (returns total instead of avg) → 0.5 partial148    s2, _ = TASKS["task_2_medium"]["grader"](conn2, """149        WITH recent AS (150            SELECT u.plan, u.id AS uid, COUNT(e.id) AS evts151            FROM users u152            LEFT JOIN events e ON e.user_id = u.id AND e.ts >= '2024-01-01'153            GROUP BY u.plan, u.id154        )155        SELECT plan,156               SUM(evts) AS total_events,157               CAST(SUM(evts) AS REAL) AS avg_events_per_user158        FROM recent GROUP BY plan ORDER BY plan159    """)160    assert 0.0 < s2 < 1.0, f"Expected partial (wrong avg), got {s2}"161    ok(f"task_2_medium: partial answer (correct total, wrong avg) → {s2:.3f}")162except Exception as e:163    fail(f"Partial answers: {e}")164 165# ──────────────────────────────────────────────────────────────────────────────166# 6+7. Environment lifecycle167# ──────────────────────────────────────────────────────────────────────────────168 169print("\n── 6+7. Environment lifecycle ──")170try:171    from app.environment import SQLEnvironment172    from app.models import SQLAction173 174    env = SQLEnvironment("task_1_easy")175    obs = env.reset()176    assert obs.task_id == "task_1_easy"177    assert obs.steps_taken == 0178    assert not obs.done179    ok("reset() returns clean observation")180 181    obs2, r, done, info = env.step(SQLAction(mode="sql", query="SELECT * FROM customers LIMIT 3"))182    assert obs2.steps_taken == 1183    assert obs2.last_result is not None184    assert obs2.last_result.row_count == 3185    assert r.score >= 0.0186    assert not done187    ok(f"step(sql) works, result has 3 rows, reward={r.score}")188 189    st = env.state()190    assert st.steps_taken == 1191    assert len(st.query_history) == 1192    ok("state() returns correct step count and history")193 194    obs3, r3, done3, info3 = env.step(SQLAction(195        mode="submit",196        query="SELECT DISTINCT c.name, c.email FROM customers c JOIN orders o ON c.id = o.customer_id"197    ))198    assert done3, "Episode should be done after submit"199    assert r3.score == 1.0, f"Expected 1.0, got {r3.score}"200    assert r3.is_final201    ok(f"submit() → score=1.0, done=True, is_final=True")202 203    # Step after done should be no-op204    obs4, r4, done4, _ = env.step(SQLAction(mode="sql", query="SELECT 1"))205    assert done4206    assert r4.score == 0.0207    ok("step() after done is a no-op")208except Exception as e:209    fail(f"Environment lifecycle: {e}\n{traceback.format_exc()}")210 211# ──────────────────────────────────────────────────────────────────────────────212# 8. Hint injection213# ──────────────────────────────────────────────────────────────────────────────214 215print("\n── 8. Hint injection after 3 failures ──")216try:217    from app.environment import SQLEnvironment218    from app.models import SQLAction219 220    env = SQLEnvironment("task_3_hard")221    env.reset()222    hint_seen = False223    for i in range(5):224        obs, r, done, _ = env.step(SQLAction(mode="sql", query="SELECT 1"))225        if obs.hint:226            hint_seen = True227            ok(f"Hint appeared at step {i+1}: '{obs.hint[:60]}…'")228            break229    assert hint_seen, "Hint never appeared after 5 wrong steps"230except Exception as e:231    fail(f"Hint injection: {e}")232 233# ──────────────────────────────────────────────────────────────────────────────234# 9. Max steps exhaustion235# ──────────────────────────────────────────────────────────────────────────────236 237print("\n── 9. Episode terminates at max_steps ──")238try:239    from app.environment import SQLEnvironment240    from app.models import SQLAction241 242    env = SQLEnvironment("task_1_easy")243    env.reset()244    done = False245    for i in range(env.max_steps + 2):246        _, _, done, _ = env.step(SQLAction(mode="sql", query="SELECT 1"))247        if done:248            ok(f"Episode terminated at step {i+1} (max={env.max_steps})")249            break250    assert done, "Episode should have terminated"251except Exception as e:252    fail(f"Max steps: {e}")253 254# ──────────────────────────────────────────────────────────────────────────────255# Summary256# ──────────────────────────────────────────────────────────────────────────────257 258print()259if _FAILURES:260    print(f"\033[31m✗ {len(_FAILURES)} check(s) FAILED:\033[0m")261    for f in _FAILURES:262        print(f"  - {f}")263    sys.exit(1)264else:265    print("\033[32m✓ All checks passed — submission is valid!\033[0m")266