Grish2114/email-prioritization-env
0
1from fastapi import FastAPI, HTTPException, Body2from fastapi.responses import HTMLResponse3from pydantic import BaseModel4import os5import litellm6import re7 8from env import EmailEnvironment, Action, Priority, Department, EnvironmentState, list_tasks, Observation, Reward9from env.graders import grade10from typing import List, Optional11 12app = FastAPI(title="Email Prioritization Env")13 14env = EmailEnvironment()15 16class ResetRequest(BaseModel):17 task_id: int18 seed: int = 4219 20class StepRequest(BaseModel):21 task_id: int22 priority: str23 department: str24 reply: str = ""25 26class GraderRequest(BaseModel):27 task_id: int28 email_id: str29 priority: str30 department: str31 reply: str = ""32 33class BaselineRequest(BaseModel):34 api_key: str35 model: str = "claude-3-5-sonnet-20241022"36 tasks: list[int] = [1, 2, 3]37 seed: int = 4238 39@app.get("/health")40def health():41 return {"status": "ok"}42 43@app.post("/reset")44def reset_env(req: ResetRequest):45 obs = env.reset(task_id=req.task_id)46 return obs.model_dump()47 48@app.post("/step")49def step_env(req: StepRequest):50 try:51 action = Action(52 priority=Priority(req.priority.lower()),53 department=Department(req.department.lower()),54 reply=req.reply55 )56 except Exception as e:57 raise HTTPException(status_code=400, detail=str(e))58 59 obs, reward, done, info = env.step(action)60 return {61 "observation": obs.model_dump() if obs else None,62 "reward": reward.model_dump(),63 "done": done,64 "info": info65 }66 67@app.get("/state")68def get_state(task_id: int = 1):69 return env.state().model_dump()70 71@app.get("/tasks")72def get_tasks_endpoint():73 return {74 "tasks": list_tasks(),75 "action_schema": Action.model_json_schema()76 }77 78@app.post("/grader")79def run_grader(req: GraderRequest):80 try:81 action = Action(82 priority=Priority(req.priority.lower()),83 department=Department(req.department.lower()),84 reply=req.reply85 )86 except Exception as e:87 raise HTTPException(status_code=400, detail=str(e))88 89 e_env = EmailEnvironment(task_id=req.task_id)90 email = next((e for e in e_env.all_emails if e["email_id"] == req.email_id), None)91 if not email:92 raise HTTPException(status_code=404, detail="Email not found")93 94 reward = grade(req.task_id, action, email, cumulative_score=0.0)95 return reward.model_dump()96 97@app.post("/baseline")98def run_baseline(req: BaselineRequest):99 results = []100 101 for task_id in req.tasks:102 b_env = EmailEnvironment(task_id=task_id, seed=req.seed)103 obs = b_env.reset()104 task_desc = b_env.task.description105 106 task_results = {107 "task_id": task_id,108 "steps": []109 }110 111 while not b_env.is_done:112 prompt = f"You are an expert email triage agent. {task_desc}. Email from {obs.sender} with subject: {obs.subject}. Body: {obs.body}. Reply with PRIORITY: urgent|normal|low, DEPARTMENT: billing|technical|hr|general, REPLY: [text for task 3 only]"113 114 try:115 message = litellm.completion(116 model=req.model,117 api_key=req.api_key,118 max_tokens=1000,119 messages=[120 {"role": "user", "content": prompt}121 ]122 )123 text = message.choices[0].message.content124 125 priority_val = "normal"126 department_val = "general"127 reply_val = ""128 129 p_match = re.search(r'(?i)priority\s*(?:\*\*)?\s*:\s*(?:\*\*)?\s*(urgent|normal|low)', text)130 if p_match:131 priority_val = p_match.group(1).lower()132 133 d_match = re.search(r'(?i)department\s*(?:\*\*)?\s*:\s*(?:\*\*)?\s*(billing|technical|hr|general)', text)134 if d_match:135 department_val = d_match.group(1).lower()136 137 r_match = re.search(r'(?i)reply\s*(?:\*\*)?\s*:\s*(.*)', text, re.DOTALL)138 if r_match:139 reply_val = r_match.group(1).strip()140 141 action = Action(142 priority=Priority(priority_val) if priority_val in Priority._value2member_map_ else Priority.normal,143 department=Department(department_val) if department_val in Department._value2member_map_ else Department.general,144 reply=reply_val145 )146 except Exception as e:147 print(f"Agent error: {e}")148 action = Action(priority=Priority.normal, department=Department.general)149 150 next_obs, reward, done, info = b_env.step(action)151 152 task_results["steps"].append({153 "observation": obs.model_dump(),154 "action": action.model_dump(),155 "reward": reward.model_dump()156 })157 158 obs = next_obs159 160 task_results["final_info"] = info161 results.append(task_results)162 163 return results164 165@app.get("/")166def serve_index():167 path = os.path.join(os.path.dirname(__file__), "..", "static", "index.html")168 with open(path, "r", encoding="utf-8") as f:169 return HTMLResponse(content=f.read())170 