ritvik360 commited on
Commit
ae4bbb1
·
verified ·
1 Parent(s): a069777

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. client.py +7 -10
  2. server/environment.py +5 -2
client.py CHANGED
@@ -51,15 +51,12 @@ class NL2SQLEnv:
51
  return self._parse_result(resp.json())
52
 
53
  def _parse_result(self, payload: Dict[str, Any]) -> StepResult:
54
- obs_data = payload.get("observation", payload)
55
 
56
- # SAFETY CHECK: If reward or score is None/null, default to 0.0
57
- raw_reward = obs_data.get("reward")
58
- safe_reward = float(raw_reward) if raw_reward is not None else 0.0
59
 
60
- raw_score = obs_data.get("score")
61
- safe_score = float(raw_score) if raw_score is not None else 0.0
62
-
63
  obs = NL2SQLObservation(
64
  question=obs_data.get("question", ""),
65
  schema_context=obs_data.get("schema_context", ""),
@@ -70,12 +67,12 @@ class NL2SQLEnv:
70
  result_columns=obs_data.get("result_columns", []),
71
  step=obs_data.get("step", 0),
72
  max_steps=obs_data.get("max_steps", 5),
73
- done=obs_data.get("done", False),
74
  reward=safe_reward,
75
- score=safe_score,
76
  )
77
  return StepResult(
78
  observation=obs,
79
  reward=safe_reward,
80
- done=obs.done,
81
  )
 
51
  return self._parse_result(resp.json())
52
 
53
  def _parse_result(self, payload: Dict[str, Any]) -> StepResult:
54
+ obs_data = payload.get("observation", {})
55
 
56
+ # Read standard OpenEnv top-level keys safely
57
+ safe_reward = float(payload.get("reward", 0.0))
58
+ safe_done = bool(payload.get("done", False))
59
 
 
 
 
60
  obs = NL2SQLObservation(
61
  question=obs_data.get("question", ""),
62
  schema_context=obs_data.get("schema_context", ""),
 
67
  result_columns=obs_data.get("result_columns", []),
68
  step=obs_data.get("step", 0),
69
  max_steps=obs_data.get("max_steps", 5),
70
+ done=safe_done,
71
  reward=safe_reward,
72
+ score=float(obs_data.get("score", 0.0)),
73
  )
74
  return StepResult(
75
  observation=obs,
76
  reward=safe_reward,
77
+ done=safe_done,
78
  )
server/environment.py CHANGED
@@ -152,7 +152,7 @@ class NL2SQLEnvironment(Environment):
152
  self._last_obs = obs
153
  return obs
154
 
155
- def step(self, action: NL2SQLAction) -> NL2SQLObservation:
156
  """Execute the agent's SQL and return graded observation."""
157
  if self._task is None or self._example is None:
158
  # Called before reset — auto-reset
@@ -210,8 +210,11 @@ class NL2SQLEnvironment(Environment):
210
  reward=reward,
211
  score=score,
212
  )
 
213
  self._last_obs = obs
214
- return obs
 
 
215
 
216
  @property
217
  def state(self) -> NL2SQLState:
 
152
  self._last_obs = obs
153
  return obs
154
 
155
+ def step(self, action: NL2SQLAction) -> tuple[NL2SQLObservation, float, bool, dict]:
156
  """Execute the agent's SQL and return graded observation."""
157
  if self._task is None or self._example is None:
158
  # Called before reset — auto-reset
 
210
  reward=reward,
211
  score=score,
212
  )
213
+
214
  self._last_obs = obs
215
+ # CRITICAL FIX: Return the standard OpenEnv 4-part tuple!
216
+ info = {"score": score, "error": error}
217
+ return obs, reward, done, info
218
 
219
  @property
220
  def state(self) -> NL2SQLState: