SavK1 commited on
Commit
74d0de6
·
1 Parent(s): 860d7e4

made minor type check changes, open env has strict type checking

Browse files
Files changed (3) hide show
  1. client.py +3 -4
  2. models.py +1 -1
  3. server/pm_ops_environment.py +34 -17
client.py CHANGED
@@ -1,6 +1,6 @@
1
  """OpenEnv client wrapper for PM-Ops environment."""
2
  import os
3
- from openenv.core.client import EnvClient
4
  from models import PMOpsAction, PMOpsObservation
5
 
6
 
@@ -8,7 +8,6 @@ class PMOpsEnv(EnvClient):
8
  action_class = PMOpsAction
9
  observation_class = PMOpsObservation
10
 
11
- def __init__(self, base_url: str = None, token: str = None):
12
  url = base_url or os.getenv("API_BASE_URL", "https://adityaguntur-pm-ops.hf.space")
13
- tok = token or os.getenv("HF_TOKEN")
14
- super().__init__(base_url=url, token=tok)
 
1
  """OpenEnv client wrapper for PM-Ops environment."""
2
  import os
3
+ from openenv.core.env_client import EnvClient
4
  from models import PMOpsAction, PMOpsObservation
5
 
6
 
 
8
  action_class = PMOpsAction
9
  observation_class = PMOpsObservation
10
 
11
+ def __init__(self, base_url: str | None = None, token: str | None= None):
12
  url = base_url or os.getenv("API_BASE_URL", "https://adityaguntur-pm-ops.hf.space")
13
+ super().__init__(base_url=url)
 
models.py CHANGED
@@ -46,4 +46,4 @@ class PMOpsState(State):
46
  codebase: Dict[str, Any] = Field(default_factory=dict)
47
  step_count: int = Field(default=0)
48
  finished: bool = Field(default=False)
49
- episode_id: str = Field(default="")
 
46
  codebase: Dict[str, Any] = Field(default_factory=dict)
47
  step_count: int = Field(default=0)
48
  finished: bool = Field(default=False)
49
+ episode_id: str | None = Field(default="")
server/pm_ops_environment.py CHANGED
@@ -54,12 +54,18 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
54
  difficulty = rng.choice(_DIFFICULTY_POOL)
55
  task_type = rng.choice(_TASK_TYPES)
56
 
 
 
 
57
  for attempt in range(10):
58
  org = generate_org_config(seed + attempt, difficulty)
59
  scenario = generate_scenario(task_type, org, seed + attempt)
60
  if _oracle_check(scenario):
61
  break
62
 
 
 
 
63
  channels = list(org["oncall_channels"].values())
64
  noise = org.get("noise_channels", [])
65
 
@@ -87,14 +93,16 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
87
  )
88
 
89
  def step(self, action: PMOpsAction, timeout_s: Optional[float] = None, **kwargs) -> PMOpsObservation:
90
- if self._ticketing is None:
91
  raise RuntimeError("Call reset() before step()")
92
 
 
 
93
  if self._done:
