Meta-Hackathon1/Customer_Support_Env
0
1from typing import Optional2 3from fastapi import FastAPI4from pydantic import BaseModel5 6from dataset import easy_cases, hard_cases, medium_cases7from environment import SupportEnv8 9app = FastAPI()10 11# --- Models ---12class Action(BaseModel):13 message: str14 15class ResetRequest(BaseModel):16 task: str = "easy"17 18class Observation(BaseModel):19 echoed_message: str20 21class StepResponse(BaseModel):22 observation: Observation23 reward: float24 done: bool25 info: dict = {}26 27# --- Parser ---28def parse_action(message):29 message = message.lower()30 31 if "refund" in message:32 return "refund"33 elif "replace" in message:34 return "replace"35 elif "reject" in message:36 return "reject"37 elif "proof" in message:38 return "ask_proof"39 else:40 return "reject"41 42# --- TASKS ---43TASKS = {"easy": easy_cases, "medium": medium_cases, "hard": hard_cases}44 45env = None46 47# --- Routes ---48 49# RESET50@app.post("/reset")51def reset(req: Optional[ResetRequest] = None):52 global env53 54 task_name = req.task if req and req.task in TASKS else "easy"55 56 cases = TASKS[task_name]57 env = SupportEnv(cases)58 59 obs = env.reset()60 61 return {62 "observation": {"echoed_message": obs}, 63 "reward": 0.0, 64 "done": False, 65 "info": {}66 }67 68# STATE69@app.get("/state")70def get_state():71 global env72 if env is None:73 return {"error": "Environment not initialized. Call /reset first."}74 75 return {"observation": {"echoed_message": env._get_obs()}}76 77# STEP78@app.post("/step", response_model=StepResponse)79def step(action: Action):80 global env81 82 # Safety check if step is called before reset83 if env is None:84 return {85 "observation": {"echoed_message": "Error: Environment not initialized."},86 "reward": 0.0,87 "done": True,88 "info": {}89 }90 91 parsed_action = parse_action(action.message)92 93 obs, reward, done = env.step(parsed_action)94 95 response = {96 "observation": {"echoed_message": obs}, 97 "reward": reward, 98 "done": done, 99 "info": {}100 }101 102 if done:103 response["info"]["score"] = float(env.get_score())104 105 return response106 107# --- Main ---108def main():109 import uvicorn110 uvicorn.run(app, host="0.0.0.0", port=7860)111 112if __name__ == "__main__":113 main()