BHaritha/msme-openenv
1
1"""2models.py - Typed Pydantic models for OpenEnv spec compliance.3Observation, Action, Reward — all strictly typed.4"""5from pydantic import BaseModel, Field6from typing import Optional, Dict, Any, List, Literal7 8# ── Reward ────────────────────────────────────9class MSMEReward(BaseModel):10 score: float = Field(..., gt=0.0001, lt=0.999, description="Reward in (0.0001, 0.999)")11 breakdown: Dict[str, Any] = Field(default_factory=dict)12 reason: str = ""13 14# ── Actions ───────────────────────────────────15class Task1Action(BaseModel):16 label: Literal["delayed_payment", "partial_payment", "payment_denial"]17 18class Task2Action(BaseModel):19 claimant: str20 opponent: str21 amount: int = Field(..., gt=0)22 due_date: str23 days_overdue: int = Field(..., ge=0)24 25class Task3Action(BaseModel):26 letter: str = Field(27 ...,28 min_length=500,29 description="Full formal demand letter text; target at least 150 words.",30 )31 32class MSMEAction(BaseModel):33 """Union action — agent fills whichever field matches current task."""34 label: Optional[str] = None # Task 135 claimant: Optional[str] = None # Task 236 opponent: Optional[str] = None # Task 237 amount: Optional[int] = None # Task 238 due_date: Optional[str] = None # Task 239 days_overdue: Optional[int] = None # Task 240 letter: Optional[str] = None # Task 341 42 def to_dict(self) -> dict:43 return {k: v for k, v in self.model_dump().items() if v is not None}44 45# ── Observations ─────────────────────────────46class EmailObs(BaseModel):47 subject: str48 body: str49 50class Task1Observation(BaseModel):51 task: Literal["classify_dispute"]52 description: str53 email: EmailObs54 valid_labels: List[str]55 action_format: Dict[str, str]56 note: str = ""57 58class Task2Observation(BaseModel):59 task: Literal["extract_facts"]60 description: str61 email: EmailObs62 action_format: Dict[str, str]63 note: str = ""64 65class DisputeContext(BaseModel):66 claimant: str67 opponent: str68 amount: int69 invoice_no: str70 invoice_date: str71 due_date: str72 days_overdue: int73 dispute_type: str74 evidence: List[str]75 76class Task3Observation(BaseModel):77 task: Literal["draft_demand_letter"]78 description: str79 context: Dict[str, Any]80 action_format: Dict[str, str]81 note: str = ""82 83# ── Episode state ─────────────────────────────84class EpisodeState(BaseModel):85 mode: str86 seed: Optional[int]87 step_count: int88 done: bool89 last_reward: Optional[float]90 task_id: Optional[int] = None91 scenario_id: Optional[str] = None92 current_task: Optional[Any] = None93 episode_rewards: Dict[str, Any] = Field(default_factory=dict)94 average_reward: Optional[float] = None95 observation: Optional[Dict[str, Any]] = None96 97# ── API request/response models ───────────────98class ResetRequest(BaseModel):99 task_id: int = Field(1, ge=1, le=3)100 seed: Optional[int] = None101 mode: Literal["single", "sequential"] = "single"102 103class StepRequest(BaseModel):104 action: Dict[str, Any]105 106class ResetResponse(BaseModel):107 observation: Dict[str, Any]108 state: Dict[str, Any]109 110class StepResponse(BaseModel):111 state: Dict[str, Any]112 reward: float = Field(..., gt=0.001, lt=0.999)113 done: bool114 info: Dict[str, Any]115 task_completed: Optional[int] = None116 next_task: Optional[int] = None117 episode_avg: Optional[float] = None118 