Team Ai
Apppublic

SHUBHAMOS/meta-pytorch-hackathon

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
environment.py416 linesDownload Raw Back to server
1"""2SHUBHAMOS: AI Email Operations & Triage Environment3environment.py — EmailTriageEnv OpenEnv-compliant environment class4 5Implements:6  reset(task_config) → Observation7  step(action)       → (Observation, float, bool, dict)8  state()            → State9"""10 11from __future__ import annotations12import random13from datetime import datetime, timezone14from typing import Any, Dict, List, Optional, Tuple15 16from .models import (17    Action,18    ActionType,19    Email,20    EmailCategory,21    Observation,22    ObservationEmail,23    Priority,24    Sentiment,25    State,26    StepInfo,27)28 29# ── Seed data for realistic email generation ────────────────────────────────────30 31_SENDERS = [32    "alice.johnson@acmecorp.com",33    "bob.smith@enterprise.io",34    "support@bigclient.net",35    "billing@financeops.com",36    "ceo@urgentco.org",37    "info@spamblaster.biz",38    "noreply@marketing.cloud",39    "james.t@retailpartner.com",40    "helpdesk@internal.corp",41    "carol.w@medicalgroup.org",42]43 44_SUBJECTS_BY_CATEGORY: Dict[EmailCategory, List[str]] = {45    "spam": [46        "You've won a $500 gift card!",47        "URGENT: Claim your prize now",48        "Free iPhone — limited offer",49        "Make money from home — guaranteed",50        "Your account has been selected!",51    ],52    "general_inquiry": [53        "Question about your services",54        "Request for product information",55        "Partnership inquiry",56        "Media inquiry regarding your company",57        "How do I get started?",58        "Can you help me understand your pricing?",59    ],60    "billing_issue": [61        "Invoice discrepancy — please review",62        "Double charged on my account",63        "Refund request for order #44821",64        "Payment failed — what do I do?",65        "Billing error on last statement",66        "I was overcharged this month",67    ],68    "urgent_complaint": [69        "URGENT: Service outage affecting my business",70        "Complete system failure — need immediate help",71        "Critical bug causing data loss",72        "Unacceptable service — legal action pending",73        "Production down — customers impacted",74    ],75}76 77_BODIES_BY_CATEGORY: Dict[EmailCategory, List[str]] = {78    "spam": [79        "Congratulations! You have been selected to receive a $500 gift card. Click here to claim your reward before it expires in 24 hours. No purchase necessary. Act now!",80        "We have a special offer just for you! Make $5000/week from the comfort of your home. No experience needed. Join thousands of successful members. Click to get started today!",81        "LIMITED TIME OFFER: You have been pre-approved for a FREE iPhone 15 Pro. Simply complete a short survey and the device will be shipped to your door. Claim now!",82        "Your email was randomly selected for our VIP program. As a VIP member you get access to exclusive deals, cash prizes, and luxury vacations. Join for free today!",83    ],84    "general_inquiry": [85        "Hello, I came across your company online and I'm interested in learning more about your product offerings. Could you please send me a product catalog and pricing information? I'm particularly interested in your enterprise solutions. Thank you for your time.",86        "Hi there, I'm evaluating potential vendors for our upcoming project and your company came highly recommended. I'd love to schedule a 30-minute call to discuss how you might be able to help us. What's your availability this week?",87        "I'm a journalist writing an article about innovation in your industry. I would love to get a quote from your team about recent developments. Could you connect me with your PR contact? Deadline is Friday.",88        "We are a small business looking to expand our operations and are interested in your services. Before we proceed, I have a few questions about your support model and SLA commitments. Can someone from your sales team reach out?",89        "Hello, I recently attended your webinar on digital transformation and had a few follow-up questions. Could you share the recording and slide deck? Also, I'd love to speak with a product specialist about implementation timelines.",90    ],91    "billing_issue": [92        "I received my invoice for last month and I believe there is an error. I was charged $1,200 but my contract clearly states the monthly fee is $950. Can you please review and send a corrected invoice as soon as possible? Invoice #INV-2024-0892.",93        "I was double charged on March 15th. My bank statement shows two transactions of $299 from your company. Please refund one of these charges immediately. I can provide bank statement screenshots if needed. This needs to be resolved urgently.",94        "I would like to request a refund for my recent purchase (Order #44821). I was charged for a premium plan but I never actually used the service. Per your refund policy, I should be eligible for a full refund within 30 days. Please process this ASAP.",95        "My credit card payment failed but I was still charged. I now see a pending charge of $450 on my card but I never received a confirmation. Can you check your system and either confirm the payment was received or release the hold?",96        "There is a discrepancy on my account statement. I was on the $99/month plan but was charged $149 this cycle. I did not authorize any plan upgrade. Please revert this change and refund the difference of $50.",97    ],98    "urgent_complaint": [99        "Our production environment has been down for the past 2 hours and your support team has been completely unresponsive. This is causing significant financial losses for our company — approximately $10,000 per hour. We have 500+ customers impacted and I need an immediate response. If this isn't resolved in the next hour we will be forced to seek an emergency injunction.",100        "I am writing to formally complain about the catastrophic failure of your service last night. Our entire database was corrupted due to what appears to be a bug in your latest update. We lost 3 days worth of data that cannot be recovered. Our legal team is reviewing our options. I expect a formal response within 24 hours.",101        "CRITICAL: Your system has been sending duplicate notifications to all our users for the last 4 hours. Our customers are furious and several have already cancelled their subscriptions. I have escalated this internally to our CEO. Your on-call team needs to fix this immediately. I am requesting a full incident report within 48 hours.",102        "I have been waiting 3 weeks for my issue to be resolved and I am at the end of my patience. Every time I contact support I get a different answer. This is completely unacceptable for an enterprise customer paying $50,000/year. I am now involving our procurement department to review the contract renewal.",103    ],104}105 106_SENTIMENT_BY_CATEGORY: Dict[EmailCategory, Sentiment] = {107    "spam": "neutral",108    "general_inquiry": "positive",109    "billing_issue": "negative",110    "urgent_complaint": "negative",111}112 113_PRIORITY_BY_CATEGORY: Dict[EmailCategory, Priority] = {114    "spam": "low",115    "general_inquiry": "medium",116    "billing_issue": "medium",117    "urgent_complaint": "high",118}119 120# Occasional overrides for variety121_PRIORITY_OVERRIDES = {122    "billing_issue": {"high": 0.3, "medium": 0.7},  # 30% chance of high priority billing123    "general_inquiry": {"low": 0.2, "medium": 0.8},124}125 126 127def _generate_emails(count: int, seed: int) -> List[Email]:128    """Generate a deterministic list of realistic emails using the given seed."""129    rng = random.Random(seed)130 131    # Category distribution132    categories: List[EmailCategory] = ["spam", "general_inquiry", "billing_issue", "urgent_complaint"]133    weights = [0.2, 0.35, 0.3, 0.15]134 135    emails: List[Email] = []136    base_ts = datetime(2024, 3, 1, 9, 0, 0, tzinfo=timezone.utc)137 138    for i in range(count):139        category: EmailCategory = rng.choices(categories, weights=weights, k=1)[0]140 141        # Pick subject and body142        subject = rng.choice(_SUBJECTS_BY_CATEGORY[category])143        body = rng.choice(_BODIES_BY_CATEGORY[category])144        sender = rng.choice(_SENDERS)145 146        # Timestamp spread over 5 business days147        offset_minutes = rng.randint(0, 5 * 8 * 60)148        ts = base_ts.replace(149            minute=base_ts.minute + offset_minutes % 60,150            hour=base_ts.hour + (offset_minutes // 60) % 8,151        )152 153        # Priority (with occasional overrides for realism)154        if category in _PRIORITY_OVERRIDES:155            opts = _PRIORITY_OVERRIDES[category]156            priority_labels = list(opts.keys())157            priority_weights = list(opts.values())158            gt_priority: Priority = rng.choices(priority_labels, weights=priority_weights, k=1)[0]159        else:160            gt_priority = _PRIORITY_BY_CATEGORY[category]161 162        sentiment: Sentiment = _SENTIMENT_BY_CATEGORY[category]163 164        emails.append(Email(165            id=f"email_{i+1:03d}",166            subject=subject,167            body=body,168            sender=sender,169            timestamp=ts,170            gt_category=category,171            gt_priority=gt_priority,172            sentiment=sentiment,173        ))174 175    return emails176 177 178class EmailTriageEnv:179    """180    SHUBHAMOS OpenEnv-compliant environment for email triage.181 182    Usage:183        env = EmailTriageEnv()184        obs = env.reset(task_config)185        for _ in range(max_steps):186            action = agent.act(obs)187            obs, reward, done, info = env.step(action)188            if done:189                break190        final_state = env.state()191    """192 193    def __init__(self) -> None:194        self._emails: List[Email] = []195        self._step_count: int = 0196        self._max_steps: int = 0197        self._done: bool = False198        self._total_reward: float = 0.0199        self._action_history: List[Dict[str, Any]] = []200        self._seed: int = 0201        self._task_id: str = ""202        self._initialized: bool = False203 204        # Import reward engine lazily to allow Phase 2 to plug in205        self._reward_engine = None206 207    # ── Public OpenEnv API ────────────────────────────────────────────────────208 209    def reset(self, task_config: Dict[str, Any]) -> Observation:210        """211        Reset the environment for a new episode.212 213        Args:214            task_config: Dict with keys:215                - seed (int): random seed for reproducibility216                - email_count (int): number of emails to generate217                - max_steps (int): episode step limit218                - task_id (str, optional): "easy"/"medium"/"hard"219 220        Returns:221            Observation: initial state visible to agent222        """223        self._seed = int(task_config.get("seed", 42))224        email_count = int(task_config.get("email_count", 8))225        self._max_steps = int(task_config.get("max_steps", 50))226        self._task_id = str(task_config.get("task_id", "easy"))227 228        # Seed Python stdlib RNG for any stochastic resets229        random.seed(self._seed)230 231        # Generate deterministic email set232        self._emails = _generate_emails(email_count, self._seed)233        self._step_count = 0234        self._done = False235        self._total_reward = 0.0236        self._action_history = []237        self._initialized = True238 239        return self._build_observation()240 241    def step(self, action: Action) -> Tuple[Observation, float, bool, Dict[str, Any]]:242        """243        Apply an action to the environment.244 245        Args:246            action: Action object with action_type and email_id (+ optional payload)247 248        Returns:249            (observation, reward, done, info)250        """251        if not self._initialized:252            raise RuntimeError("Call reset() before step()")253        if self._done:254            raise RuntimeError("Episode is done. Call reset() to start a new episode.")255 256        # Validate action257        valid, error_msg = self._validate_action(action)258 259        step_reward = 0.0260        if valid:261            # Dispatch action to update email state262            self._dispatch_action(action)263 264            # Calculate reward (reward engine plugged in by Phase 2)265            if self._reward_engine is not None:266                step_reward = self._reward_engine.score_action(action, self._emails)267            else:268                # Stub reward: small negative to incentivize efficiency269                step_reward = -0.01270 271        # Increment step272        self._step_count += 1273        self._total_reward += step_reward274 275        # Record action history276        self._action_history.append({277            "step": self._step_count,278            "action_type": action.action_type,279            "email_id": action.email_id,280            "valid": valid,281            "error": error_msg,282            "reward": step_reward,283        })284 285        # Check termination (D-13)286        all_processed = all(e.is_processed() for e in self._emails)287        step_limit_reached = self._step_count >= self._max_steps288        self._done = all_processed or step_limit_reached289 290        obs = self._build_observation()291        info = StepInfo(292            action_type=action.action_type,293            email_id=action.email_id,294            valid=valid,295            error=error_msg,296            step_reward=step_reward,297            done=self._done,298        ).model_dump()299 300        return obs, step_reward, self._done, info301 302    def state(self) -> State:303        """304        Return the complete internal state including ground truth labels.305        Used by graders — never expose this to the agent during evaluation.306        """307        return State(308            emails=list(self._emails),309            action_history=list(self._action_history),310            step_count=self._step_count,311            max_steps=self._max_steps,312            done=self._done,313            total_reward=self._total_reward,314            seed=self._seed,315            task_id=self._task_id,316        )317 318    # ── Private helpers ───────────────────────────────────────────────────────319 320    def _validate_action(self, action: Action) -> Tuple[bool, Optional[str]]:321        """Validate action target and required payload fields."""322        # Check email exists323        email = self._find_email(action.email_id)324        if email is None:325            return False, f"Email '{action.email_id}' not found in inbox"326 327        # Check email is not already terminally processed328        if email.is_processed():329            return False, f"Email '{action.email_id}' is already processed (resolved/escalated/ignored)"330 331        # Payload validation332        if action.action_type == "classify_email" and action.category is None:333            return False, "classify_email requires 'category' field"334        if action.action_type == "set_priority" and action.level is None:335            return False, "set_priority requires 'level' field"336        if action.action_type == "draft_reply" and (action.text is None or not action.text.strip()):337            return False, "draft_reply requires non-empty 'text' field"338 339        return True, None340 341    def _dispatch_action(self, action: Action) -> None:342        """Mutate email state based on action type."""343        email = self._find_email(action.email_id)344        if email is None:345            return346 347        if action.action_type == "classify_email":348            email.category = action.category349 350        elif action.action_type == "set_priority":351            email.priority = action.level352 353        elif action.action_type == "draft_reply":354            email.reply_drafted = True355            email.reply_text = action.text356 357        elif action.action_type == "mark_resolved":358            email.resolved = True359 360        elif action.action_type == "escalate_email":361            email.escalated = True362 363        elif action.action_type == "ignore_email":364            email.ignored = True365 366    def _find_email(self, email_id: str) -> Optional[Email]:367        """Find an email by ID or return None."""368        for e in self._emails:369            if e.id == email_id:370                return e371        return None372 373    def _build_observation(self) -> Observation:374        """Build the agent-facing Observation from current internal state."""375        obs_emails = []376        for e in self._emails:377            obs_emails.append(ObservationEmail(378                id=e.id,379                subject=e.subject,380                sender=e.sender,381                timestamp=e.timestamp,382                sentiment=e.sentiment,383                body_preview=e.body[:200],  # D-08: 200-char limit384                category=e.category,385                priority=e.priority,386                resolved=e.resolved,387                escalated=e.escalated,388                ignored=e.ignored,389                reply_drafted=e.reply_drafted,390            ))391 392        resolved_count = sum(1 for e in self._emails if e.resolved)393        escalated_count = sum(1 for e in self._emails if e.escalated)394        ignored_count = sum(1 for e in self._emails if e.ignored)395        pending_count = len(self._emails) - resolved_count - escalated_count - ignored_count396 397        elapsed_ratio = (398            self._step_count / self._max_steps if self._max_steps > 0 else 0.0399        )400 401        return Observation(402            emails=obs_emails,403            total_emails=len(self._emails),404            pending_count=pending_count,405            resolved_count=resolved_count,406            escalated_count=escalated_count,407            ignored_count=ignored_count,408            step_count=self._step_count,409            max_steps=self._max_steps,410            elapsed_ratio=elapsed_ratio,411        )412 413    def attach_reward_engine(self, engine: Any) -> None:414        """Attach a reward engine (called by Phase 2). engine must implement score_action()."""415        self._reward_engine = engine416