CoolFace
Apppublic

Tsah00/sql-env

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
test_env.py641 linesDownload Raw Back to root
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