File size: 3,473 Bytes
5a94bfd
2c8c77b
09885b6
5a94bfd
 
 
 
 
 
 
 
 
296e2a5
1aecc95
5a94bfd
2c8c77b
 
09885b6
2c8c77b
5a94bfd
 
 
 
2c8c77b
5a94bfd
 
1aecc95
2c8c77b
1aecc95
2c8c77b
5a94bfd
2c8c77b
 
5a94bfd
 
 
2c8c77b
 
 
 
 
 
 
 
 
 
c59cc5e
 
 
2c8c77b
 
 
 
5a94bfd
 
 
 
 
 
 
 
2c8c77b
5a94bfd
 
2c8c77b
 
5a94bfd
 
2c8c77b
 
 
5a94bfd
 
2c8c77b
5a94bfd
 
 
2c8c77b
5a94bfd
 
09885b6
2c8c77b
 
 
5a94bfd
 
 
 
 
 
2c8c77b
5a94bfd
 
 
 
 
 
09885b6
 
5a94bfd
c59cc5e
 
 
 
 
 
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
import sqlite3
import time
from typing import Optional, Dict, Any, Tuple, List
from .models import SQLObservation, SQLAction, SQLReward
from .tasks import TASKS, load_fixtures
from .graders import grade_sql

class SQLReviewEnv:
    def __init__(self):
        self.conn = None
        self.task_id = "syntax-fix"
        self.step_count = 0
        self.max_steps = 8
        self.last_reward = 0.01
        self.done = False
        self.last_error: Optional[str] = None
        self._query_start_time = 0
        self.history: List[float] = []

    async def reset(self, task_id: str = "syntax-fix") -> SQLObservation:
        self.task_id = task_id
        if self.task_id not in TASKS:
            self.task_id = "syntax-fix"

        self.step_count = 0
        self.done = False
        self.last_reward = 0.01
        self.last_error = None
        self.history = [0.01]

        if self.conn:
            self.conn.close()
            
        self.conn = sqlite3.connect(":memory:")
        self.conn.row_factory = sqlite3.Row
        
        # Absolute execution speed limits and RAM constraints
        self.conn.execute("PRAGMA max_page_count = 10000;")
        
        def timeout_handler():
            if time.time() - self._query_start_time > 1.5:
                # Return 1 to violently abort the massive sqlite block
                return 1
            return 0
            
        self.conn.set_progress_handler(timeout_handler, 1000)

        # Anchor start time before fixtures load so the timeout is valid immediately
        self._query_start_time = time.time()
        load_fixtures(self.conn, self.task_id)

        task_data = TASKS[self.task_id]

        return SQLObservation(
            task_id=self.task_id,
            db_schema=task_data["db_schema"],
            query=task_data["query"],
            error_message=None,
            expected_hint=task_data["expected_hint"],
            step=self.step_count
        )

    async def step(self, action: SQLAction) -> SQLReward:
        if self.done:
            return SQLReward(value=self.last_reward, breakdown={}, done=True, info={"error": self.last_error or ""})

        self.step_count += 1
        
        # Anchor the time bounds logic prior to grader evaluation
        self._query_start_time = time.time()

        task_data = TASKS[self.task_id]
        expected_sql = task_data["validation_query"]

        reward_val, breakdown, done, info = grade_sql(
            self.task_id, self.conn, action.sql, expected_sql, self.step_count, self.max_steps
        )

        self.last_reward = reward_val
        self.done = done
        self.history.append(reward_val)
        # Persist the latest error message for next observation if needed
        self.last_error = info.get("error") or info.get("validation_error") or info.get("plan_error") or None

        return SQLReward(
            value=reward_val,
            breakdown=breakdown,
            done=done,
            info=info
        )

    def state(self) -> dict:
        return {
            "task_id": self.task_id,
            "current_step": self.step_count,
            "max_steps": self.max_steps,
            "last_reward": self.last_reward,
            "done": self.done,
            "history": self.history
        }

    def close(self) -> None:
        """Cleanly close the SQLite connection. Called on server shutdown."""
        if self.conn:
            self.conn.close()
            self.conn = None