Team Ai
Apppublic

Mahathi4554/sql-query-debugging

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
test_env.py110 linesDownload Raw Back to Tests
1"""2Test suite for SQL Query Debugging OpenEnv.3Run: python tests/test_env.py4"""5import sys6import os7sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))8 9from environment import SQLEnv, Action, compute_reward, execute_sql10from tasks import TASKS11 12env = SQLEnv()13passed = 014failed = 015 16def check(label, condition, detail=""):17    global passed, failed18    if condition:19        print(f"  PASS  {label}")20        passed += 121    else:22        print(f"  FAIL  {label}" + (f" — {detail}" if detail else ""))23        failed += 124 25print("\n=== Task definitions ===")26check("3 tasks registered", len(TASKS) == 3)27for tid, task in TASKS.items():28    check(f"  {tid} has expected_rows", len(task.expected_rows) > 0)29    check(f"  {tid} has solution_query", len(task.solution_query) > 0)30    db = task.setup_db()31    check(f"  {tid} DB setup OK", os.path.exists(db))32 33print("\n=== SQLite executor ===")34t1 = TASKS["task_syntax_fix"]35db = t1.setup_db()36r = execute_sql(db, "SELECT name FROM employees WHERE department='Engineering';")37check("Valid query executes", r["executed"])38check("Valid query gets rows", len(r["rows"]) == 4)39r2 = execute_sql(db, "SELCT name FORM employees;")40check("Broken query fails gracefully", not r2["executed"])41 42print("\n=== Reward engine ===")43t1_db = t1.setup_db()44r_bad = execute_sql(t1_db, "SELCT * FORM employees")45reward_bad = compute_reward(t1, "SELCT * FORM employees", r_bad, attempt=0)46check("Bad syntax → reward < 0.15", reward_bad.value < 0.15)47 48r_good = execute_sql(t1_db, t1.solution_query)49reward_good = compute_reward(t1, t1.solution_query, r_good, attempt=0)50check("Solution query → reward >= 0.7", reward_good.value >= 0.7, f"got {reward_good.value}")51 52print("\n=== Episode: task_syntax_fix ===")53obs = env.reset(task_id="task_syntax_fix")54check("reset() returns Observation", obs.task_id == "task_syntax_fix")55check("broken_query present", len(obs.broken_query) > 0)56check("attempt starts at 0", obs.attempt == 0)57 58result = env.step(Action(sql_query="SELCT name FORM employees;"))59check("bad step returns StepResult", result.reward.value < 0.2)60check("attempt increments", result.observation.attempt == 1)61check("not done on bad attempt", not result.done)62 63result2 = env.step(Action(sql_query=t1.solution_query))64check("correct solution → reward >= 0.7", result2.reward.value >= 0.7, f"got {result2.reward.value}")65check("done=True on correct solution", result2.done)66 67state = env.state()68check("state() returns EnvState", state.task_id == "task_syntax_fix")69check("state.done is True", state.done)70 71print("\n=== Episode: task_logic_bug ===")72t2 = TASKS["task_logic_bug"]73obs2 = env.reset(task_id="task_logic_bug")74check("reset for task2 works", obs2.task_id == "task_logic_bug")75res2 = env.step(Action(sql_query=t2.solution_query))76check("task2 solution → reward >= 0.7", res2.reward.value >= 0.7, f"got {res2.reward.value}")77 78print("\n=== Episode: task_optimization ===")79t3 = TASKS["task_optimization"]80obs3 = env.reset(task_id="task_optimization")81check("reset for task3 works", obs3.task_id == "task_optimization")82res3 = env.step(Action(sql_query=t3.solution_query))83check("task3 solution → reward >= 0.7", res3.reward.value >= 0.7, f"got {res3.reward.value}")84 85print("\n=== Grader (standalone) ===")86for tid, task in TASKS.items():87    score = env.grade(tid, task.solution_query)88    check(f"grader({tid}) solution >= 0.7", score >= 0.7, f"got {score}")89    broken_score = env.grade(tid, "SELECT 1;")90    check(f"grader({tid}) wrong query < 0.5", broken_score < 0.5, f"got {broken_score}")91 92print("\n=== Hint system ===")93env.reset(task_id="task_syntax_fix")94env.step(Action(sql_query="SELECT 1;"))95env.step(Action(sql_query="SELECT 2;"))96r_hint = env.step(Action(sql_query="SELECT 3;"))97check("Hint revealed after 2+ failed attempts", r_hint.observation.hint is not None)98 99print("\n=== Max attempts / done flag ===")100env.reset(task_id="task_syntax_fix")101for i in range(5):102    r = env.step(Action(sql_query="SELECT 1;"))103final_done = r.done104check("Episode ends after max_attempts", final_done)105 106print(f"\n{'='*40}")107print(f"Results: {passed} passed, {failed} failed")108print('='*40)109if failed > 0:110    sys.exit(1)