Aditya Guntur commited on
Commit
e48dce9
·
1 Parent(s): 42466e3

Fix: inherit from openenv Environment base class to resolve reset_async/concurrency errors

Browse files
Files changed (1) hide show
  1. server/pm_ops_environment.py +29 -15
server/pm_ops_environment.py CHANGED
@@ -1,9 +1,9 @@
1
  import copy
2
  import random
3
- import uuid
4
  from typing import Any, Dict, Optional
5
 
6
- from models import PMOpsAction, PMOpsObservation
 
7
  from server.apps.ticketing import TicketingApp
8
  from server.apps.codebase import CodebaseApp
9
  from server.apps.chat import ChatApp
@@ -34,10 +34,11 @@ def _oracle_check(scenario: Dict[str, Any]) -> bool:
34
  return bool(channel) and (bool(team) or scenario["type"] == "release_notes")
35
 
36
 
37
- class PMOpsEnvironment:
38
  SUPPORTS_CONCURRENT_SESSIONS = True
39
 
40
  def __init__(self):
 
41
  self._ticketing: Optional[TicketingApp] = None
42
  self._codebase: Optional[CodebaseApp] = None
43
  self._chat: Optional[ChatApp] = None
@@ -46,8 +47,9 @@ class PMOpsEnvironment:
46
  self._step_count: int = 0
47
  self._done: bool = False
48
 
49
- def reset(self) -> PMOpsObservation:
50
- seed = random.randint(0, 2 ** 31)
 
51
  rng = random.Random(seed)
52
  difficulty = rng.choice(_DIFFICULTY_POOL)
53
  task_type = rng.choice(_TASK_TYPES)
@@ -84,7 +86,7 @@ class PMOpsEnvironment:
84
  done=False,
85
  )
86
 
87
- def step(self, action: PMOpsAction) -> PMOpsObservation:
88
  if self._ticketing is None:
89
  raise RuntimeError("Call reset() before step()")
90
 
@@ -129,6 +131,18 @@ class PMOpsEnvironment:
129
  done=done,
130
  )
131
 
 
 
 
 
 
 
 
 
 
 
 
 
132
  def _dispatch(self, action: PMOpsAction) -> Dict[str, Any]:
133
  at = action.action_type
134
  args = action.args or {}
@@ -156,12 +170,12 @@ class PMOpsEnvironment:
156
  if at.startswith("ticketing."):
157
  op = at.split(".", 1)[1]
158
  handlers = {
159
- "create_ticket": self._ticketing.create_ticket,
160
- "update_ticket": self._ticketing.update_ticket,
161
- "get_ticket": self._ticketing.get_ticket,
162
- "list_tickets": self._ticketing.list_tickets,
163
- "assign_ticket": self._ticketing.assign_ticket,
164
- "comment_ticket": self._ticketing.comment_ticket,
165
  "transition_ticket": self._ticketing.transition_ticket,
166
  }
167
  if op not in handlers:
@@ -182,10 +196,10 @@ class PMOpsEnvironment:
182
  if at.startswith("chat."):
183
  op = at.split(".", 1)[1]
184
  handlers = {
185
- "post_message": self._chat.post_message,
186
- "read_channel": self._chat.read_channel,
187
  "list_channels": self._chat.list_channels,
188
- "search": self._chat.search,
189
  }
190
  if op not in handlers:
191
  return {"ok": False, "error": f"Unknown chat action: {op}"}
 
1
  import copy
2
  import random
 
3
  from typing import Any, Dict, Optional
4
 
5
+ from openenv.core.env_server.interfaces import Environment
6
+ from models import PMOpsAction, PMOpsObservation, PMOpsState
7
  from server.apps.ticketing import TicketingApp
8
  from server.apps.codebase import CodebaseApp
9
  from server.apps.chat import ChatApp
 
34
  return bool(channel) and (bool(team) or scenario["type"] == "release_notes")
35
 
36
 
37
+ class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]):
38
  SUPPORTS_CONCURRENT_SESSIONS = True
39
 
40
  def __init__(self):
41
+ super().__init__()
42
  self._ticketing: Optional[TicketingApp] = None
43
  self._codebase: Optional[CodebaseApp] = None
44
  self._chat: Optional[ChatApp] = None
 
47
  self._step_count: int = 0
48
  self._done: bool = False
49
 
50
+ def reset(self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs) -> PMOpsObservation:
51
+ if seed is None:
52
+ seed = random.randint(0, 2 ** 31)
53
  rng = random.Random(seed)
54
  difficulty = rng.choice(_DIFFICULTY_POOL)
55
  task_type = rng.choice(_TASK_TYPES)
 
86
  done=False,
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
 
 
131
  done=done,
132
  )
133
 
134
+ @property
135
+ def state(self) -> PMOpsState:
136
+ return PMOpsState(
137
+ org_config=self._org_config or {},
138
+ task_config=self._scenario or {},
139
+ ticketing=self._ticketing.snapshot() if self._ticketing else {},
140
+ chat=self._chat.snapshot() if self._chat else {},
141
+ codebase={},
142
+ step_count=self._step_count,
143
+ finished=self._done,
144
+ )
145
+
146
  def _dispatch(self, action: PMOpsAction) -> Dict[str, Any]:
147
  at = action.action_type
148
  args = action.args or {}
 
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:
 
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}"}