sc-likes-to-code/openenv-customer-support-env
0
1from server.tasks import load_task2from server.grader import evaluate_action3from models import Observation, Ticket, Reward4 5 6# Max steps per difficulty level7TASK_MAX_STEPS = {8 "easy": 6,9 "medium": 6,10 "hard": 8,11}12 13 14class SupportEnv:15 def __init__(self):16 self.state_data = None17 self.current_task = None18 self.task_level = "easy"19 20 # ── reset ────────────────────────────────────────────────────────────────21 def reset(self, task: str = "easy") -> Observation:22 self.task_level = task23 self.current_task = load_task(task)24 self.state_data = {25 "history": [],26 "step_count": 0,27 "asked_info": False,28 "max_steps": TASK_MAX_STEPS.get(task, 8),29 "task_level": task,30 }31 return self._get_observation()32 33 # ── step ─────────────────────────────────────────────────────────────────34 def step(self, action):35 # safety: auto-init if called before reset36 if self.state_data is None or self.current_task is None:37 self.reset(self.task_level)38 39 self.state_data["step_count"] += 140 step = self.state_data["step_count"]41 42 # grade the action43 reward: Reward = evaluate_action(self.current_task, action, self.state_data)44 reward.score = round(max(min(float(reward.score), 0.99), 0.01), 2)45 46 # update memory flags47 if action.action_type == "ask":48 self.state_data["asked_info"] = True49 50 # store full history entry (action_type needed by grader repeat-check)51 self.state_data["history"].append({52 "user": self.current_task["tickets"][0]["text"],53 "agent": action.content or "",54 "action_type": action.action_type,55 })56 57 done = self._is_done(action, reward, step)58 59 return self._get_observation(), reward, done, {}60 61 # ── state ────────────────────────────────────────────────────────────────62 def state(self) -> dict:63 if self.state_data is None:64 return {"status": "not_initialized"}65 return {66 "task_level": self.state_data.get("task_level"),67 "step_count": self.state_data.get("step_count"),68 "max_steps": self.state_data.get("max_steps"),69 "asked_info": self.state_data.get("asked_info"),70 "history_len": len(self.state_data.get("history", [])),71 "history": self.state_data.get("history", []),72 }73 74 # ── close ────────────────────────────────────────────────────────────────75 def close(self):76 """Cleanup hook — called by inference.py after episode ends."""77 self.state_data = None78 self.current_task = None79 80 # ── internal helpers ─────────────────────────────────────────────────────81 def _is_done(self, action, reward: Reward, step: int) -> bool:82 max_steps = self.state_data.get("max_steps", 8)83 expected = self.current_task.get("expected", {})84 85 # always end if max steps reached86 if step >= max_steps:87 return True88 89 # always end on very high reward (near-perfect episode)90 if reward.score >= 0.95:91 return True92 93 # EASY / MEDIUM: done after a valid respond action94 if self.task_level in ("easy", "medium"):95 if action.action_type == "respond" and reward.score >= 0.4:96 return True97 98 # HARD: done only after respond/escalate AND agent has asked for info99 elif self.task_level == "hard":100 needs_info = expected.get("needs_info", False)101 if action.action_type in ("respond", "escalate"):102 if not needs_info or self.state_data.get("asked_info"):103 return True104 105 return False106 107 def _get_observation(self) -> Observation:108 tickets = [Ticket(**t) for t in self.current_task["tickets"]]109 return Observation(110 tickets=tickets,111 current_ticket_id=tickets[0].id if tickets else None,112 history=self.state_data["history"],113 )