Spaces:
Sleeping
Sleeping
File size: 1,237 Bytes
341aec9 e965a47 341aec9 5716a3a e965a47 5716a3a e965a47 5716a3a e965a47 5716a3a e965a47 5716a3a e965a47 341aec9 | 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 | from typing import List, Optional, Any, Callable
from openenv.core.env_server.types import Action, Observation, State
from pydantic import Field, BaseModel
class SQLAction(Action):
corrected_query: str
class SQLObservation(Observation):
task_id: str
broken_query: str
schema_context: Optional[str] = None
error_hint: Optional[str] = None
step_number: int
previous_attempt: Optional[str] = None
feedback: Optional[str] = None
reward: float = 0.0
done: bool = False
class SQLState(State):
task_id: str
difficulty: str
step_count: int
max_steps: int
done: bool
last_reward: float
rewards_history: List[float]
class SQLReward(BaseModel):
value: float = Field(gt=0.0, lt=1.0)
reason: str
class SQLTask(BaseModel):
task_id: str
difficulty: str
broken_query: str
canonical_answer: str
schema_context: Optional[str] = None
error_hint: Optional[str] = None
max_steps: int = 5
grader: Optional[Any] = Field(default=None, exclude=True)
model_config = {"arbitrary_types_allowed": True}
class StepResult(BaseModel):
observation: SQLObservation
reward: float
done: bool
info: dict = Field(default_factory=dict) |