CoolFace
Apppublic

Grish2114/email-prioritization-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
server.py170 linesDownload Raw Back to api
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