CoolFace
Apppublic

sc-likes-to-code/openenv-customer-support-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
your_environment.py113 linesDownload Raw Back to server
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        )