Tsah00/sql-env
0
1"""2test_env.py - Comprehensive test suite for the SQL Query Learning Environment.3 4Tests:5 1. Database seeding and schema6 2. All 9 tasks (3 easy, 3 medium, 3 hard) with reference solutions7 3. Grader correctness and reward ranges8 4. Environment lifecycle (reset, step, state)9 5. Partial credit scoring10 6. Error handling (bad SQL, wrong difficulty, empty queries)11 7. Multi-difficulty sweep12 8. Inference script dry-run13 9. FastAPI HTTP endpoints (integration)14 10. WebSocket endpoint15"""16 17from __future__ import annotations18 19import json20import sqlite321import sys22import os23import time24import threading25 26# Make root importable27sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))28 29import pytest30 31from models import SQLAction, SQLObservation, SQLState32from server.tasks import (33 TASKS, SCHEMA_INFO, seed_database, grade,34 get_all_tasks_by_difficulty,35)36from server.sql_environment import SQLEnvironment37 38 39# ===========================================================================40# Fixtures41# ===========================================================================42 43@pytest.fixture44def db():45 """In-memory SQLite database seeded with test data."""46 conn = sqlite3.connect(":memory:")47 conn.row_factory = sqlite3.Row48 seed_database(conn)49 yield conn50 conn.close()51 52 53@pytest.fixture54def env():55 """Fresh SQLEnvironment instance."""56 e = SQLEnvironment()57 yield e58 e.close()59 60 61# ===========================================================================62# 1. Database / Schema Tests63# ===========================================================================64 65class TestDatabase:66 67 def test_tables_exist(self, db):68 tables = {row[0] for row in db.execute(69 "SELECT name FROM sqlite_master WHERE type='table'"70 ).fetchall()}71 assert {"customers", "products", "orders", "order_items"} <= tables72 73 def test_customers_seeded(self, db):74 count = db.execute("SELECT COUNT(*) FROM customers").fetchone()[0]75 assert count >= 1076 77 def test_products_seeded(self, db):78 count = db.execute("SELECT COUNT(*) FROM products").fetchone()[0]79 assert count >= 1080 81 def test_orders_seeded(self, db):82 count = db.execute("SELECT COUNT(*) FROM orders").fetchone()[0]83 assert count >= 2084 85 def test_order_items_seeded(self, db):86 count = db.execute("SELECT COUNT(*) FROM order_items").fetchone()[0]87 assert count >= 3088 89 def test_categories_exist(self, db):90 cats = {row[0] for row in db.execute("SELECT DISTINCT category FROM products").fetchall()}91 assert len(cats) >= 392 93 def test_order_statuses(self, db):94 statuses = {row[0] for row in db.execute("SELECT DISTINCT status FROM orders").fetchall()}95 assert "completed" in statuses96 97 def test_schema_info_nonempty(self):98 assert len(SCHEMA_INFO) > 10099 assert "customers" in SCHEMA_INFO100 assert "products" in SCHEMA_INFO101 assert "orders" in SCHEMA_INFO102 103 104# ===========================================================================105# 2. Task Registry Tests106# ===========================================================================107 108class TestTaskRegistry:109 110 def test_nine_tasks_total(self):111 assert len(TASKS) == 9112 113 def test_three_per_difficulty(self):114 for diff in ("easy", "medium", "hard"):115 tasks = get_all_tasks_by_difficulty(diff)116 assert len(tasks) == 3, f"Expected 3 {diff} tasks, got {len(tasks)}"117 118 def test_all_tasks_have_required_fields(self):119 required = {"id", "difficulty", "description", "reference_sql",120 "required_keywords", "expected_columns"}121 for tid, task in TASKS.items():122 for field in required:123 assert field in task, f"Task {tid} missing field: {field}"124 125 def test_descriptions_nonempty(self):126 for tid, task in TASKS.items():127 assert len(task["description"]) > 20, f"Task {tid} has short description"128 129 def test_reference_sql_nonempty(self):130 for tid, task in TASKS.items():131 assert len(task["reference_sql"].strip()) > 10, f"Task {tid} missing SQL"132 133 134# ===========================================================================135# 3. Grader Tests - Reference Solutions Score 1.0136# ===========================================================================137 138class TestGraderReferenceSolutions:139 """The reference solution for every task must score >= 0.9."""140 141 def _grade_reference(self, task_id: str, db) -> float:142 task = TASKS[task_id]143 reward, msg, agent_rows, expected_rows = grade(task_id, task["reference_sql"], db)144 return reward145 146 def test_easy_1_reference(self, db):147 assert self._grade_reference("easy_1", db) >= 0.9148 149 def test_easy_2_reference(self, db):150 assert self._grade_reference("easy_2", db) >= 0.9151 152 def test_easy_3_reference(self, db):153 assert self._grade_reference("easy_3", db) >= 0.9154 155 def test_medium_1_reference(self, db):156 assert self._grade_reference("medium_1", db) >= 0.9157 158 def test_medium_2_reference(self, db):159 assert self._grade_reference("medium_2", db) >= 0.9160 161 def test_medium_3_reference(self, db):162 assert self._grade_reference("medium_3", db) >= 0.9163 164 def test_hard_1_reference(self, db):165 assert self._grade_reference("hard_1", db) >= 0.9166 167 def test_hard_2_reference(self, db):168 assert self._grade_reference("hard_2", db) >= 0.9169 170 def test_hard_3_reference(self, db):171 assert self._grade_reference("hard_3", db) >= 0.9172 173 174# ===========================================================================175# 4. Grader Tests - Incorrect Queries Score Low176# ===========================================================================177 178class TestGraderIncorrectQueries:179 180 def test_wrong_table(self, db):181 reward, msg, _, _ = grade("easy_1", "SELECT name FROM products", db)182 assert reward < 0.5183 184 def test_syntax_error(self, db):185 reward, msg, _, _ = grade("easy_1", "SELEKT * FRUM customers", db)186 assert reward <= 0.01 # minimum reward — never exactly 0.0187 assert "error" in msg.lower() or "query error" in msg.lower()188 189 def test_empty_query(self, db):190 reward, msg, _, _ = grade("easy_2", "", db)191 assert reward <= 0.01 # minimum reward — never exactly 0.0192 193 def test_wrong_filter(self, db):194 # Should be USA, returns UK instead - partial credit at most195 reward, msg, _, _ = grade(196 "easy_1",197 "SELECT name, email FROM customers WHERE country = 'UK'",198 db199 )200 assert reward < 0.8201 202 def test_reward_range(self, db):203 for task_id, task in TASKS.items():204 reward, _, _, _ = grade(task_id, task["reference_sql"], db)205 assert 0.0 < reward < 1.0, f"Reward out of open interval (0,1) for {task_id}: {reward}"206 207 def test_unknown_task_returns_zero(self, db):208 reward, msg, _, _ = grade("nonexistent_task", "SELECT 1", db)209 assert reward <= 0.01 # minimum reward — never exactly 0.0210 assert "unknown" in msg.lower()211 212 213# ===========================================================================214# 5. Partial Credit Tests215# ===========================================================================216 217class TestPartialCredit:218 219 def test_partial_result_gets_partial_credit(self, db):220 # Returns only some USA customers221 reward_full, _, _, _ = grade(222 "easy_1",223 "SELECT name, email FROM customers WHERE country = 'USA'",224 db225 )226 reward_partial, _, _, _ = grade(227 "easy_1",228 "SELECT name, email FROM customers WHERE country = 'USA' LIMIT 1",229 db230 )231 assert reward_partial < reward_full232 assert reward_partial > 0.0233 234 def test_extra_columns_dont_break_grader(self, db):235 # Adding extra columns - result still matches on name/email fields236 reward, msg, _, _ = grade(237 "easy_1",238 "SELECT name, email, city FROM customers WHERE country = 'USA'",239 db240 )241 # Should get some reward (has name+email in result, plus extra city)242 assert reward >= 0.0 # grader is lenient243 244 245# ===========================================================================246# 6. Environment Lifecycle Tests247# ===========================================================================248 249class TestEnvironmentLifecycle:250 251 def test_reset_returns_observation(self, env):252 obs = env.reset()253 assert isinstance(obs, SQLObservation)254 assert obs.task_description != ""255 assert obs.schema_info != ""256 assert obs.reward == 0.0 # reset always gives reward=0.0 (not graded)257 assert obs.done is False258 259 def test_reset_difficulty_easy(self, env):260 obs = env.reset(difficulty="easy")261 assert env.state.current_difficulty == "easy"262 263 def test_reset_difficulty_medium(self, env):264 obs = env.reset(difficulty="medium")265 assert env.state.current_difficulty == "medium"266 267 def test_reset_difficulty_hard(self, env):268 obs = env.reset(difficulty="hard")269 assert env.state.current_difficulty == "hard"270 271 def test_state_after_reset(self, env):272 env.reset()273 state = env.state274 assert isinstance(state, SQLState)275 assert state.episode_id != ""276 assert state.step_count == 0277 assert state.total_reward == 0.0278 279 def test_step_increments_step_count(self, env):280 env.reset(difficulty="easy")281 env.step(SQLAction(query="SELECT 1"))282 assert env.state.step_count == 1283 env.step(SQLAction(query="SELECT 2"))284 assert env.state.step_count == 2285 286 def test_step_before_reset_returns_error(self, env):287 obs = env.step(SQLAction(query="SELECT 1"))288 assert obs.error != "" or obs.done is True289 290 def test_step_updates_total_reward(self, env):291 env.reset(difficulty="easy")292 task = TASKS["easy_1"]293 env.step(SQLAction(query=task["reference_sql"], difficulty="easy"))294 assert env.state.total_reward > 0.0295 296 def test_step_with_correct_query(self, env):297 env.reset(difficulty="easy", task_id="easy_1")298 task = TASKS["easy_1"]299 obs = env.step(SQLAction(query=task["reference_sql"], difficulty="easy"))300 assert obs.reward >= 0.9301 assert len(obs.result) > 0302 303 def test_step_with_error_query(self, env):304 env.reset(difficulty="easy")305 obs = env.step(SQLAction(query="INVALID SQL !!!", difficulty="easy"))306 assert obs.reward <= 0.01 # minimum reward for errors307 308 def test_done_after_max_steps(self, env):309 env.reset(difficulty="easy")310 for _ in range(20):311 obs = env.step(SQLAction(query="SELECT 1", difficulty="easy"))312 if obs.done:313 break314 assert obs.done is True315 316 def test_multiple_resets_give_fresh_state(self, env):317 env.reset()318 env.step(SQLAction(query="SELECT 1"))319 first_id = env.state.episode_id320 321 env.reset()322 second_id = env.state.episode_id323 324 assert first_id != second_id325 assert env.state.step_count == 0326 327 328# ===========================================================================329# 7. All 9 Tasks Solvable End-to-End330# ===========================================================================331 332class TestAllTasksSolvable:333 """Each task should yield reward >= 0.9 when given the reference solution."""334 335 @pytest.mark.parametrize("task_id", list(TASKS.keys()))336 def test_task_solvable(self, task_id):337 env = SQLEnvironment()338 try:339 obs = env.reset(task_id=task_id,340 difficulty=TASKS[task_id]["difficulty"])341 task = TASKS[task_id]342 obs = env.step(SQLAction(343 query=task["reference_sql"],344 difficulty=task["difficulty"],345 task_id=task_id,346 ))347 assert obs.reward >= 0.9, (348 f"Task {task_id} scored {obs.reward}: {obs.message}"349 )350 finally:351 env.close()352 353 354# ===========================================================================355# 8. Multi-Difficulty Sweep356# ===========================================================================357 358class TestMultiDifficultySweep:359 360 def test_easy_sweep(self, env):361 obs = env.reset(difficulty="easy")362 total = 0.0363 for task_id in ["easy_1", "easy_2", "easy_3"]:364 task = TASKS[task_id]365 obs = env.step(SQLAction(366 query=task["reference_sql"],367 difficulty="easy",368 task_id=task_id,369 ))370 total += obs.reward371 if obs.done:372 break373 assert total > 0.0374 375 def test_medium_sweep(self, env):376 obs = env.reset(difficulty="medium")377 total = 0.0378 for task_id in ["medium_1", "medium_2", "medium_3"]:379 task = TASKS[task_id]380 obs = env.step(SQLAction(381 query=task["reference_sql"],382 difficulty="medium",383 task_id=task_id,384 ))385 total += obs.reward386 if obs.done:387 break388 assert total > 0.0389 390 def test_hard_sweep(self, env):391 obs = env.reset(difficulty="hard")392 total = 0.0393 for task_id in ["hard_1", "hard_2", "hard_3"]:394 task = TASKS[task_id]395 obs = env.step(SQLAction(396 query=task["reference_sql"],397 difficulty="hard",398 task_id=task_id,399 ))400 total += obs.reward401 if obs.done:402 break403 assert total > 0.0404 405 406# ===========================================================================407# 9. Inference Script Tests408# ===========================================================================409 410class TestInferenceScript:411 412 def test_inference_run_task_easy(self):413 """Run a single easy task with fallback heuristic (no LLM API)."""414 from inference import run_task415 from openai import OpenAI416 client = OpenAI(base_url="http://localhost:1", api_key="test")417 score = run_task("easy_1", "easy", client)418 assert 0.0 <= score <= 1.0419 assert score >= 0.9420 421 def test_inference_run_task_medium(self):422 from inference import run_task423 from openai import OpenAI424 client = OpenAI(base_url="http://localhost:1", api_key="test")425 score = run_task("medium_1", "medium", client)426 assert 0.0 <= score <= 1.0427 assert score >= 0.9428 429 def test_inference_run_task_hard(self):430 from inference import run_task431 from openai import OpenAI432 client = OpenAI(base_url="http://localhost:1", api_key="test")433 score = run_task("hard_1", "hard", client)434 assert 0.0 <= score <= 1.0435 assert score >= 0.9436 437 def test_inference_all_tasks(self):438 from inference import run_task, ALL_TASKS439 from openai import OpenAI440 client = OpenAI(base_url="http://localhost:1", api_key="test")441 for task_id, difficulty in ALL_TASKS:442 score = run_task(task_id, difficulty, client)443 assert 0.0 <= score <= 1.0, f"Task {task_id} score out of range: {score}"444 445 446# ===========================================================================447# 10. FastAPI HTTP Endpoint Tests (integration)448# ===========================================================================449 450class TestFastAPIEndpoints:451 452 @pytest.fixture(scope="class")453 def test_client(self):454 from fastapi.testclient import TestClient455 from server.app import app456 with TestClient(app) as client:457 yield client458 459 def test_health_returns_200(self, test_client):460 resp = test_client.get("/health")461 assert resp.status_code == 200462 data = resp.json()463 assert data["status"] == "ok"464 465 def test_reset_returns_observation(self, test_client):466 resp = test_client.post("/reset", json={"difficulty": "easy"})467 assert resp.status_code == 200468 data = resp.json()469 assert "task_description" in data470 assert "schema_info" in data471 assert data["reward"] == 0.0 # reset gives reward=0.0 (not graded)472 473 def test_step_returns_reward(self, test_client):474 test_client.post("/reset", json={"difficulty": "easy"})475 resp = test_client.post("/step", json={476 "query": "SELECT name, email FROM customers WHERE country = 'USA'",477 "difficulty": "easy",478 })479 assert resp.status_code == 200480 data = resp.json()481 assert "reward" in data482 assert 0.0 <= data["reward"] <= 1.0483 484 def test_state_returns_metadata(self, test_client):485 test_client.post("/reset", json={"difficulty": "easy"})486 resp = test_client.get("/state")487 assert resp.status_code == 200488 data = resp.json()489 assert "episode_id" in data490 assert "step_count" in data491 492 def test_step_with_bad_sql(self, test_client):493 test_client.post("/reset", json={"difficulty": "easy"})494 resp = test_client.post("/step", json={495 "query": "THIS IS NOT SQL",496 "difficulty": "easy",497 })498 assert resp.status_code == 200499 data = resp.json()500 assert data["reward"] <= 0.01 # minimum reward for bad SQL501 502 def test_tasks_endpoint(self, test_client):503 resp = test_client.get("/tasks")504 assert resp.status_code == 200505 data = resp.json()506 assert "easy" in data507 assert "medium" in data508 assert "hard" in data509 assert len(data["easy"]) == 3510 511 def test_schema_endpoint(self, test_client):512 resp = test_client.get("/schema")513 assert resp.status_code == 200514 data = resp.json()515 assert "schema" in data516 assert "customers" in data["schema"]517 518 def test_root_returns_html(self, test_client):519 resp = test_client.get("/")520 assert resp.status_code == 200521 assert "text/html" in resp.headers["content-type"]522 523 def test_full_easy_episode(self, test_client):524 """Complete episode solving easy_1 with reference solution."""525 test_client.post("/reset", json={"difficulty": "easy", "task_id": "easy_1"})526 resp = test_client.post("/step", json={527 "query": TASKS["easy_1"]["reference_sql"],528 "difficulty": "easy",529 "task_id": "easy_1",530 })531 assert resp.status_code == 200532 data = resp.json()533 assert data["reward"] >= 0.9534 535 def test_medium_join_query(self, test_client):536 test_client.post("/reset", json={"difficulty": "medium", "task_id": "medium_1"})537 resp = test_client.post("/step", json={538 "query": TASKS["medium_1"]["reference_sql"],539 "difficulty": "medium",540 "task_id": "medium_1",541 })542 data = resp.json()543 assert data["reward"] >= 0.9544 545 def test_hard_cte_query(self, test_client):546 test_client.post("/reset", json={"difficulty": "hard", "task_id": "hard_1"})547 resp = test_client.post("/step", json={548 "query": TASKS["hard_1"]["reference_sql"],549 "difficulty": "hard",550 "task_id": "hard_1",551 })552 data = resp.json()553 assert data["reward"] >= 0.9554 555 556# ===========================================================================557# 11. Reward Consistency Tests558# ===========================================================================559 560class TestRewardConsistency:561 562 def test_reward_deterministic(self, db):563 """Same query on same data should always give same reward."""564 task_id = "easy_1"565 query = TASKS[task_id]["reference_sql"]566 rewards = [grade(task_id, query, db)[0] for _ in range(5)]567 assert len(set(rewards)) == 1, "Reward is not deterministic"568 569 def test_correct_beats_incorrect(self, db):570 correct_reward, _, _, _ = grade(571 "easy_1",572 "SELECT name, email FROM customers WHERE country = 'USA'",573 db574 )575 wrong_reward, _, _, _ = grade(576 "easy_1",577 "SELECT name, email FROM customers WHERE country = 'Japan'",578 db579 )580 assert correct_reward > wrong_reward581 582 def test_reward_increases_with_more_matches(self, db):583 """Returning more matching rows should give higher reward."""584 # All USA customers585 full_reward, _, _, _ = grade(586 "easy_1",587 "SELECT name, email FROM customers WHERE country = 'USA'",588 db589 )590 # Only one USA customer591 partial_reward, _, _, _ = grade(592 "easy_1",593 "SELECT name, email FROM customers WHERE country = 'USA' LIMIT 1",594 db595 )596 assert full_reward >= partial_reward597 598 599# ===========================================================================600# Entry point601# ===========================================================================602 603if __name__ == "__main__":604 # Quick smoke test without pytest605 print("Running smoke tests...")606 import sqlite3607 608 conn = sqlite3.connect(":memory:")609 conn.row_factory = sqlite3.Row610 seed_database(conn)611 612 print(f" Tasks registered: {len(TASKS)}")613 for diff in ["easy", "medium", "hard"]:614 tasks = get_all_tasks_by_difficulty(diff)615 print(f" {diff.capitalize()} tasks: {len(tasks)}")616 617 # Test all reference solutions618 passed = 0619 failed = 0620 for task_id, task in TASKS.items():621 reward, msg, _, _ = grade(task_id, task["reference_sql"], conn)622 status = "PASS" if reward >= 0.9 else "FAIL"623 if reward >= 0.9:624 passed += 1625 else:626 failed += 1627 print(f" [{status}] {task_id}: reward={reward:.4f} - {msg[:50]}")628 629 conn.close()630 print(f"\nSmoke test: {passed} passed, {failed} failed")631 632 # Test environment lifecycle633 env = SQLEnvironment()634 obs = env.reset(difficulty="easy")635 print(f"\nEnvironment reset: task={obs.task_description[:60]}...")636 obs = env.step(SQLAction(query=TASKS["easy_1"]["reference_sql"], difficulty="easy"))637 print(f"Step result: reward={obs.reward}, rows={len(obs.result)}")638 env.close()639 640 print("\nAll smoke tests complete.")641 