Ashraf01k's picture
fix(openenv): eradicate exact 0.0 scores triggering Phase 2 bot boundary crash loops
1aecc95
Raw History Blame Contribute Delete
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