vinayaknandi05/sql-optimization-openenv
0
1"""2Tests for SQLOptimizationEnv — verifies all 3 tasks have working graders3and scores in the 0.0–1.0 range.4"""5import sys, os6sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))7 8import pytest9from env.environment import SQLOptimizationEnv, SQLAction10 11 12# ── Fixtures ──────────────────────────────────────────────────────────────────13 14@pytest.fixture15def env():16 return SQLOptimizationEnv()17 18 19# ── Task: easy ────────────────────────────────────────────────────────────────20 21def test_easy_reset(env):22 obs = env.reset(task_id="task_easy")23 assert obs.task_description24 assert obs.original_query == "SELECT * FROM employees;"25 assert obs.step == 026 assert not obs.done27 28 29def test_easy_bad_query_scores_low(env):30 env.reset(task_id="task_easy")31 obs, reward, done, info = env.step(SQLAction(query="SELECT * FROM employees;"))32 assert 0.0 <= info["score"] <= 1.033 assert info["score"] < 0.6, "SELECT * with no WHERE should score below 0.6"34 35 36def test_easy_good_query_scores_high(env):37 env.reset(task_id="task_easy")38 obs, reward, done, info = env.step(SQLAction(39 query="SELECT id, name, department, salary FROM employees WHERE department = 'Engineering' ORDER BY salary DESC;",40 message="Selected specific columns, filtered by department, added ORDER BY"41 ))42 assert 0.0 <= info["score"] <= 1.043 assert info["score"] >= 0.7, f"Optimized query should score >= 0.7, got {info['score']}"44 45 46# ── Task: medium ──────────────────────────────────────────────────────────────47 48def test_medium_reset(env):49 obs = env.reset(task_id="task_medium")50 assert "N+1" in obs.task_description or "JOIN" in obs.task_description or "project" in obs.task_description.lower()51 assert obs.step == 052 53 54def test_medium_bad_query_scores_low(env):55 env.reset(task_id="task_medium")56 obs, reward, done, info = env.step(SQLAction(57 query="SELECT name, department, salary, (SELECT COUNT(*) FROM project_assignments pa WHERE pa.employee_id = e.id) AS project_count FROM employees e;"58 ))59 assert 0.0 <= info["score"] <= 1.060 # Correlated subquery still returns correct columns, so some partial credit OK61 assert info["score"] < 0.8, "Unoptimized N+1 query should score below 0.8"62 63 64def test_medium_good_query_scores_high(env):65 env.reset(task_id="task_medium")66 obs, reward, done, info = env.step(SQLAction(67 query="""68 SELECT e.name, e.department, e.salary, COUNT(pa.project_id) AS project_count69 FROM employees e70 LEFT JOIN project_assignments pa ON e.id = pa.employee_id71 GROUP BY e.id, e.name, e.department, e.salary72 ORDER BY project_count DESC;73 """,74 message="Replaced correlated subquery with LEFT JOIN + GROUP BY"75 ))76 assert 0.0 <= info["score"] <= 1.077 assert info["score"] >= 0.75, f"Optimized JOIN query should score >= 0.75, got {info['score']}"78 79 80# ── Task: hard ────────────────────────────────────────────────────────────────81 82def test_hard_reset(env):83 obs = env.reset(task_id="task_hard")84 assert obs.task_description85 assert obs.step == 086 87 88def test_hard_broken_query_fails(env):89 env.reset(task_id="task_hard")90 obs, reward, done, info = env.step(SQLAction(91 query="SELECT name, department, salary, AVG(salary) FROM employees WHERE salary > AVG(salary);"92 ))93 assert 0.0 <= info["score"] <= 1.094 assert info["score"] < 0.5, "Broken aggregation query should score below 0.5"95 96 97def test_hard_good_query_scores_high(env):98 env.reset(task_id="task_hard")99 obs, reward, done, info = env.step(SQLAction(100 query="""101 SELECT e.name, e.department, e.salary, dept_avg.avg_salary AS dept_avg_salary102 FROM employees e103 JOIN (104 SELECT department, AVG(salary) AS avg_salary105 FROM employees106 GROUP BY department107 ) dept_avg ON e.department = dept_avg.department108 WHERE e.salary > dept_avg.avg_salary109 ORDER BY e.department, e.salary DESC;110 """,111 message="Used subquery to compute dept avg, filtered employees above avg, added ORDER BY"112 ))113 assert 0.0 <= info["score"] <= 1.0114 assert info["score"] >= 0.65, f"Correct aggregation query should score >= 0.65, got {info['score']}"115 116 117# ── General ───────────────────────────────────────────────────────────────────118 119def test_reward_in_range(env):120 env.reset(task_id="task_easy")121 _, reward, _, _ = env.step(SQLAction(query="SELECT id, name FROM employees WHERE department='HR' ORDER BY name;"))122 assert -1.0 <= reward <= 1.0123 124 125def test_done_on_max_steps(env):126 env.reset(task_id="task_easy")127 done = False128 for _ in range(10):129 _, _, done, _ = env.step(SQLAction(query="SELECT 1;"))130 if done:131 break132 assert done, "Episode should be done after MAX_STEPS"133 134 135def test_state_returns_dict(env):136 env.reset(task_id="task_medium")137 state = env.state()138 assert isinstance(state, dict)139 assert "task_id" in state140 assert "score" in state141 142 143def test_score_range_all_tasks(env):144 """All graders must produce scores in 0.0–1.0."""145 for task_id in ["task_easy", "task_medium", "task_hard"]:146 env.reset(task_id=task_id)147 _, _, _, info = env.step(SQLAction(query="SELECT 1;"))148 assert 0.0 <= info["score"] <= 1.0, f"{task_id} score out of range"149 150 151if __name__ == "__main__":152 pytest.main([__file__, "-v"])153 