CoolFace
Apppublic

Yalpha/AGRITECH-META

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py187 linesDownload Raw Back to root
1"""2app.py – FastAPI server for AgriDecisionEnv v33Exposes OpenEnv-compatible REST API for HuggingFace Spaces validation.4 5Endpoints:6  GET  /          → health check (200 OK)7  POST /reset     → reset environment, return initial observation8  POST /step      → take action, return (observation, reward, done, info)9  GET  /state     → return current environment state10  POST /inference → run full inference.py episode, return logs11"""12import os, sys, subprocess, threading13from dotenv import load_dotenv14 15load_dotenv()16from typing import Optional17 18sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))19 20from fastapi import FastAPI, HTTPException21from fastapi.responses import JSONResponse, PlainTextResponse22from pydantic import BaseModel23from env import AgriEnv24from models import Action, Observation25 26app = FastAPI(title="AgriDecisionEnv-v3", version="3.0.0")27 28# Global env instance (one session at a time — sufficient for validator)29_env: Optional[AgriEnv] = None30_lock = threading.Lock()31 32 33def _get_env() -> AgriEnv:34    global _env35    if _env is None:36        raise HTTPException(status_code=400, detail="Call /reset first.")37    return _env38 39 40# ── Health check ──────────────────────────────────────────────────────────────41 42@app.get("/")43def root():44    return {"status": "ok", "env": "AgriDecisionEnv-v3", "version": "3.0.0"}45 46 47@app.get("/health")48def health():49    return {"status": "ok"}50 51 52# ── OpenEnv Core API ──────────────────────────────────────────────────────────53 54class ResetRequest(BaseModel):55    scenario: str = "default"56    seed: int = 4257 58 59@app.post("/reset")60def reset(req: ResetRequest = ResetRequest()):61    global _env62    with _lock:63        _env = AgriEnv(scenario=req.scenario, seed=req.seed)64        obs = _env.reset()65    return obs.model_dump()66 67 68class StepRequest(BaseModel):69    crop: str = "wheat"70    fertilizer: float = 0.371    irrigation: float = 0.472 73 74@app.post("/step")75def step(req: StepRequest):76    with _lock:77        env = _get_env()78        action = Action(79            crop=req.crop,80            fertilizer=req.fertilizer,81            irrigation=req.irrigation,82        )83        try:84            obs, reward, done, info = env.step(action)85        except RuntimeError as e:86            raise HTTPException(status_code=400, detail=str(e))87 88    return {89        "observation": obs.model_dump(),90        "reward":      round(float(reward), 4),91        "done":        done,92        "info":        info,93    }94 95 96@app.get("/state")97def state():98    with _lock:99        env = _get_env()100        return env.state().model_dump()101 102 103# ── Inference endpoint — runs inference.py and returns structured logs ─────────104 105@app.post("/inference")106def run_inference(task: str = "hard", scenario: str = "default"):107    env_vars = os.environ.copy()108    env_vars["AGRI_TASK"]    = task109    env_vars["AGRI_SCENARIO"] = scenario110 111    try:112        result = subprocess.run(113            [sys.executable, "inference.py"],114            capture_output=True,115            text=True,116            timeout=300,117            env=env_vars,118        )119        logs   = result.stdout.strip()120        errors = result.stderr.strip()121 122        # Parse score from [END] line123        score = None124        for line in logs.splitlines():125            if line.startswith("[END]"):126                for part in line.split():127                    if part.startswith("score="):128                        try:129                            score = float(part.split("=")[1])130                        except ValueError:131                            pass132 133        return {134            "task":    task,135            "logs":    logs,136            "errors":  errors or None,137            "score":   score,138            "success": result.returncode == 0,139        }140    except subprocess.TimeoutExpired:141        raise HTTPException(status_code=504, detail="Inference timed out (>300s)")142    except Exception as e:143        raise HTTPException(status_code=500, detail=str(e))144 145 146# ── Task grader endpoints ─────────────────────────────────────────────────────147 148@app.post("/grade/easy")149def grade_easy():150    from tasks.easy import run_easy_task151    action = Action(crop="wheat", fertilizer=0.4, irrigation=0.5)152    score  = run_easy_task(action)153    return {"task": "easy", "score": score}154 155 156@app.post("/grade/medium")157def grade_medium():158    from tasks.medium import run_medium_task159    actions = [160        Action(crop="wheat", fertilizer=0.3, irrigation=0.5),161        Action(crop="rice",  fertilizer=0.4, irrigation=0.6),162        Action(crop="wheat", fertilizer=0.2, irrigation=0.4),163    ]164    score = run_medium_task(actions)165    return {"task": "medium", "score": score}166 167 168@app.post("/grade/hard")169def grade_hard():170    from tasks.hard import run_hard_task171    actions = [172        Action(crop="wheat", fertilizer=0.3, irrigation=0.5),173        Action(crop="rice",  fertilizer=0.5, irrigation=0.6),174        Action(crop="wheat", fertilizer=0.2, irrigation=0.4),175        Action(crop="none",  fertilizer=0.0, irrigation=0.2),176        Action(crop="wheat", fertilizer=0.3, irrigation=0.4),177    ]178    score = run_hard_task(actions)179    return {"task": "hard", "score": score}180 181 182# ── Entry point ───────────────────────────────────────────────────────────────183 184if __name__ == "__main__":185    import uvicorn186    uvicorn.run("app:app", host="0.0.0.0", port=7860, reload=False)187