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