harshal15122003 commited on
Commit
e29bcb3
·
verified ·
1 Parent(s): 9fcc8c8

Update server/app.py

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