Spaces:
Sleeping
Sleeping
Aditya Guntur
Fix: inherit from openenv Environment base class to resolve reset_async/concurrency errors
e48dce9 | import copy | |
| import random | |
| from typing import Any, Dict, Optional | |
| from openenv.core.env_server.interfaces import Environment | |
| from models import PMOpsAction, PMOpsObservation, PMOpsState | |
| from server.apps.ticketing import TicketingApp | |
| from server.apps.codebase import CodebaseApp | |
| from server.apps.chat import ChatApp | |
| from server.world.org_generator import generate_org_config | |
| from server.world.scenario_gen import generate_scenario | |
| from server.tasks.triage_task import TriageTask | |
| from server.tasks.incident_routing_task import IncidentRoutingTask | |
| from server.tasks.release_notes_task import ReleaseNotesTask | |
| from server.tasks.dep_update_task import DepUpdateTask | |
| MAX_STEPS = 40 | |
| _TASK_TYPES = ["triage", "incident_routing", "release_notes", "dep_update"] | |
| _DIFFICULTY_POOL = ["easy", "medium", "medium", "hard"] | |
| _TASK_GRADERS = { | |
| "triage": TriageTask(), | |
| "incident_routing": IncidentRoutingTask(), | |
| "release_notes": ReleaseNotesTask(), | |
| "dep_update": DepUpdateTask(), | |
| } | |
| def _oracle_check(scenario: Dict[str, Any]) -> bool: | |
| expected = scenario.get("expected", {}) | |
| if not expected: | |
| return False | |
| channel = expected.get("channel") or expected.get("channels") | |
| team = expected.get("team") or expected.get("teams_to_notify") | |
| return bool(channel) and (bool(team) or scenario["type"] == "release_notes") | |
| class PMOpsEnvironment(Environment[PMOpsAction, PMOpsObservation, PMOpsState]): | |
| SUPPORTS_CONCURRENT_SESSIONS = True | |
| def __init__(self): | |
| super().__init__() | |
| self._ticketing: Optional[TicketingApp] = None | |
| self._codebase: Optional[CodebaseApp] = None | |
| self._chat: Optional[ChatApp] = None | |
| self._org_config: Optional[Dict[str, Any]] = None | |
| self._scenario: Optional[Dict[str, Any]] = None | |
| self._step_count: int = 0 | |
| self._done: bool = False | |
| def reset(self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs) -> PMOpsObservation: | |
| if seed is None: | |
| seed = random.randint(0, 2 ** 31) | |
| rng = random.Random(seed) | |
| difficulty = rng.choice(_DIFFICULTY_POOL) | |
| task_type = rng.choice(_TASK_TYPES) | |
| for attempt in range(10): | |
| org = generate_org_config(seed + attempt, difficulty) | |
| scenario = generate_scenario(task_type, org, seed + attempt) | |
| if _oracle_check(scenario): | |
| break | |
| channels = list(org["oncall_channels"].values()) | |
| noise = org.get("noise_channels", []) | |
| self._org_config = org | |
| self._scenario = scenario | |
| self._ticketing = TicketingApp(org) | |
| self._codebase = CodebaseApp(seed, org["services"]) | |
| self._chat = ChatApp(channels, noise) | |
| self._step_count = 0 | |
| self._done = False | |
| return PMOpsObservation( | |
| step=0, | |
| max_steps=MAX_STEPS, | |
| task_brief=scenario["brief"], | |
| last_action_result={ | |
| "ok": True, | |
| "data": "Environment ready. Call meta.read_runbook to learn this org's conventions.", | |
| }, | |
| app_state_deltas={"ticketing": [], "chat": [], "codebase": []}, | |
| steps_remaining=MAX_STEPS, | |
| token_budget_remaining=12000, | |
| reward=0.0, | |
| done=False, | |
| ) | |
| def step(self, action: PMOpsAction, timeout_s: Optional[float] = None, **kwargs) -> PMOpsObservation: | |
| if self._ticketing is None: | |
| raise RuntimeError("Call reset() before step()") | |
| if self._done: | |
| return PMOpsObservation( | |
| step=self._step_count, | |
| max_steps=MAX_STEPS, | |
| task_brief=self._scenario["brief"], | |
| last_action_result={"ok": False, "error": "Episode already finished"}, | |
| app_state_deltas={"ticketing": [], "chat": [], "codebase": []}, | |
| steps_remaining=0, | |
| token_budget_remaining=0, | |
| reward=0.0, | |
| done=True, | |
| ) | |
| self._step_count += 1 | |
| result = self._dispatch(action) | |
| done = action.action_type == "meta.finish" or self._step_count >= MAX_STEPS | |
| reward = 0.0 | |
| if done: | |
| reward = self._grade() | |
| self._done = True | |
| side_effects = result.get("side_effects", []) | |
| deltas = { | |
| "ticketing": side_effects if any("ticket" in s for s in side_effects) else [], | |
| "chat": side_effects if any("message" in s for s in side_effects) else [], | |
| "codebase": [], | |
| } | |
| return PMOpsObservation( | |
| step=self._step_count, | |
| max_steps=MAX_STEPS, | |
| task_brief=self._scenario["brief"], | |
| last_action_result=result, | |
| app_state_deltas=deltas, | |
| steps_remaining=max(0, MAX_STEPS - self._step_count), | |
| token_budget_remaining=max(0, 12000 - self._step_count * 300), | |
| reward=reward, | |
| done=done, | |
| ) | |
| def state(self) -> PMOpsState: | |
| return PMOpsState( | |
| org_config=self._org_config or {}, | |
| task_config=self._scenario or {}, | |
| ticketing=self._ticketing.snapshot() if self._ticketing else {}, | |
| chat=self._chat.snapshot() if self._chat else {}, | |
| codebase={}, | |
| step_count=self._step_count, | |
| finished=self._done, | |
| ) | |
| def _dispatch(self, action: PMOpsAction) -> Dict[str, Any]: | |
| at = action.action_type | |
| args = action.args or {} | |
| if at == "meta.noop": | |
| return {"ok": True, "data": "No operation."} | |
| if at == "meta.read_runbook": | |
| return { | |
| "ok": True, | |
| "data": { | |
| "org_config": copy.deepcopy(self._org_config), | |
| "hint": ( | |
| "Use label_taxonomy for valid ticket labels, " | |
| "priority_levels for valid priorities, " | |
| "team_map[service] to find the owning team, " | |
| "oncall_channels[service] to find the channel to notify." | |
| ), | |
| }, | |
| } | |
| if at == "meta.finish": | |
| return {"ok": True, "data": "Episode finishing. Score will be computed."} | |
| if at.startswith("ticketing."): | |
| op = at.split(".", 1)[1] | |
| handlers = { | |
| "create_ticket": self._ticketing.create_ticket, | |
| "update_ticket": self._ticketing.update_ticket, | |
| "get_ticket": self._ticketing.get_ticket, | |
| "list_tickets": self._ticketing.list_tickets, | |
| "assign_ticket": self._ticketing.assign_ticket, | |
| "comment_ticket": self._ticketing.comment_ticket, | |
| "transition_ticket": self._ticketing.transition_ticket, | |
| } | |
| if op not in handlers: | |
| return {"ok": False, "error": f"Unknown ticketing action: {op}"} | |
| return handlers[op](args) | |
| if at.startswith("codebase."): | |
| op = at.split(".", 1)[1] | |
| handlers = { | |
| "list_commits": self._codebase.list_commits, | |
| "get_commit": self._codebase.get_commit, | |
| "list_prs": self._codebase.list_prs, | |
| } | |
| if op not in handlers: | |
| return {"ok": False, "error": f"Unknown codebase action: {op}"} | |
| return handlers[op](args) | |
| if at.startswith("chat."): | |
| op = at.split(".", 1)[1] | |
| handlers = { | |
| "post_message": self._chat.post_message, | |
| "read_channel": self._chat.read_channel, | |
| "list_channels": self._chat.list_channels, | |
| "search": self._chat.search, | |
| } | |
| if op not in handlers: | |
| return {"ok": False, "error": f"Unknown chat action: {op}"} | |
| return handlers[op](args) | |
| return {"ok": False, "error": f"Unknown action_type: {at}"} | |
| def _grade(self) -> float: | |
| final_state = { | |
| "ticketing": self._ticketing.snapshot(), | |
| "chat": self._chat.snapshot(), | |
| } | |
| grader = _TASK_GRADERS.get(self._scenario["type"]) | |
| if not grader: | |
| return 0.0 | |
| return grader.grade(final_state, self._org_config, self._scenario) | |
| def close(self) -> None: | |
| self._ticketing = None | |
| self._codebase = None | |
| self._chat = None | |