Spaces:
Sleeping
Sleeping
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
|