Hariprita/nl2sql-openenv
0
1from dataclasses import dataclass, field2from typing import Optional3import uuid4 5try:6 from openenv.core.env_server.interfaces import Environment as BaseEnvironment7 from openenv.core.env_server.types import State as BaseState8 HAS_OPENENV = True9except ImportError:10 HAS_OPENENV = False11 12 class BaseEnvironment:13 pass14 15 @dataclass16 class BaseState:17 episode_id: str = ""18 step_count: int = 019 20 21@dataclass22class SQLAction:23 sql_query: str24 25 def to_dict(self):26 return {"sql_query": self.sql_query}27 28 29@dataclass30class SQLObservation:31 schema: str # Database DDL shown to the agent32 question: str # Natural language question33 result: str # Query execution result or error message34 reward: float # 0.0 to 1.035 done: bool36 feedback: str # Human-readable explanation of reward37 task_id: str # Which task is active38 task_difficulty: str # easy / medium / hard39 attempt: int # Which attempt (1, 2, or 3)40 max_attempts: int # Always 341 hint: str = ""42 43 @property44 def goal(self):45 """Alias so inference scripts can use observation.goal"""46 return self.question47 48 @property49 def last_action_error(self):50 return self.result.startswith("ERROR:") if self.result else False51 52 @property53 def url(self):54 """Stub for compatibility with generic inference scripts"""55 return f"task://{self.task_id}/attempt/{self.attempt}"56 57 58@dataclass59class SQLState:60 episode_id: str = field(default_factory=lambda: str(uuid.uuid4()))61 step_count: int = 062 current_task_id: str = ""63 current_task_idx: int = 064 total_tasks: int = 365 cumulative_reward: float = 0.066 