pm-ops / server /pm_ops_environment.py
Aditya Guntur
Fix: inherit from openenv Environment base class to resolve reset_async/concurrency errors
e48dce9
Raw
History Blame Contribute Delete
8.49 kB
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,
)
@property
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