BHaritha/msme-openenv
1
1"""2environment.py - MSME Payment Dispute OpenEnv Environment3Implements reset(), step(), state() per OpenEnv spec.4 5FIX 1: Sequential episode mode — tasks chain together.6 Wrong Task 1 classification bleeds into Task 3 context,7 making this a true multi-step environment.8"""9from .tasks import get_scenario10from .graders import grade_task1, grade_task2, grade_task3, grade_task3_multiturn11 12class MSMEDisputeEnv:13 def __init__(self):14 self._mode = "single" # "single" or "sequential"15 self._task_id = None16 self._scenario = None17 self._done = False18 self._step_count = 019 self._seed = None20 self._last_reward = None21 self._episode_rewards = []22 23 # Sequential mode state24 self._seq_step = 0 # which task we're on (1,2,3)25 self._seq_scenarios = {} # {1: scenario, 2: scenario, 3: scenario}26 self._seq_results = {} # {1: grader_result, 2: grader_result}27 self._agent_t1_label = None # agent's Task1 answer bleeds into Task328 self._agent_t2_facts = None # agent's Task2 answer bleeds into Task329 self._t1_fallback_used = False # True if agent gave invalid Task1 label30 31 # Multi-turn Task 3 state (up to 3 turns: draft → revise → final)32 self._t3_turn = 133 self._t3_scenario = None34 self._t3_done = False35 36 # ─────────────────────────────────────────37 # PUBLIC API38 # ─────────────────────────────────────────39 40 def reset(self, task_id: int = 1, seed: int = None, mode: str = "single") -> dict:41 """42 Start a fresh episode.43 44 mode="single" — one task, one step (original behaviour)45 mode="sequential" — all 3 tasks chained, agent's earlier46 answers affect later task context47 task_id — used in single mode only48 seed — fixed seed for reproducibility49 """50 self._seed = seed51 self._done = False52 self._step_count = 053 self._last_reward = None54 self._episode_rewards = []55 self._agent_t1_label = None56 self._agent_t2_facts = None57 self._seq_results = {}58 self._t1_fallback_used = False59 self._t3_turn = 160 self._t3_done = False61 62 if mode == "sequential":63 self._mode = "sequential"64 self._task_id = None65 self._seq_step = 166 # Pre-load all 3 scenarios with same seed for consistency67 self._seq_scenarios = {68 1: get_scenario(1, seed),69 2: get_scenario(2, seed),70 3: get_scenario(3, seed),71 }72 self._scenario = self._seq_scenarios[1]73 return self._build_obs(task_id=1)74 75 else:76 if task_id not in (1, 2, 3):77 raise ValueError("task_id must be 1, 2, or 3")78 self._mode = "single"79 self._task_id = task_id80 self._scenario = get_scenario(task_id, seed)81 return self._build_obs(task_id=task_id)82 83 def step(self, action: dict) -> dict:84 """85 Submit an action.86 87 Single mode: one step → done.88 Sequential mode: three steps → done after Task 3.89 Agent's Task1 label and Task2 facts are90 injected into the Task3 observation so91 a wrong classification genuinely hurts92 the final letter quality score.93 """94 if self._scenario is None:95 raise RuntimeError("Call reset() before step()")96 if self._done:97 raise RuntimeError("Episode done. Call reset() to start a new one.")98 99 if self._mode == "sequential":100 return self._step_sequential(action)101 else:102 return self._step_single(action)103 104 def state(self) -> dict:105 """Return current environment state."""106 if self._scenario is None:107 return {"status": "not_started"}108 109 base = {110 "mode": self._mode,111 "seed": self._seed,112 "step_count": self._step_count,113 "done": self._done,114 "last_reward": self._last_reward,115 }116 117 if self._mode == "sequential":118 base.update({119 "current_task": self._seq_step if not self._done else "done",120 "episode_rewards": self._seq_results,121 "average_reward": (122 round(sum(r["score"] for r in self._seq_results.values()) /123 len(self._seq_results), 3)124 if self._seq_results else None125 ),126 "observation": self._build_obs(self._seq_step) if not self._done else None,127 })128 else:129 base.update({130 "task_id": self._task_id,131 "scenario_id": self._scenario.get("id"),132 "observation": self._build_obs(self._task_id),133 })134 135 return base136 137 # ─────────────────────────────────────────138 # PRIVATE HELPERS139 # ─────────────────────────────────────────140 141 def _step_single(self, action: dict) -> dict:142 result = self._grade(action, self._task_id, self._scenario)143 self._last_reward = result["score"]144 self._step_count += 1145 146 # Multi-turn: Task 3 allows up to 3 turns (draft → revise → final)147 if self._task_id == 3 and result.get("needs_revision") and self._t3_turn < 3:148 self._t3_turn += 1149 done = False # not done yet — agent should revise150 self._done = False151 else:152 self._t3_done = True153 done = True154 self._done = True155 156 return {157 "state": self.state(),158 "reward": result["score"],159 "done": done,160 "info": result,161 "feedback": result.get("feedback", []),162 "message": result.get("message", "")163 }164 165 def _step_sequential(self, action: dict) -> dict:166 current = self._seq_step167 scenario = self._seq_scenarios[current]168 169 result = self._grade(action, current, scenario)170 self._seq_results[current] = result171 self._last_reward = result["score"]172 self._step_count += 1173 174 # Store agent answers so they bleed into Task 3 context.175 # If Task 1 label is invalid/broken, fall back to "delayed_payment"176 # (most common type) so Task 3 context is always gradeable.177 VALID_LABELS = {"delayed_payment", "partial_payment", "payment_denial"}178 if current == 1:179 raw_label = str(action.get("label", "")).strip().lower()180 self._agent_t1_label = raw_label if raw_label in VALID_LABELS else "delayed_payment"181 if raw_label not in VALID_LABELS:182 # Record that a fallback was used — this is visible in state()183 self._t1_fallback_used = True184 elif current == 2:185 self._agent_t2_facts = action186 187 # Advance or finish188 if current < 3:189 self._seq_step += 1190 self._scenario = self._seq_scenarios[self._seq_step]191 done = False192 else:193 self._done = True194 done = True195 196 avg = round(197 sum(r["score"] for r in self._seq_results.values()) / len(self._seq_results), 3198 )199 200 return {201 "state": self.state(),202 "reward": result["score"],203 "done": done,204 "task_completed": current,205 "next_task": current + 1 if not done else None,206 "episode_avg": avg if done else None,207 "info": result,208 }209 210 def _build_obs(self, task_id: int) -> dict:211 if self._mode == "sequential":212 s = self._seq_scenarios.get(task_id, {})213 else:214 s = self._scenario215 216 if task_id == 1:217 return {218 "task": "classify_dispute",219 "description": "Classify the dispute type from the email.",220 "email": s["email"],221 "valid_labels": ["delayed_payment", "partial_payment", "payment_denial"],222 "action_format": {"label": "<one of the valid_labels>"},223 "note": "Your classification here will affect the context given in Task 3."224 }225 226 elif task_id == 2:227 return {228 "task": "extract_facts",229 "description": "Extract structured facts from the formal notice.",230 "email": s["email"],231 "action_format": {232 "claimant": "<company name>",233 "opponent": "<company name>",234 "amount": "<integer in rupees>",235 "due_date": "<date string>",236 "days_overdue": "<integer>"237 },238 "note": "Your extracted facts will be used as context in Task 3."239 }240 241 elif task_id == 3:242 ctx = dict(s["context"]) # copy243 244 # FIX 1 CORE: Inject agent's earlier answers into Task 3 context.245 # If agent got Task 1 wrong, they get the wrong dispute_type injected.246 # This means a wrong Task 1 genuinely hurts Task 3 letter quality.247 if self._mode == "sequential":248 if self._agent_t1_label:249 ctx["dispute_type_from_agent"] = self._agent_t1_label250 ctx["dispute_type"] = self._agent_t1_label # overwrites ground truth251 if self._agent_t2_facts:252 for field in ["claimant", "opponent", "amount", "due_date", "days_overdue"]:253 if field in self._agent_t2_facts:254 ctx[field] = self._agent_t2_facts[field]255 256 return {257 "task": "draft_demand_letter",258 "description": "Draft a formal MSME payment demand letter using the context below.",259 "context": ctx,260 "action_format": {"letter": "<full demand letter text, minimum 150 words>"},261 "note": "Context was built from your Task 1 and Task 2 answers." if self._mode == "sequential" else ""262 }263 264 def _grade(self, action: dict, task_id: int, scenario: dict) -> dict:265 if task_id == 1:266 return grade_task1(action, {"label": scenario["label"]})267 elif task_id == 2:268 return grade_task2(action, scenario["ground_truth"])269 elif task_id == 3:270 # Build grading scenario with possibly agent-modified context271 grading_scenario = dict(scenario)272 if self._mode == "sequential" and self._agent_t1_label:273 ctx = dict(scenario["context"])274 ctx["dispute_type"] = self._agent_t1_label275 grading_scenario = {**scenario, "context": ctx}276 # Use multi-turn grader — tracks which revision turn we're on277 return grade_task3_multiturn(action, grading_scenario, turn=self._t3_turn)278 