Team Ai
Apppublic

Codexzzz/sql-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
client.py74 linesDownload Raw Back to sql_env
1# from openenv.core.env_client import EnvClient2# from openenv.core.client_types import StepResult3# from openenv.core.env_server.types import State4# from .models import SqlAction, SqlObservation5 6 7# class SqlEnv(EnvClient[SqlAction, SqlObservation, State]):8 9#     def _step_payload(self, action: SqlAction) -> dict:10#         return {"sql_query": action.sql_query}11 12#     def _parse_result(self, payload: dict) -> StepResult[SqlObservation]:13#         obs_data = payload.get("observation", {})14#         obs = SqlObservation(15#             task_description  = obs_data.get("task_description",   ""),16#             schema_info       = obs_data.get("schema_info",         ""),17#             query_result      = obs_data.get("query_result",        []),18#             error_message     = obs_data.get("error_message",       ""),19#             feedback          = obs_data.get("feedback",            ""),20#             score_breakdown   = obs_data.get("score_breakdown",     {}),21#             attempts_remaining= obs_data.get("attempts_remaining",   0),22#             done              = payload.get("done",                False),23#             reward            = payload.get("reward",              0.0),24#         )25#         return StepResult(26#             observation=obs,27#             reward=payload.get("reward", 0.0),28#             done=payload.get("done", False),29#         )30 31#     def _parse_state(self, payload: dict) -> State:32#         return State(33#             episode_id=payload.get("episode_id"),34#             step_count=payload.get("step_count", 0),35#         )36 37 38from openenv.core.env_client import EnvClient39from openenv.core.client_types import StepResult40from openenv.core.env_server.types import State41from .models import SqlAction, SqlObservation42 43 44class SqlEnv(EnvClient[SqlAction, SqlObservation, State]):45 46    def _step_payload(self, action: SqlAction) -> dict:47        return {"sql_query": action.sql_query}48 49    def _parse_result(self, payload: dict) -> StepResult[SqlObservation]:50        obs_data = payload.get("observation", {})51 52        obs = SqlObservation(53            task_description   = obs_data.get("task_description",   ""),54            schema_info        = obs_data.get("schema_info",         ""),55            query_result       = obs_data.get("query_result",        []),56            error_message      = obs_data.get("error_message",       ""),57            feedback           = obs_data.get("feedback",            ""),58            score_breakdown    = obs_data.get("score_breakdown",     {}),59            attempts_remaining = obs_data.get("attempts_remaining",   0),60            done               = obs_data.get("done",   payload.get("done",   False)),61            reward             = obs_data.get("reward", payload.get("reward", 0.0)),62        )63 64        return StepResult(65            observation = obs,66            reward      = payload.get("reward", 0.0),67            done        = payload.get("done",   False),68        )69 70    def _parse_state(self, payload: dict) -> State:71        return State(72            episode_id = payload.get("episode_id"),73            step_count = payload.get("step_count", 0),74        )