Spaces:
Sleeping
Sleeping
File size: 3,687 Bytes
f74885f d5f4938 f74885f 0f2bf94 f74885f 0f2bf94 d5f4938 0f2bf94 d5f4938 f74885f d5f4938 0f2bf94 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 f74885f d5f4938 0f2bf94 3ad39e3 675fdb5 3ad39e3 d5f4938 0f2bf94 9fcc8c8 0f2bf94 3ad39e3 0f2bf94 d5f4938 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | import uvicorn
from fastapi import FastAPI
from pydantic import BaseModel
from typing import List, Dict, Any
from env import EmailSortingEnv
app = FastAPI(
title="Email Sorting OpenEnv",
description="Real-world email sorting environment for RL agents",
version="1.0.0"
)
env = EmailSortingEnv()
# ============================================
# PYDANTIC MODELS — typed Observation, Action, Reward
# ============================================
class EmailModel(BaseModel):
subject: str
body: str
sender: str
class Observation(BaseModel):
email: EmailModel
step: int
max_steps: int
total_reward: float
done: bool
valid_actions: List[str]
class StepRequest(BaseModel):
action: str
class StepResponse(BaseModel):
observation: Observation
reward: float
done: bool
info: Dict[str, Any]
class ResetResponse(BaseModel):
observation: Observation
class GraderTask(BaseModel):
task_id: str
score: float
class GradersResponse(BaseModel):
tasks: List[GraderTask]
average_score: float
# ============================================
# API ENDPOINTS
# ============================================
@app.get("/health")
def health_check():
return {"status": "ok", "message": "Email Sorting Environment is running"}
@app.get("/")
def root():
return {
"name": "Email Sorting OpenEnv",
"version": "1.0.0",
"description": "Sort emails as spam, important, or promotion",
"endpoints": ["/reset", "/step", "/state", "/graders", "/health"]
}
@app.post("/reset", response_model=ResetResponse)
def reset():
"""Reset environment and return initial observation."""
state = env.reset()
return ResetResponse(observation=Observation(**state))
@app.post("/step", response_model=StepResponse)
def step(request: StepRequest):
"""Take action and return next observation, reward, done, info."""
next_state, reward, done, info = env.step(request.action)
return StepResponse(
observation=Observation(**next_state),
reward=reward,
done=done,
info=info
)
@app.get("/state", response_model=ResetResponse)
def get_state():
"""Return current observation without taking action."""
return ResetResponse(observation=Observation(**env.state()))
@app.get("/graders", response_model=GradersResponse)
def run_graders():
"""Run all graders and return scores."""
from graders import grade_easy_sorting, grade_medium_sorting, grade_hard_sorting
tasks = [
GraderTask(task_id="easy_sorting", score=grade_easy_sorting()),
GraderTask(task_id="medium_sorting", score=grade_medium_sorting()),
GraderTask(task_id="hard_sorting", score=grade_hard_sorting()),
]
avg = round(sum(t.score for t in tasks) / len(tasks), 4)
return GradersResponse(tasks=tasks, average_score=avg)
@app.get("/graders/easy_sorting")
def grade_easy():
from graders import grade_easy_sorting
return {"task_id": "easy_sorting", "score": grade_easy_sorting()}
@app.get("/graders/medium_sorting")
def grade_medium():
from graders import grade_medium_sorting
return {"task_id": "medium_sorting", "score": grade_medium_sorting()}
@app.get("/graders/hard_sorting")
def grade_hard():
from graders import grade_hard_sorting
return {"task_id": "hard_sorting", "score": grade_hard_sorting()}
# ============================================
# START SERVER
# ============================================
if __name__ == "__main__":
print("Starting Email Sorting Environment Server...")
print("Server running at http://localhost:7860")
uvicorn.run(app, host="0.0.0.0", port=7860)
|