File size: 7,671 Bytes
1a1713a
e965a47
1a1713a
e965a47
c4a66be
f6bd747
e965a47
 
6c6f994
e965a47
 
 
 
 
 
 
 
 
 
 
f52aed3
e965a47
 
 
 
 
 
 
 
 
f52aed3
 
 
 
 
 
7a6f18c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3fa3f1b
7a6f18c
 
fd851a9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f52aed3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fd851a9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c41f6ba
c7d40b5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e965a47
 
 
 
 
 
 
 
 
 
3fa3f1b
e965a47
 
 
 
 
 
 
 
 
 
 
 
f6bd747
e965a47
 
 
f6bd747
c4a66be
 
 
1a1713a
 
 
 
 
 
66cc86e
 
 
 
 
 
 
c41f6ba
 
 
66cc86e
 
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
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
"""
FastAPI server using openenv.core base classes — required for validator.
"""
import random
from typing import Optional
from openenv.core.env_server.http_server import create_app
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State

try:
    from sql_env.models import SQLAction, SQLObservation, SQLState
    from sql_env.tasks import TASK_SETS
    from sql_env.grader import grade, generate_feedback
except ImportError:
    from models import SQLAction, SQLObservation, SQLState
    from tasks import TASK_SETS
    from grader import grade, generate_feedback


class SQLCorrectionEnvironment(Environment):
    SUPPORTS_CONCURRENT_SESSIONS = True

    def __init__(self):
        self._difficulty = "easy"
        self._current_task = None
        self._step_count = 0
        self._done = False
        self._last_reward = 0.0
        self._rewards_history = []

    def reset(self, seed=None, episode_id=None, **kwargs) -> SQLObservation:
        actual_difficulty = (
            kwargs.get("task_id") or
            kwargs.get("difficulty") or
            "easy"
        )
        self._difficulty = actual_difficulty
        tasks = TASK_SETS.get(actual_difficulty, TASK_SETS["easy"])
        self._current_task = random.choice(tasks)
        self._step_count = 0
        self._done = False
        self._last_reward = 0.0
        self._rewards_history = []
        return SQLObservation(
            task_id=self._current_task.task_id,
            broken_query=self._current_task.broken_query,
            schema_context=self._current_task.schema_context,
            error_hint=self._current_task.error_hint,
            step_number=0,
            previous_attempt=None,
            feedback=None,
            reward=0.001,
            done=False,
        )

    def step(self, action: SQLAction) -> SQLObservation:
        if self._current_task is None:
            self.reset()
        self._step_count += 1
        reward_obj = grade(action, self._current_task)
        reward = reward_obj.value
        self._last_reward = reward
        self._rewards_history.append(reward)
        done = (reward >= 0.95) or (self._step_count >= self._current_task.max_steps)
        self._done = done
        feedback = generate_feedback(action, self._current_task, reward_obj)
        return SQLObservation(
            task_id=self._current_task.task_id,
            broken_query=self._current_task.broken_query,
            schema_context=self._current_task.schema_context,
            error_hint=self._current_task.error_hint,
            step_number=self._step_count,
            previous_attempt=action.corrected_query,
            feedback=feedback,
            reward=reward,
            done=done,
        )

    @property
    def state(self) -> SQLState:
        if self._current_task is None:
            return SQLState(
                task_id="none",
                difficulty="none",
                step_count=0,
                max_steps=0,
                done=False,
                last_reward=0.0,
                rewards_history=[],
            )
        return SQLState(
            task_id=self._current_task.task_id,
            difficulty=self._difficulty,
            step_count=self._step_count,
            max_steps=self._current_task.max_steps,
            done=self._done,
            last_reward=self._last_reward,
            rewards_history=self._rewards_history,
        )
    def step(self, action: SQLAction) -> SQLObservation:
        if self._current_task is None:
            self.reset()
        self._step_count += 1
        reward_obj = grade(action, self._current_task)
        reward = reward_obj.value
        self._last_reward = reward
        self._rewards_history.append(reward)
        done = (reward >= 0.95) or (self._step_count >= self._current_task.max_steps)
        self._done = done
        feedback = generate_feedback(action, self._current_task, reward_obj)
        return SQLObservation(
            task_id=self._current_task.task_id,
            broken_query=self._current_task.broken_query,
            schema_context=self._current_task.schema_context,
            error_hint=self._current_task.error_hint,
            step_number=self._step_count,
            previous_attempt=action.corrected_query,
            feedback=feedback,
            reward=reward,
            done=done,
        )

    def state(self) -> SQLState:
        if self._current_task is None:
            return SQLState(
                task_id="none",
                difficulty="none",
                step_count=0,
                max_steps=0,
                done=False,
                last_reward=0.0,
                rewards_history=[],
            )
        return SQLState(
            task_id=self._current_task.task_id,
            difficulty=self._difficulty,
            step_count=self._step_count,
            max_steps=self._current_task.max_steps,
            done=self._done,
            last_reward=self._last_reward,
            rewards_history=self._rewards_history,
        )
    
def step(self, action: SQLAction) -> SQLObservation:
    # Auto-reset if no task loaded (create_app may use fresh instances)
    if self._current_task is None:
        self.reset()
    
    self._step_count += 1
    reward_obj = grade(action, self._current_task)
    reward = reward_obj.value
    self._last_reward = reward
    self._rewards_history.append(reward)
    done = (reward >= 0.95) or (self._step_count >= self._current_task.max_steps)
    self._done = done
    feedback = generate_feedback(action, self._current_task, reward_obj)
    return SQLObservation(
        task_id=self._current_task.task_id,
        broken_query=self._current_task.broken_query,
        schema_context=self._current_task.schema_context,
        error_hint=self._current_task.error_hint,
        step_number=self._step_count,
        previous_attempt=action.corrected_query,
        feedback=feedback,
        reward=reward,
        done=done,
    )

    @property
    def state(self) -> SQLState:
        if self._current_task is None:
            return SQLState(
                task_id="none",
                difficulty="none",
                step_count=0,
                max_steps=0,
                done=False,
                last_reward=0.001,
                rewards_history=[],
            )
        return SQLState(
            task_id=self._current_task.task_id,
            difficulty=self._difficulty,
            step_count=self._step_count,
            max_steps=self._current_task.max_steps,
            done=self._done,
            last_reward=self._last_reward,
            rewards_history=self._rewards_history,
        )

app = create_app(
    SQLCorrectionEnvironment,
    SQLAction,
    SQLObservation,
    env_name="sql-correction-env",
)


def main():
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=7860)


if __name__ == "__main__":
    main()
from fastapi import Request

@app.get("/tasks")
async def list_tasks():
    return {
        "tasks": [
            {"id": "easy", "difficulty": "easy", "description": "Fix a single syntax error.", "steps": 5, "ideal_action": "correct_sql", "has_grader": True, "grader": "sql_env.grader.grade"},
            {"id": "medium", "difficulty": "medium", "description": "Fix multiple errors.", "steps": 5, "ideal_action": "correct_sql", "has_grader": True, "grader": "sql_env.grader.grade"},
            {"id": "hard", "difficulty": "hard", "description": "Fix complex multi-join queries.", "steps": 4, "ideal_action": "correct_sql", "has_grader": True, "grader": "sql_env.grader.grade"},
        ]
    }