VarshaChukka/SQL_Debugging_REPL_Environment
0
1import sqlite3
2from app.tasks import get_task
3from app.grader import grade
4
5class SQLEnv:
6 def __init__(self):
7 self.conn = None
8 self.current_task = None
9
10 def reset(self, task_id=0):
11 self.conn = sqlite3.connect(":memory:")
12 self.current_task = get_task(task_id)
13
14 for stmt in self.current_task["setup"]:
15 self.conn.execute(stmt)
16
17 observation = {
18 "problem": self.current_task["problem"],
19 "db_schema": self.current_task["schema"],
20 "last_error": None,
21 "query_result": None
22 }
23
24 return observation
25
26 def step(self, action):
27 try:
28 cursor = self.conn.execute(action.sql_query)
29 result = cursor.fetchall()
30
31 score, feedback = grade(result, self.current_task)
32
33 observation = {
34 "problem": self.current_task["problem"],
35 "db_schema": self.current_task["schema"],
36 "last_error": None,
37 "query_result": str(result)
38 }
39
40 reward = {
41 "score": float(score),
42 "feedback": feedback
43 }
44
45 done = score == 1.0
46 info = {}
47
48 return observation, reward, done, info
49
50 except Exception as e:
51 observation = {
52 "problem": self.current_task["problem"],
53 "db_schema": self.current_task["schema"],
54 "last_error": str(e),
55 "query_result": None
56 }
57
58 reward = {
59 "score": 0.0, # ⚠️ FIXED (no negative)
60 "feedback": "SQL Error"
61 }
62
63 return observation, reward, False, {}
64
65 def state(self):
66 return {
67 "problem": self.current_task["problem"],
68 "db_schema": self.current_task["schema"]
69 }