Team Ai
Apppublic

training-monkey/dataoncallenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py154 linesDownload Raw Back to api
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