krishna-hub's picture
init
85b7ac8
Raw
History Blame
4.26 kB
"""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)