AMD21/codefixerenv
0
1import os2import sys3sys.path.insert(0, os.path.dirname(__file__))4 5from fastapi import FastAPI, Request6from fastapi.middleware.cors import CORSMiddleware7from fastapi.responses import JSONResponse8from pydantic import BaseModel9from typing import Optional10import uvicorn11 12from env.environment import CodeFixerEnv13from env.models import Action, Observation, Reward, StepResult14from tasks.tasks import TASKS15 16app = FastAPI(17 title="CodeFixerEnv",18 description="OpenEnv environment for code-debugging RL tasks.",19 version="1.0.0",20 docs_url="/docs",21 redoc_url="/redoc",22)23 24app.add_middleware(25 CORSMiddleware,26 allow_origins=["*"],27 allow_methods=["*"],28 allow_headers=["*"],29)30 31env = CodeFixerEnv()32 33class ResetRequest(BaseModel):34 difficulty: Optional[str] = "easy"35 36class StepRequest(BaseModel):37 action_type: str38 action_content: str39 40class StepResponse(BaseModel):41 observation: Observation42 reward: Reward43 done: bool44 info: dict45 46class TaskInfo(BaseModel):47 id: str48 difficulty: str49 context: str50 51@app.get("/health")52def health():53 return {"status": "ok", "environment": "CodeFixerEnv", "version": "1.0.0"}54 55@app.post("/reset")56async def reset(request: Request):57 difficulty = "easy"58 try:59 body = await request.body()60 if body:61 try:62 import json63 data = json.loads(body)64 if isinstance(data, dict):65 difficulty = data.get("difficulty", "easy")66 elif isinstance(data, str):67 difficulty = data68 except Exception:69 difficulty = "easy"70 except Exception:71 difficulty = "easy"72 73 try:74 obs = env.reset(difficulty=difficulty)75 return JSONResponse(content=obs.model_dump())76 except ValueError as e:77 return JSONResponse(status_code=400, content={"detail": str(e)})78 79@app.post("/step")80async def step(request: Request):81 try:82 body = await request.body()83 import json84 data = json.loads(body)85 action = Action(type=data["action_type"], content=data.get("action_content", ""))86 obs, reward, done, info = env.step(action)87 return JSONResponse(content={88 "observation": obs.model_dump(),89 "reward": reward.model_dump(),90 "done": done,91 "info": info,92 })93 except RuntimeError as e:94 return JSONResponse(status_code=400, content={"detail": str(e)})95 except Exception as e:96 return JSONResponse(status_code=500, content={"detail": str(e)})97 98@app.get("/state")99def state():100 return env.state()101 102@app.get("/tasks")103def list_tasks():104 return [105 {"id": t["id"], "difficulty": t["difficulty"], "context": t["context"]}106 for _, (t, _) in TASKS.items()107 ]108 109if __name__ == "__main__":110 port = int(os.environ.get("PORT", 7860))111 uvicorn.run(app, host="0.0.0.0", port=port)