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() @app.post("/reset", response_model=TrafficState) 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 @app.post("/step", response_model=StepResult) def step(request: StepRequest): observation, reward, done, info = ENV.step(request.action) return StepResult(observation=observation, reward=reward, done=done, info=info) @app.get("/state", response_model=TrafficState) def state(): return ENV.get_state() @app.get("/tasks") def tasks(): return list_tasks() @app.post("/grader") 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 @app.get("/baseline") 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