Spaces:
Sleeping
Sleeping
File size: 4,188 Bytes
1a1713a e965a47 1a1713a e965a47 c4a66be f6bd747 e965a47 6c6f994 e965a47 f6bd747 e965a47 f6bd747 c4a66be 1a1713a 66cc86e | 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 121 122 123 | """
FastAPI server using openenv.core base classes — required for validator.
"""
import random
from typing import Optional
from openenv.core.env_server.http_server import create_app
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State
try:
from sql_env.models import SQLAction, SQLObservation, SQLState
from sql_env.tasks import TASK_SETS
from sql_env.grader import grade, generate_feedback
except ImportError:
from models import SQLAction, SQLObservation, SQLState
from tasks import TASK_SETS
from grader import grade, generate_feedback
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._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._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,
reward=0.0,
done=False,
)
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._rewards_history.append(reward)
done = (reward >= 0.95) or (self._step_count >= self._current_task.max_steps)
self._done = done
feedback = generate_feedback(action, self._current_task, reward_obj)
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=action.corrected_query,
feedback=feedback,
reward=reward,
done=done,
)
@property
def state(self) -> SQLState:
if self._current_task is None:
return SQLState(
task_id="none",
difficulty="none",
step_count=0,
max_steps=0,
done=False,
last_reward=0.0,
rewards_history=[],
)
return SQLState(
task_id=self._current_task.task_id,
difficulty=self._difficulty,
step_count=self._step_count,
max_steps=self._current_task.max_steps,
done=self._done,
last_reward=self._last_reward,
rewards_history=self._rewards_history,
)
app = create_app(
SQLCorrectionEnvironment,
SQLAction,
SQLObservation,
env_name="sql-correction-env",
max_concurrent_envs=10,
)
def main():
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)
if __name__ == "__main__":
main()
from fastapi import Request
@app.get("/tasks")
async def list_tasks():
"""Return graded tasks in openenv validator format."""
from sql_env.grader import grade
return {
"tasks": [
{"id": "easy", "difficulty": "easy", "description": "Fix a single syntax error.", "steps": 5, "ideal_action": "correct_sql", "has_grader": True},
{"id": "medium", "difficulty": "medium", "description": "Fix multiple errors.", "steps": 5, "ideal_action": "correct_sql", "has_grader": True},
{"id": "hard", "difficulty": "hard", "description": "Fix complex multi-join queries.", "steps": 4, "ideal_action": "correct_sql", "has_grader": True},
]
}
|