sravaniamere commited on
Commit
154f3d9
·
1 Parent(s): 95707b2

fix grader serialization and tasks endpoint

Browse files
Files changed (2) hide show
  1. sql_env/models.py +2 -3
  2. sql_env/server.py +28 -6
sql_env/models.py CHANGED
@@ -25,10 +25,9 @@ class SQLTask(BaseModel):
25
  schema_context: Optional[str] = None
26
  error_hint: Optional[str] = None
27
  max_steps: int = 5
28
- grader: Optional[Callable] = None # ← add this
29
 
30
- class Config:
31
- arbitrary_types_allowed = True # ← required for Callable in Pydantic
32
 
33
  class StepResult(BaseModel):
34
  observation: SQLObservation
 
25
  schema_context: Optional[str] = None
26
  error_hint: Optional[str] = None
27
  max_steps: int = 5
28
+ grader: Optional[Any] = Field(default=None, exclude=True)
29
 
30
+ model_config = {"arbitrary_types_allowed": True}
 
31
 
32
  class StepResult(BaseModel):
33
  observation: SQLObservation
sql_env/server.py CHANGED
@@ -99,13 +99,35 @@ async def health():
99
 
100
  @app.get("/tasks")
101
  async def list_tasks():
102
- """Return graded tasks by difficulty (RL validator format)."""
103
- graded_tasks = {
104
- diff: [task.__dict__ for task in tasks if task.grader is not None]
105
- for diff, tasks in ALL_TASKS.items()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
  }
107
- return graded_tasks # {"easy": [tasks], "medium": [tasks], "hard": [tasks]}
108
-
109
 
110
  @app.get("/")
111
  async def root():
 
99
 
100
  @app.get("/tasks")
101
  async def list_tasks():
102
+ """Return graded tasks in openenv validator format."""
103
+ return {
104
+ "tasks": [
105
+ {
106
+ "name": "easy",
107
+ "difficulty": "easy",
108
+ "description": "Fix a single syntax error. Error hint provided.",
109
+ "max_steps": 5,
110
+ "has_grader": True,
111
+ "grader": "grade",
112
+ },
113
+ {
114
+ "name": "medium",
115
+ "difficulty": "medium",
116
+ "description": "Fix multiple errors across keywords and clauses. No hint.",
117
+ "max_steps": 5,
118
+ "has_grader": True,
119
+ "grader": "grade",
120
+ },
121
+ {
122
+ "name": "hard",
123
+ "difficulty": "hard",
124
+ "description": "Fix many errors in complex multi-join queries. Schema provided.",
125
+ "max_steps": 4,
126
+ "has_grader": True,
127
+ "grader": "grade",
128
+ },
129
+ ]
130
  }
 
 
131
 
132
  @app.get("/")
133
  async def root():