Spaces:
Sleeping
Sleeping
| #!/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() | |