Spaces:
Sleeping
Sleeping
Download sql_env/env.py from pheonixBond007/sql-review-env: direct link, hf CLI and curl.
- Browser
- Download file 3.47 kB
-
https://huggingface.co/spaces/pheonixBond007/sql-review-env/resolve/main/sql_env/env.py
- Command line
-
hf download hf://spaces/pheonixBond007/sql-review-env/sql_env/env.py
-
curl -L -o env.py https://huggingface.co/spaces/pheonixBond007/sql-review-env/resolve/main/sql_env/env.py
3.47 kB
| 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 | |