SHUBHAMOS/meta-pytorch-hackathon
0
1"""2SHUBHAMOS: AI Email Operations & Triage Environment3models.py — All typed Pydantic v2 data models4 5These models define the data contract between environment, agents, and graders.6- Email: full internal representation with ground truth labels7- ObservationEmail: agent-facing truncated view (no ground truth)8- Action: structured agent action with type dispatch9- Observation: what the agent receives each step10- State: full internal state (includes ground truth, used by graders)11- StepInfo: per-step metadata12"""13 14from __future__ import annotations15from datetime import datetime16from typing import Literal, Optional, List, Dict, Any17from pydantic import BaseModel, Field18 19# ── Type Aliases ───────────────────────────────────────────────────────────────20 21EmailCategory = Literal["spam", "general_inquiry", "billing_issue", "urgent_complaint", "tech_support"]22Priority = Literal["low", "medium", "high", "unknown"]23Sentiment = Literal["positive", "neutral", "negative"]24ActionType = Literal[25 "classify_email",26 "set_priority",27 "draft_reply",28 "mark_resolved",29 "escalate_email",30 "ignore_email",31]32 33# ── Internal Email Model (full, with ground truth) ─────────────────────────────34 35class Email(BaseModel):36 """Full internal email representation. Ground truth labels are stored here37 and never exposed to the agent via Observation."""38 39 id: str = Field(..., description="Unique email identifier e.g. 'email_001'")40 subject: str = Field(..., description="Email subject line")41 body: str = Field(..., description="Full email body text")42 sender: str = Field(..., description="Sender email address")43 timestamp: datetime = Field(..., description="Arrival timestamp")44 45 # Ground truth labels — hidden from agent, used by graders46 gt_category: EmailCategory = Field(..., description="True category for grading")47 gt_priority: Priority = Field(..., description="True priority for grading (no 'unknown')")48 sentiment: Sentiment = Field(..., description="Body sentiment")49 50 # Agent-assigned labels (start as None / unknown)51 category: Optional[EmailCategory] = Field(None, description="Agent-assigned category")52 priority: Priority = Field("unknown", description="Agent-assigned priority")53 54 # Status flags (mutually exclusive terminal states: resolved/escalated/ignored)55 resolved: bool = Field(False, description="Marked resolved by agent")56 escalated: bool = Field(False, description="Escalated by agent")57 ignored: bool = Field(False, description="Ignored by agent")58 reply_drafted: bool = Field(False, description="Reply has been drafted")59 reply_text: Optional[str] = Field(None, description="Drafted reply text")60 61 def is_processed(self) -> bool:62 """True if email has been terminally acted on (resolved/escalated/ignored)."""63 return self.resolved or self.escalated or self.ignored64 65 66# ── Agent-Facing Email View (no ground truth) ─────────────────────────────────67 68class ObservationEmail(BaseModel):69 """Truncated email view exposed to the agent. Contains no ground truth labels."""70 71 id: str72 subject: str73 sender: str74 timestamp: datetime75 sentiment: Sentiment76 body_preview: str = Field(..., description="First 200 chars of body only")77 category: Optional[EmailCategory] = Field(None, description="Agent-assigned category so far")78 priority: Priority = Field("unknown", description="Agent-assigned priority so far")79 resolved: bool = False80 escalated: bool = False81 ignored: bool = False82 reply_drafted: bool = False83 84 85# ── Action Model ───────────────────────────────────────────────────────────────86 87class Action(BaseModel):88 """Structured agent action. action_type determines which payload fields apply.89 90 Action types and their required fields:91 - classify_email: email_id, category (required)92 - set_priority: email_id, level (required)93 - draft_reply: email_id, text (required)94 - mark_resolved: email_id only95 - escalate_email: email_id only96 - ignore_email: email_id only97 """98 99 action_type: ActionType = Field(..., description="Which action to perform")100 email_id: str = Field(..., description="Target email ID")101 category: Optional[EmailCategory] = Field(None, description="Used by classify_email")102 level: Optional[Priority] = Field(None, description="Used by set_priority")103 text: Optional[str] = Field(None, description="Used by draft_reply")104 105 106# ── Observation Model (returned by step/reset) ────────────────────────────────107 108class Observation(BaseModel):109 """What the agent sees after each step or on reset. No ground truth."""110 111 emails: List[ObservationEmail] = Field(..., description="Current inbox state")112 total_emails: int = Field(..., description="Total emails in this episode")113 pending_count: int = Field(..., description="Emails not yet terminally processed")114 resolved_count: int = Field(..., description="Emails successfully resolved")115 escalated_count: int = Field(..., description="Emails escalated")116 ignored_count: int = Field(..., description="Emails ignored")117 step_count: int = Field(..., description="Steps taken so far (0-indexed at reset)")118 max_steps: int = Field(..., description="Maximum steps before forced termination")119 elapsed_ratio: float = Field(..., description="step_count / max_steps, 0.0 to 1.0")120 121 122# ── Internal State (complete, used by graders) ────────────────────────────────123 124class State(BaseModel):125 """Complete internal environment state including ground truth. Graders use this."""126 127 emails: List[Email] = Field(..., description="All emails with ground truth labels")128 action_history: List[Dict[str, Any]] = Field(129 default_factory=list,130 description="Ordered list of all actions taken this episode"131 )132 step_count: int = Field(0, description="Steps taken in current episode")133 max_steps: int = Field(0, description="Max steps for current task")134 done: bool = Field(False, description="Whether episode has terminated")135 total_reward: float = Field(0.0, description="Cumulative reward this episode")136 seed: int = Field(0, description="Random seed used for this episode")137 task_id: str = Field("", description="Task identifier: easy / medium / hard")138 139 140# ── Per-Step Info ─────────────────────────────────────────────────────────────141 142class StepInfo(BaseModel):143 """Structured info returned in the info dict from step()."""144 145 action_type: str = Field(..., description="Action that was applied")146 email_id: str = Field(..., description="Email that was targeted")147 valid: bool = Field(..., description="Whether the action was valid")148 error: Optional[str] = Field(None, description="Error message if invalid")149 step_reward: float = Field(0.0, description="Reward earned this step")150 done: bool = Field(False, description="Whether episode ended after this step")151 