Spaces:
Running
Running
File size: 7,103 Bytes
1165a28 072c8d8 1165a28 5c2f240 1165a28 072c8d8 5c2f240 1165a28 5c2f240 1165a28 072c8d8 1165a28 | 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 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 | """
Cat Rescue β OpenEnv FastAPI Server
=====================================
Wraps CatRescueEnv, CatRescueGrader, and CatRescueRewards into a
standard OpenEnv HTTP API.
Endpoints
---------
POST /reset β start a new episode (calls env.reset())
POST /step β take one action (calls env.step(action))
GET /state β read current state (calls env.state())
POST /grade β score an episode (calls grader.grade(episode_log))
Run locally
-----------
uvicorn server:app --host 0.0.0.0 --port 7860
HuggingFace Spaces expects port 7860, which is why we default to it.
"""
from __future__ import annotations
import os
from typing import Any, Dict, List, Optional
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse
from pydantic import BaseModel, Field
# ---------------------------------------------------------------------------
# Import the three core modules written by the team
# ---------------------------------------------------------------------------
from environment import CatRescueEnv # YOUR file
from grader import CatRescueGrader # teammate's file
from rewards import CatRescueRewards # teammate's file
# ---------------------------------------------------------------------------
# FastAPI app setup
# ---------------------------------------------------------------------------
app = FastAPI(
title="Cat Rescue OpenEnv",
description=(
"Grid-based AI environment for the Meta Γ PyTorch OpenEnv Hackathon. "
"An agent navigates a grid to rescue trapped cats."
),
version="1.0.0",
)
# Allow cross-origin requests (needed for HF Spaces iframes / web agents)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
# ---------------------------------------------------------------------------
# Shared singleton instances
# A single environment instance is kept alive between requests so that
# the agent can call /reset once and then repeatedly call /step.
# ---------------------------------------------------------------------------
env = CatRescueEnv(level=1) # default level; /reset can change it
grader = CatRescueGrader()
rewards = CatRescueRewards()
# ---------------------------------------------------------------------------
# Request / Response schemas (Pydantic keeps the API self-documenting)
# ---------------------------------------------------------------------------
class ResetRequest(BaseModel):
"""Body for POST /reset."""
level: int = Field(default=1, ge=1, le=3, description="Difficulty level (1, 2, or 3).")
max_steps: int = Field(default=200, ge=1, description="Max steps before episode ends.")
seed: Optional[int] = Field(default=None, description="Random seed for reproducibility.")
class StepRequest(BaseModel):
"""Body for POST /step."""
action: int = Field(
...,
ge=0,
le=3,
description="Action to take. 0=UP, 1=DOWN, 2=LEFT, 3=RIGHT.",
)
class StepResponse(BaseModel):
"""Response from POST /step."""
observation: Dict[str, Any]
reward: float
done: bool
info: Dict[str, Any]
class GradeRequest(BaseModel):
"""Body for POST /grade. episode_log is a list of step records."""
episode_log: List[Dict[str, Any]] = Field(
...,
description=(
"Ordered list of step dicts, each containing at least "
"'action', 'reward', 'done', and 'info' keys."
),
)
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
# Absolute path to the directory that contains server.py
# (works both locally and inside a Docker container)
_HERE = os.path.dirname(os.path.abspath(__file__))
@app.get("/", tags=["UI"])
def root() -> FileResponse:
"""
Serve the game UI.
Returns index.html so players can open the root URL in a browser.
"""
return FileResponse(os.path.join(_HERE, "index.html"), media_type="text/html")
@app.post("/reset", tags=["OpenEnv"])
def reset(body: ResetRequest) -> Dict[str, Any]:
"""
**POST /reset** β Start a fresh episode.
Re-initialises the environment with the requested level, max_steps,
and optional seed, then returns the initial observation.
This follows the OpenEnv standard: always call /reset before /step.
"""
global env # replace the singleton with a freshly configured instance
env = CatRescueEnv(
level=body.level,
max_steps=body.max_steps,
seed=body.seed,
)
# env.reset() is already called inside __init__; call again for clarity
# and to return the canonical initial observation.
observation = env.reset()
return {"observation": observation}
@app.post("/step", tags=["OpenEnv"], response_model=StepResponse)
def step(body: StepRequest) -> StepResponse:
"""
**POST /step** β Execute one action.
The agent POSTs `{"action": <int>}` and receives back:
- `observation` β the new grid state
- `reward` β reward for this transition
- `done` β whether the episode has ended
- `info` β diagnostic dict (event, agent_pos, cats_rescued β¦)
Raises 400 if the episode is already finished (call /reset first).
"""
if env.done:
raise HTTPException(
status_code=400,
detail="Episode is already done. Call POST /reset to start a new one.",
)
observation, reward, done, info = env.step(body.action)
return StepResponse(observation=observation, reward=reward, done=done, info=info)
@app.get("/state", tags=["OpenEnv"])
def state() -> Dict[str, Any]:
"""
**GET /state** β Read the current environment state without advancing it.
Returns the same observation dict that /step returns, but does NOT
consume a step or change any game state.
"""
return {"observation": env.state()}
@app.post("/grade", tags=["Grading"])
def grade(body: GradeRequest) -> Dict[str, Any]:
"""
**POST /grade** β Score a completed episode.
Accepts the full `episode_log` (list of step records collected by
the agent) and delegates scoring to CatRescueGrader.
The grader returns a score dict; exact keys depend on grader.py.
"""
try:
result = grader.grade(body.episode_log)
except Exception as exc:
raise HTTPException(status_code=500, detail=f"Grader error: {exc}") from exc
return {"grade": result}
# ---------------------------------------------------------------------------
# Entry point β lets you run `python server.py` directly during development.
# For production / HF Spaces use: uvicorn server:app --host 0.0.0.0 --port 7860
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import uvicorn
uvicorn.run("server:app", host="0.0.0.0", port=7860, reload=True)
|