Team Ai
Apppublic

vinayaknandi05/sql-optimization-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
test_environment.py153 linesDownload Raw Back to tests
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