94
  return PMOpsObservation(
95
  step=self._step_count,
96
  max_steps=MAX_STEPS,
97
- task_brief=self._scenario["brief"],
98
  last_action_result={"ok": False, "error": "Episode already finished"},
99
  app_state_deltas={"ticketing": [], "chat": [], "codebase": []},
100
  steps_remaining=0,
@@ -122,7 +130,7 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
122
  return PMOpsObservation(
123
  step=self._step_count,
124
  max_steps=MAX_STEPS,
125
- task_brief=self._scenario["brief"],
126
  last_action_result=result,
127
  app_state_deltas=deltas,
128
  steps_remaining=max(0, MAX_STEPS - self._step_count),
@@ -147,6 +155,12 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
147
  at = action.action_type
148
  args = action.args or {}
149
 
 
 
 
 
 
 
150
  if at == "meta.noop":
151
  return {"ok": True, "data": "No operation."}
152
 
@@ -170,13 +184,13 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
170
  if at.startswith("ticketing."):
171
  op = at.split(".", 1)[1]
172
  handlers = {
173
- "create_ticket": self._ticketing.create_ticket,
174
- "update_ticket": self._ticketing.update_ticket,
175
- "get_ticket": self._ticketing.get_ticket,
176
- "list_tickets": self._ticketing.list_tickets,
177
- "assign_ticket": self._ticketing.assign_ticket,
178
- "comment_ticket": self._ticketing.comment_ticket,
179
- "transition_ticket": self._ticketing.transition_ticket,
180
  }
181
  if op not in handlers:
182
  return {"ok": False, "error": f"Unknown ticketing action: {op}"}
@@ -185,9 +199,9 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
185
  if at.startswith("codebase."):
186
  op = at.split(".", 1)[1]
187
  handlers = {
188
- "list_commits": self._codebase.list_commits,
189
- "get_commit": self._codebase.get_commit,
190
- "list_prs": self._codebase.list_prs,
191
  }
192
  if op not in handlers:
193
  return {"ok": False, "error": f"Unknown codebase action: {op}"}
@@ -196,10 +210,10 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
196
  if at.startswith("chat."):
197
  op = at.split(".", 1)[1]
198
  handlers = {
199
- "post_message": self._chat.post_message,
200
- "read_channel": self._chat.read_channel,
201
- "list_channels": self._chat.list_channels,
202
- "search": self._chat.search,
203
  }
204
  if op not in handlers:
205
  return {"ok": False, "error": f"Unknown chat action: {op}"}
@@ -208,6 +222,9 @@ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
208
  return {"ok": False, "error": f"Unknown action_type: {at}"}
209
 
210
  def _grade(self) -> float:
 
 
 
211
  final_state = {
212
  "ticketing": self._ticketing.snapshot(),
213
  "chat": self._chat.snapshot(),
 
54
  difficulty = rng.choice(_DIFFICULTY_POOL)
55
  task_type = rng.choice(_TASK_TYPES)
56
 
57
+ org: Optional[Dict[str, Any]] = None
58
+ scenario: Optional[Dict[str, Any]] = None
59
+
60
  for attempt in range(10):
61
  org = generate_org_config(seed + attempt, difficulty)
62
  scenario = generate_scenario(task_type, org, seed + attempt)
63
  if _oracle_check(scenario):
64
  break
65
 
66
+ if org is None or scenario is None:
67
+ raise RuntimeError("Failed to initialize episode state")
68
+
69
  channels = list(org["oncall_channels"].values())
70
  noise = org.get("noise_channels", [])
71
 
 
93
  )
94
 
95
  def step(self, action: PMOpsAction, timeout_s: Optional[float] = None, **kwargs) -> PMOpsObservation:
96
+ if self._ticketing is None or self._codebase is None or self._chat is None or self._scenario is None:
97
  raise RuntimeError("Call reset() before step()")
98
 
99
+ scenario = self._scenario
100
+
101
  if self._done:
102
  return PMOpsObservation(
103
  step=self._step_count,
104
  max_steps=MAX_STEPS,
105
+ task_brief=scenario["brief"],
106
  last_action_result={"ok": False, "error": "Episode already finished"},
107
  app_state_deltas={"ticketing": [], "chat": [], "codebase": []},
108
  steps_remaining=0,
 
130
  return PMOpsObservation(
131
  step=self._step_count,
132
  max_steps=MAX_STEPS,
133
+ task_brief=scenario["brief"],
134
  last_action_result=result,
135
  app_state_deltas=deltas,
136
  steps_remaining=max(0, MAX_STEPS - self._step_count),
 
155
  at = action.action_type
156
  args = action.args or {}
157
 
158
+ ticketing = self._ticketing
159
+ codebase = self._codebase
160
+ chat = self._chat
161
+ if ticketing is None or codebase is None or chat is None:
162
+ return {"ok": False, "error": "Environment not initialized. Call reset() before step()."}
163
+
164
  if at == "meta.noop":
165
  return {"ok": True, "data": "No operation."}
166
 
 
184
  if at.startswith("ticketing."):
185
  op = at.split(".", 1)[1]
186
  handlers = {
187
+ "create_ticket": ticketing.create_ticket,
188
+ "update_ticket": ticketing.update_ticket,
189
+ "get_ticket": ticketing.get_ticket,
190
+ "list_tickets": ticketing.list_tickets,
191
+ "assign_ticket": ticketing.assign_ticket,
192
+ "comment_ticket": ticketing.comment_ticket,
193
+ "transition_ticket": ticketing.transition_ticket,
194
  }
195
  if op not in handlers:
196
  return {"ok": False, "error": f"Unknown ticketing action: {op}"}
 
199
  if at.startswith("codebase."):
200
  op = at.split(".", 1)[1]
201
  handlers = {
202
+ "list_commits": codebase.list_commits,
203
+ "get_commit": codebase.get_commit,
204
+ "list_prs": codebase.list_prs,
205
  }
206
  if op not in handlers:
207
  return {"ok": False, "error": f"Unknown codebase action: {op}"}
 
210
  if at.startswith("chat."):
211
  op = at.split(".", 1)[1]
212
  handlers = {
213
+ "post_message": chat.post_message,
214
+ "read_channel": chat.read_channel,
215
+ "list_channels": chat.list_channels,
216
+ "search": chat.search,
217
  }
218
  if op not in handlers:
219
  return {"ok": False, "error": f"Unknown chat action: {op}"}
 
222
  return {"ok": False, "error": f"Unknown action_type: {at}"}
223
 
224
  def _grade(self) -> float:
225
+ if self._ticketing is None or self._chat is None or self._org_config is None or self._scenario is None:
226
+ return 0.0
227
+
228
  final_state = {
229
  "ticketing": self._ticketing.snapshot(),
230
  "chat": self._chat.snapshot(),