| """FastAPI entrypoint implementing required environment endpoints.""" |
|
|
| from __future__ import annotations |
|
|
| from typing import Optional |
|
|
| import uvicorn |
| from fastapi import FastAPI, HTTPException |
| from fastapi.middleware.cors import CORSMiddleware |
|
|
| from server.environment import CICDDebugEnvironment |
| from server.graders import run_grader |
| from server.models import ( |
| Action, |
| BaselineRequest, |
| BaselineResponse, |
| EnvironmentInfo, |
| GraderRequest, |
| GraderResponse, |
| Observation, |
| ResetRequest, |
| ResetResponse, |
| StateResponse, |
| StepRequest, |
| StepResponse, |
| TaskInfo, |
| ) |
| from server.tasks.task_registry import TASK_REGISTRY |
|
|
| app = FastAPI( |
| title="CI/CD Debug Environment", |
| description="OpenEnv-style environment for Docker + GitHub Actions debugging", |
| version="1.0.0", |
| ) |
|
|
| app.add_middleware( |
| CORSMiddleware, |
| allow_origins=["*"], |
| allow_credentials=True, |
| allow_methods=["*"], |
| allow_headers=["*"], |
| ) |
|
|
| env: Optional[CICDDebugEnvironment] = None |
|
|
|
|
| @app.get("/") |
| async def root(): |
| return {"status": "healthy", "environment": "cicd-debug-env"} |
|
|
|
|
| @app.post("/reset", response_model=ResetResponse) |
| async def reset(request: Optional[ResetRequest] = None): |
| global env |
|
|
| request = request or ResetRequest() |
| env = CICDDebugEnvironment() |
| try: |
| observation = env.reset( |
| task_id=request.task_id, |
| scenario_id=request.scenario_id, |
| seed=request.seed, |
| ) |
| except ValueError as exc: |
| raise HTTPException(status_code=400, detail=str(exc)) from exc |
|
|
| return ResetResponse( |
| observation=observation, |
| info={ |
| "task_id": env.current_task_id, |
| "scenario_id": env.current_scenario_id, |
| "difficulty": env.current_difficulty, |
| }, |
| ) |
|
|
|
|
| @app.post("/step", response_model=StepResponse) |
| async def step(request: StepRequest): |
| global env |
|
|
| if env is None: |
| raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.") |
|
|
| observation, reward, done, info = env.step(request.action) |
| return StepResponse(observation=observation, reward=reward, done=done, info=info) |
|
|
|
|
| @app.get("/state", response_model=StateResponse) |
| async def get_state(): |
| global env |
|
|
| if env is None: |
| raise HTTPException(status_code=400, detail="Environment not initialized. Call /reset first.") |
|
|
| return StateResponse( |
| observation=env.get_observation(), |
| episode_reward=env.episode_reward, |
| steps_taken=env.step_count, |
| done=env.done, |
| ) |
|
|
|
|
| @app.get("/info", response_model=EnvironmentInfo) |
| async def get_info(): |
| tasks = [ |
| TaskInfo( |
| id=task_id, |
| name=task_cls.NAME, |
| description=task_cls.DESCRIPTION, |
| difficulty=task_cls.DIFFICULTY, |
| num_scenarios=len(task_cls.SCENARIOS), |
| ) |
| for task_id, task_cls in TASK_REGISTRY.items() |
| ] |
| return EnvironmentInfo( |
| tasks=tasks, |
| max_steps=10, |
| action_space=Action.model_json_schema(), |
| observation_space=Observation.model_json_schema(), |
| ) |
|
|
|
|
| @app.get("/tasks") |
| async def get_tasks(): |
| return { |
| "tasks": [ |
| { |
| "id": task_id, |
| "name": task_cls.NAME, |
| "description": task_cls.DESCRIPTION, |
| "difficulty": task_cls.DIFFICULTY.value, |
| } |
| for task_id, task_cls in TASK_REGISTRY.items() |
| ] |
| } |
|
|
|
|
| @app.post("/grader", response_model=GraderResponse) |
| async def grade(request: GraderRequest): |
| result = run_grader(task_id=request.task_id, trajectory=request.trajectory) |
| return GraderResponse(result=result) |
|
|
|
|
| @app.post("/baseline", response_model=BaselineResponse) |
| async def run_baseline(request: Optional[BaselineRequest] = None): |
| request = request or BaselineRequest() |
|
|
| from baseline_runner import run_baseline_episodes |
|
|
| results = run_baseline_episodes(task_id=request.task_id, num_episodes=request.num_episodes) |
| aggregate = sum(r.score for r in results) / len(results) if results else 0.0 |
| return BaselineResponse(results=results, aggregate_score=aggregate) |
|
|
|
|
| if __name__ == "__main__": |
| uvicorn.run(app, host="0.0.0.0", port=7860) |
|
|