Spaces:
Sleeping
Sleeping
| from fastapi import FastAPI, HTTPException | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from typing import Dict, Any | |
| import uvicorn | |
| from .env import EmergencyResourceEnv | |
| from .models import Observation, Action, Reward, State | |
| from .tasks import EasyTask, MediumTask, HardTask | |
| from .graders import TaskGrader | |
| app = FastAPI(title="Emergency Resource Allocation Environment") | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], | |
| allow_credentials=True, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| # Global environment instance | |
| env = EmergencyResourceEnv() | |
| current_task = None | |
| current_grader = None | |
| trajectory = [] | |
| async def root(): | |
| return {"message": "Emergency Resource Allocation Environment", "status": "running"} | |
| async def reset(task: str = "easy") -> Dict[str, Any]: | |
| """Reset environment with specified task""" | |
| global env, current_task, current_grader, trajectory | |
| trajectory = [] | |
| if task == "easy": | |
| current_task = EasyTask(env) | |
| elif task == "medium": | |
| current_task = MediumTask(env) | |
| elif task == "hard": | |
| current_task = HardTask(env) | |
| else: | |
| raise HTTPException(status_code=400, detail=f"Unknown task: {task}") | |
| current_grader = TaskGrader(current_task.name, current_task.difficulty) | |
| current_task.setup() | |
| observation = env._get_observation() | |
| return observation.model_dump() # FIXED: .dict() -> .model_dump() | |
| async def step(action: Action) -> Dict[str, Any]: | |
| """Execute action and return new state""" | |
| global env, trajectory | |
| observation, reward, done, info = env.step(action) | |
| # Record step in trajectory | |
| trajectory.append((env.agent_pos, action.action_type, reward.score)) | |
| return { | |
| "observation": observation.model_dump(), # FIXED: .dict() -> .model_dump() | |
| "reward": reward.model_dump(), # FIXED: .dict() -> .model_dump() | |
| "done": done, | |
| "info": info | |
| } | |
| async def get_state(): | |
| """Get current environment state - Simplified version without .dict()""" | |
| return { | |
| "agent_position": list(env.agent_pos), | |
| "requests": [ | |
| { | |
| "id": req.id, | |
| "position": list(req.position), | |
| "priority": req.priority, | |
| "allocated": req.allocated, | |
| "created_at": req.created_at | |
| } | |
| for req in env.requests | |
| ], | |
| "time_left": env.time_left, | |
| "resources_left": env.resources_left, | |
| "current_step": env.current_step, | |
| "done": env.done, | |
| "total_reward": env.total_reward | |
| } | |
| async def get_grade() -> Dict[str, Any]: | |
| """Get final grade for current task""" | |
| global current_grader, env, trajectory | |
| if current_grader is None: | |
| raise HTTPException(status_code=400, detail="No task active") | |
| score = current_grader.grade(env, trajectory) | |
| return {"score": score, "task": current_task.name if current_task else None} | |
| if __name__ == "__main__": | |
| uvicorn.run(app, host="127.0.0.1", port=7860) |