aananda/sql-agent-env
0
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 