#!/usr/bin/env python3 """ Email Triage OpenEnv - Baseline Inference Script Runs a language model agent against all 3 email triage tasks and logs results in the required OpenEnv [START]/[STEP]/[END] format. Environment variables: MODEL_NAME : Model identifier (default: google/gemma-4-31B-it) API_BASE_URL : LLM API endpoint (default: https://router.huggingface.co/hf-inference/models/google/gemma-4-31B-it/v1) HF_TOKEN : HuggingFace API token (fallback: OPENAI_API_KEY) Usage: export HF_TOKEN="hf_..." python inference.py """ from __future__ import annotations import logging import os import sys import time from typing import Dict, List, Optional from openai import OpenAI from src.environment import EmailTriageEnv, TASK_CONFIG from src.grader import grade_task_basic, grade_task_medium, grade_task_hard from src.models import Action, Observation from src.utils import format_action_for_log, parse_action_string # ── Configuration ──────────────────────────────────────────────────────────── API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/hf-inference/models/google/gemma-4-31B-it/v1") MODEL_NAME = os.getenv("MODEL_NAME", "google/gemma-4-31B-it") API_KEY = os.getenv("HF_TOKEN") or os.getenv("OPENAI_API_KEY") or "" TEMPERATURE = 0.3 # low temperature for consistency MAX_RETRIES = 2 logging.basicConfig( level=logging.WARNING, format="%(asctime)s [%(levelname)s] %(name)s: %(message)s", stream=sys.stderr, ) logger = logging.getLogger(__name__) # Task → grader mapping GRADERS = { "basic_triage": grade_task_basic, "multi_folder_triage": grade_task_medium, "advanced_triage_with_urgency": grade_task_hard, } # ── Logging (OpenEnv format → stdout) ──────────────────────────────────────── def log_start(task_id: str, model: str) -> None: print(f"[START] task={task_id} env=email-triage model={model}", flush=True) def log_step( step: int, action_str: str, reward: float, done: bool, error: Optional[str], ) -> None: error_part = f'"{error}"' if error else "null" print( f"[STEP] step={step} action={action_str} " f"reward={reward:.2f} done={str(done).lower()} error={error_part}", flush=True, ) def log_end( task_id: str, success: bool, steps: int, score: float, rewards: List[float], ) -> None: rewards_str = ",".join(f"{r:.2f}" for r in rewards) print( f"[END] task={task_id} success={str(success).lower()} " f"steps={steps} score={score:.2f} rewards={rewards_str}", flush=True, ) # ── Agent ──────────────────────────────────────────────────────────────────── class EmailTriageAgent: """Baseline LLM agent for email triage.""" def __init__(self, model_name: str, api_base: str, api_key: str) -> None: self.client = OpenAI(api_key=api_key, base_url=api_base) self.model_name = model_name def get_action( self, observation: Observation, step_num: int, max_steps: int, ) -> str: """Query the LLM to decide which action to take. Returns the raw action string from the model. """ inbox_lines = [] for e in observation.inbox_emails[:10]: vip_tag = " [VIP]" if e.is_vip_sender else "" prio_tag = " [URGENT]" if e.priority_flag else "" inbox_lines.append( f" ID={e.id} | From: {e.sender} | " f"Subject: {e.subject[:60]}{vip_tag}{prio_tag}" ) inbox_text = "\n".join(inbox_lines) if inbox_lines else " (empty)" prompt = f"""You are an email triage assistant. Sort each email into the correct folder. INBOX ({len(observation.inbox_emails)} emails): {inbox_text} FOLDERS: {', '.join(observation.available_folders)} RULES: - Work emails (@acmecorp.com, project updates, code reviews) -> work - Invoices, expenses, budgets, purchase orders -> finance - Calendar invites, meeting notes, RSVPs -> meetings - Phishing, marketing spam, suspicious domains (.ru, .xyz, .biz) -> spam - Newsletters, notifications, automated reports -> archive - VIP/urgent emails should be prioritized and classified correctly Step {step_num}/{max_steps}. Respond with ONLY the action, e.g.: move(0, work) """ for attempt in range(MAX_RETRIES + 1): try: response = self.client.chat.completions.create( model=self.model_name, messages=[{"role": "user", "content": prompt}], max_tokens=60, temperature=TEMPERATURE, timeout=30, ) return response.choices[0].message.content.strip() except Exception as e: logger.warning( "LLM call failed (attempt %d/%d): %s", attempt + 1, MAX_RETRIES + 1, e, ) if attempt < MAX_RETRIES: time.sleep(2 ** attempt) # Fallback: move first email to work if observation.inbox_emails: return f"move({observation.inbox_emails[0].id}, work)" return "move(0, work)" def _parse_or_fallback( raw_action: str, observation: Observation ) -> Action: """Parse the LLM's raw output into an Action, with a safe fallback.""" action = parse_action_string(raw_action) if action is not None: return action # Fallback: move first available email to work logger.warning("Failed to parse action: '%s'", raw_action) if observation.inbox_emails: return Action( action_type="move", email_id=observation.inbox_emails[0].id, target_folder="work", ) return Action(action_type="move", email_id=0, target_folder="work") # ── Task Runner ────────────────────────────────────────────────────────────── def run_task(task_id: str, agent: EmailTriageAgent, seed: int = 42) -> Dict: """Run the agent on a single task and log results. Returns a summary dict with score, steps, success status. """ env = EmailTriageEnv(task_id=task_id, seed=seed) obs = env.reset() max_steps = env.max_steps log_start(task_id, MODEL_NAME) rewards: List[float] = [] steps_taken = 0 error_msg: Optional[str] = None try: for step_num in range(1, max_steps + 1): if obs.done: break raw_action = agent.get_action(obs, step_num, max_steps) action = _parse_or_fallback(raw_action, obs) action_str = format_action_for_log(action) try: obs, reward, done, info = env.step(action) reward_val = reward.value step_error = None except Exception as e: reward_val = 0.0 done = False step_error = str(e) logger.error("Step error: %s", e) rewards.append(reward_val) steps_taken = step_num log_step( step=step_num, action_str=action_str, reward=reward_val, done=done, error=step_error, ) if done: break except Exception as e: error_msg = str(e) logger.error("Task %s failed: %s", task_id, error_msg) # Compute final score using the appropriate grader grader = GRADERS.get(task_id) if grader and env.history: score = grader(env.history) elif rewards: score = sum(rewards) / len(rewards) else: score = 0.0 score = max(0.0, min(1.0, score)) success = score >= TASK_CONFIG[task_id].get("success_threshold", 0.7) log_end( task_id=task_id, success=success, steps=steps_taken, score=score, rewards=rewards, ) return { "task_id": task_id, "score": score, "steps": steps_taken, "rewards": rewards, "success": success, "error": error_msg, } # ── Main ───────────────────────────────────────────────────────────────────── def main() -> None: """Run baseline inference on all 3 tasks.""" if not API_KEY: logger.error( "No API key found. Set HF_TOKEN or OPENAI_API_KEY environment variable." ) sys.exit(1) agent = EmailTriageAgent( model_name=MODEL_NAME, api_base=API_BASE_URL, api_key=API_KEY, ) tasks = [ "basic_triage", "multi_folder_triage", "advanced_triage_with_urgency", ] results = [] for task_id in tasks: result = run_task(task_id, agent, seed=42) results.append(result) # Summary to stderr (not stdout, to keep stdout clean for log parsing) print("\n" + "=" * 60, file=sys.stderr) print("BASELINE INFERENCE SUMMARY", file=sys.stderr) print("=" * 60, file=sys.stderr) for r in results: status = "PASS" if r["success"] else "FAIL" print( f" {r['task_id']:40s} score={r['score']:.2f} [{status}]", file=sys.stderr, ) avg = sum(r["score"] for r in results) / len(results) if results else 0.0 print(f"\n Average Score: {avg:.2f}", file=sys.stderr) print("=" * 60, file=sys.stderr) if __name__ == "__main__": main()