Spaces:
Sleeping
Sleeping
| from fastapi import FastAPI, HTTPException | |
| from baseline import baseline_scores, run_baseline | |
| from env import IndianTrafficEnv | |
| from grader import grade_rollout | |
| from models import GraderRequest, ResetRequest, StepRequest, StepResult, TrafficState | |
| from tasks import list_tasks | |
| app = FastAPI( | |
| title="Indian Traffic Signal OpenEnv", | |
| version="1.0.0", | |
| description="Seedable RL environment for Indian mixed-traffic signal control.", | |
| ) | |
| ENV = IndianTrafficEnv() | |
| def reset(request: ResetRequest): | |
| try: | |
| return ENV.reset(seed=request.seed, task_id=request.task_id) | |
| except KeyError as exc: | |
| raise HTTPException(status_code=404, detail=str(exc)) from exc | |
| def step(request: StepRequest): | |
| observation, reward, done, info = ENV.step(request.action) | |
| return StepResult(observation=observation, reward=reward, done=done, info=info) | |
| def state(): | |
| return ENV.get_state() | |
| def tasks(): | |
| return list_tasks() | |
| def grader(request: GraderRequest): | |
| try: | |
| return grade_rollout( | |
| task_id=request.task_id, | |
| seed=request.seed, | |
| actions=request.actions, | |
| max_steps=request.max_steps, | |
| ) | |
| except KeyError as exc: | |
| raise HTTPException(status_code=404, detail=str(exc)) from exc | |
| def baseline(task_id: str = "all", seed: int = 42): | |
| if task_id == "all": | |
| return baseline_scores(seed=seed) | |
| try: | |
| return run_baseline(task_id=task_id, seed=seed) | |
| except KeyError as exc: | |
| raise HTTPException(status_code=404, detail=str(exc)) from exc | |