Spaces:
Paused
Paused
File size: 4,100 Bytes
60dfa24 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | """Fast smoke tests for CI (DuckDB warm-up once per process)."""
from __future__ import annotations
import os
import sys
import pytest
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
from env import SQLOptimEnv # noqa: E402
from graders import GradeMask, grade # noqa: E402
from models import Action # noqa: E402
from tasks import TASKS # noqa: E402
@pytest.fixture(scope="module")
def executor():
from executor import get_executor
return get_executor()
def test_executor_compare_task1(executor):
task = TASKS["task_1_basic_antipatterns"]
original = task["sql_query"]
optimized = (
"SELECT id, customer_id, product_id, status, total, created_at FROM orders "
"WHERE customer_id = 5000 "
"AND created_at >= DATE '2024-01-01' AND created_at < DATE '2025-01-01'"
)
r = executor.compare(original, optimized)
assert r["speedup"] >= 1.0
assert r["results_match"] is True
def test_grade_mask_changes_total():
task = TASKS["task_1_basic_antipatterns"]
action = Action(
suggestions=[
{
"issue_type": "select_star",
"line": 1,
"description": "SELECT * on large table",
"severity": "high",
"fix": "project columns",
}
],
optimized_query=(
"SELECT id, customer_id, product_id, status, total, created_at FROM orders "
"WHERE customer_id = 5000 AND created_at >= DATE '2024-01-01' "
"AND created_at < DATE '2025-01-01'"
),
summary="x" * 130,
estimated_improvement="5x",
approved=False,
)
full = grade(task, action).score
no_exec = grade(
task, action, mask=GradeMask(execution_speedup=False, result_correctness=False)
).score
assert no_exec < full
def test_fastapi_reset_step():
from fastapi.testclient import TestClient
from server.app import app
client = TestClient(app)
r = client.get("/")
assert r.status_code == 200
assert r.json()["environment"] == "sql-optim-env"
obs = client.post("/reset", json={"task_id": "task_1_basic_antipatterns"}).json()
assert obs["task_id"] == "task_1_basic_antipatterns"
step = client.post(
"/step",
json={
"suggestions": [
{
"issue_type": "select_star",
"line": 1,
"description": "avoid star",
"severity": "high",
"fix": "cols",
}
],
"optimized_query": (
"SELECT id, customer_id, product_id, status, total, created_at FROM orders "
"WHERE customer_id = 5000 "
"AND created_at >= DATE '2024-01-01' AND created_at < DATE '2025-01-01'"
),
"summary": "Rewrite removes anti-patterns and uses a sargable date range.",
"estimated_improvement": "4x",
"approved": False,
},
)
assert step.status_code == 200
body = step.json()
assert "reward" in body
assert body["reward"]["score"] > 0.5
def test_sqoptim_env_reset_step():
env = SQLOptimEnv()
obs = env.reset(task_id="task_1_basic_antipatterns")
assert obs.step_count == 0
result = env.step(
Action(
suggestions=[
{
"issue_type": "select_star",
"line": 1,
"description": "SELECT *",
"severity": "high",
"fix": "list columns",
}
],
optimized_query=(
"SELECT id, customer_id, product_id, status, total, created_at FROM orders "
"WHERE customer_id = 5000 "
"AND created_at >= DATE '2024-01-01' AND created_at < DATE '2025-01-01'"
),
summary="A" * 130,
estimated_improvement="5x",
approved=False,
)
)
assert result.reward.score > 0.4
|