RonyForAI/Mirage_DB_RL
0
1from openenv.core import EnvClient2from openenv.core.client_types import StepResult3from Mirage_RL.models import QueryAction, QueryObservation, QueryState4 5 6class QueryClient(EnvClient[QueryAction, QueryObservation, QueryState]):7 8 def _step_payload(self, action: QueryAction):9 return {10 "next_table": action.next_table,11 "join_type": action.join_type,12 "use_index": action.use_index,13 }14 15 def _parse_result(self, payload):16 obs = payload.get("observation", {})17 18 observation = QueryObservation(19 done = payload.get("done", False),20 reward = payload.get("reward", 0.0),21 tables = obs.get("tables", []),22 table_rows = obs.get("table_rows", []),23 selectivities = obs.get("selectivities", []),24 has_index = obs.get("has_index", []),25 chosen_order = obs.get("chosen_order", []),26 remaining_tables = obs.get("remaining_tables", []),27 step_number = obs.get("step_number", 0),28 current_cost = obs.get("current_cost", 0.0),29 query_context = obs.get("query_context", ""),30 intermediate_size = obs.get("intermediate_size", 1.0),31 )32 33 return StepResult(34 observation = observation,35 reward = payload.get("reward", 0.0),36 done = payload.get("done", False),37 )38 39 def _parse_state(self, payload):40 return QueryState(41 chosen_order = payload.get("chosen_order", []),42 remaining_tables = payload.get("remaining_tables", []),43 current_cost = payload.get("current_cost", 0.0),44 final_cost = payload.get("final_cost", 0.0),45 scenario_name = payload.get("scenario_name", ""),46 )