Spaces:
Sleeping
Sleeping
File size: 4,142 Bytes
1a1713a 0e2347e 1a1713a 0e2347e 1a1713a 95707b2 | 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 108 109 110 111 112 113 114 115 116 117 118 119 120 | """
FastAPI HTTP wrapper for SQLCorrectionEnv using openenv.core base classes.
"""
from openenv.core.env_server import create_fastapi_app
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State
from sql_env.models import SQLAction, SQLObservation
from sql_env.tasks import ALL_TASKS, TASK_SETS
from sql_env.grader import grade, generate_feedback
import random
class SQLCorrectionEnvironment(Environment):
def __init__(self):
super().__init__()
self._difficulty = "easy"
self._current_task = None
self._step_count = 0
self._done = False
self._last_reward = 0.0
self._previous_attempt = None
self._feedback = None
self._rewards_history = []
def reset(self, difficulty: str = "easy") -> SQLObservation:
self._difficulty = difficulty
tasks = TASK_SETS.get(difficulty, TASK_SETS["easy"])
self._current_task = random.choice(tasks)
self._step_count = 0
self._done = False
self._last_reward = 0.0
self._previous_attempt = None
self._feedback = None
self._rewards_history = []
return SQLObservation(
task_id=self._current_task.task_id,
broken_query=self._current_task.broken_query,
schema_context=self._current_task.schema_context,
error_hint=self._current_task.error_hint,
step_number=0,
previous_attempt=None,
feedback=None,
)
def step(self, action: SQLAction) -> SQLObservation:
self._step_count += 1
reward_obj = grade(action, self._current_task)
reward = reward_obj.value
self._last_reward = reward
self._previous_attempt = action.corrected_query
self._feedback = generate_feedback(action, self._current_task, reward_obj)
self._rewards_history.append(reward)
done = (reward >= 0.95) or (self._step_count >= self._current_task.max_steps)
self._done = done
return SQLObservation(
task_id=self._current_task.task_id,
broken_query=self._current_task.broken_query,
schema_context=self._current_task.schema_context,
error_hint=self._current_task.error_hint,
step_number=self._step_count,
previous_attempt=self._previous_attempt,
feedback=self._feedback,
)
@property
def state(self) -> dict:
if self._current_task is None:
return {"status": "not_initialized"}
return {
"task_id": self._current_task.task_id,
"difficulty": self._difficulty,
"step_count": self._step_count,
"done": self._done,
"last_reward": self._last_reward,
"rewards_history": self._rewards_history,
}
def get_tasks(self):
return {
"tasks": [
{
"name": "easy",
"difficulty": "easy",
"description": "Fix a single syntax error. Error hint provided.",
"max_steps": 5,
"has_grader": True,
"grader": "sql_env.grader.grade",
},
{
"name": "medium",
"difficulty": "medium",
"description": "Fix multiple errors. No hint.",
"max_steps": 5,
"has_grader": True,
"grader": "sql_env.grader.grade",
},
{
"name": "hard",
"difficulty": "hard",
"description": "Fix complex multi-join queries. Schema provided.",
"max_steps": 4,
"has_grader": True,
"grader": "sql_env.grader.grade",
},
]
}
env = SQLCorrectionEnvironment()
app = create_fastapi_app(env, SQLAction, SQLObservation)
def main():
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)
if __name__ == "__main__":
main() |