training-monkey/dataoncallenv
0
1"""FastAPI entry point for DataOnCallEnv.2"""3 4from fastapi import FastAPI, HTTPException5from fastapi.middleware.cors import CORSMiddleware6from pydantic import BaseModel7from typing import Optional8 9from models import Action, Observation, Reward, EnvState10from environment import DataOnCallEnv11from tasks import TASKS12 13app = FastAPI(14 title="DataOnCallEnv",15 description=(16 "RL environment for data pipeline debugging. "17 "The agent acts as an on-call analyst debugging broken reports. "18 "Features: partial observability, query costs, realistic logs, "19 "anti-cheat constraints, and tiered evaluation."20 ),21 version="2.0.0",22 docs_url="/docs",23)24 25app.add_middleware(26 CORSMiddleware,27 allow_origins=["*"],28 allow_methods=["*"],29 allow_headers=["*"],30)31 32# One env instance per worker process33env = DataOnCallEnv()34 35# Models36 37class ResetRequest(BaseModel):38 task_id: int = 139 40class StepRequest(BaseModel):41 tool: str42 query: str43 reasoning: Optional[str] = None44 45class StepResponse(BaseModel):46 observation: Observation47 reward: Optional[Reward] = None48 done: bool49 info: dict50 51# Endpoints52 53@app.get("/health")54def health():55 """Health check endpoint for the environment."""56 return {57 "status": "ok",58 "env": "DataOnCallEnv",59 "version": "2.0.0",60 "tasks": [1, 2, 3],61 "spec": "openenv-0.1",62 "features": [63 "partial_observability",64 "query_costs",65 "realistic_logs",66 "anti_cheat",67 "tiered_evaluation",68 ],69 }70 71@app.get("/")72def root():73 """Root redirect info."""74 return {75 "name": "DataOnCallEnv",76 "docs": "/docs",77 "health": "/health",78 "endpoints": ["/reset", "/step", "/state", "/tasks"],79 }80 81@app.get("/tasks")82def list_tasks():83 """List all available tasks with metadata."""84 return {85 "tasks": [86 {87 "id": t["id"],88 "difficulty": t["difficulty"],89 "title": t["title"],90 "optimal_steps": t["optimal_steps"],91 "optimal_cost": t["optimal_cost"],92 }93 for t in TASKS.values()94 ]95 }96 97@app.post("/reset", response_model=Observation)98def reset(req: ResetRequest= ResetRequest()):99 """100 Start a fresh episode for the given task_id (1, 2, or 3).101 Returns the initial observation containing the scenario description.102 Agent must discover tables via list_tables() — not provided in reset.103 """104 try:105 obs = env.reset(task_id=req.task_id)106 return obs107 except ValueError as e:108 raise HTTPException(status_code=400, detail=str(e))109 110@app.get("/web")111def web():112 return {"status": "ok"}113 114@app.post("/step", response_model=StepResponse)115def step(req: StepRequest):116 """117 Send one agent action. Returns observation + reward (if done).118 Call POST /reset first.119 """120 if env.task_id is None:121 raise HTTPException(122 status_code=400,123 detail="Not initialized. Call POST /reset first."124 )125 if env.done:126 raise HTTPException(127 status_code=400,128 detail="Episode over. Call POST /reset to start a new episode."129 )130 131 action = Action(132 tool=req.tool,133 query=req.query,134 reasoning=req.reasoning,135 )136 137 try:138 obs, reward, done, info = env.step(action)139 except Exception as e:140 raise HTTPException(status_code=500, detail=str(e))141 142 return StepResponse(observation=obs, reward=reward, done=done, info=info)143 144@app.get("/state", response_model=EnvState)145def state():146 """Return full current episode state."""147 return env.state()148def main():149 import uvicorn150 uvicorn.run("api.app:app", host="0.0.0.0", port=7860)151 152if __name__ == "__main__":153 main()154 