Team Ai
Apppublic

SHUBHAMOS/meta-pytorch-hackathon

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
server.py240 linesDownload Raw Back to server
1"""2SHUBHAMOS: AI Email Operations & Triage Environment3server.py — FastAPI HTTP server wrapping the OpenEnv API4 5Endpoints:6    POST /reset    — Start a new episode7    POST /step     — Apply one action, get observation + reward8    GET  /state    — Get full internal state (with ground truth)9    GET  /health   — Liveness check10    GET  /grade    — Grade the current episode11    GET  /         — Serve dashboard12 13Usage:14    uvicorn server:app --host 0.0.0.0 --port 786015    OR16    python server.py17"""18 19from __future__ import annotations20import json21import os22from pathlib import Path23from typing import Any, Dict, Optional24 25import yaml26from fastapi import FastAPI, HTTPException, Request27from fastapi.responses import HTMLResponse, JSONResponse28from fastapi.staticfiles import StaticFiles29from pydantic import BaseModel as PydanticBaseModel30 31from .environment import EmailTriageEnv32from .models import Action33from .reward import RewardEngine34from .tasks import TASKS35from .graders import EasyGrader, MediumGrader, HardGrader, PeacefulGrader, ExtremeGrader36 37# ── Setup ─────────────────────────────────────────────────────────────────────38 39app = FastAPI(40    title="SHUBHAMOS: AI Email Operations & Triage Environment",41    description="OpenEnv-compatible email triage environment with reset/step/state API",42    version="1.0.0",43)44 45# Singleton environment (one episode at a time per server instance)46_env = EmailTriageEnv()47_engine = RewardEngine()48_env.attach_reward_engine(_engine)49_initialized = False50 51GRADERS = {52    "peaceful": PeacefulGrader,53    "easy": EasyGrader,54    "medium": MediumGrader,55    "hard": HardGrader,56    "extreme": ExtremeGrader,57}58 59# ── Request / Response models ─────────────────────────────────────────────────60 61class ResetRequest(PydanticBaseModel):62    task_id: str = "easy"  # easy | medium | hard63    seed: Optional[int] = None64    email_count: Optional[int] = None65    max_steps: Optional[int] = None66 67 68class StepRequest(PydanticBaseModel):69    action_type: str70    email_id: str71    category: Optional[str] = None72    level: Optional[str] = None73    text: Optional[str] = None74 75 76# ── Endpoints ─────────────────────────────────────────────────────────────────77 78@app.get("/health")79async def health():80    """Liveness / readiness probe."""81    return {"status": "ok", "service": "SHUBHAMOS", "version": "1.0.0"}82 83 84@app.post("/reset")85async def reset(req: Optional[ResetRequest] = None):86    """87    Start a new episode.88 89    Uses task presets from openenv.yaml but allows override of seed/email_count/max_steps.90    """91    global _initialized92 93    # Handle empty body by using defaults94    if req is None:95        req = ResetRequest(task_id="easy")96 97    if req.task_id not in TASKS:98        raise HTTPException(400, f"Unknown task_id '{req.task_id}'. Choose: easy, medium, hard")99 100    task_cls = TASKS[req.task_id]101    task_config = task_cls.config()102 103    # Allow caller to override task defaults104    if req.seed is not None:105        task_config["seed"] = req.seed106    if req.email_count is not None:107        task_config["email_count"] = req.email_count108    if req.max_steps is not None:109        task_config["max_steps"] = req.max_steps110 111    obs = _env.reset(task_config)112    _initialized = True113 114    return JSONResponse(content=_obs_to_json(obs))115 116 117@app.post("/step")118async def step(req: StepRequest):119    """Apply one action and return (observation, reward, done, info)."""120    if not _initialized:121        raise HTTPException(400, "Call /reset first to start an episode")122 123    try:124        action = Action(125            action_type=req.action_type,126            email_id=req.email_id,127            category=req.category,128            level=req.level,129            text=req.text,130        )131    except Exception as e:132        raise HTTPException(422, f"Invalid action: {e}")133 134    try:135        obs, reward, done, info = _env.step(action)136    except RuntimeError as e:137        raise HTTPException(400, str(e))138 139    return JSONResponse(content={140        "observation": _obs_to_json(obs),141        "reward": reward,142        "done": done,143        "info": info,144    })145 146 147@app.get("/state")148async def state():149    """Return full internal state including ground truth labels (for graders)."""150    if not _initialized:151        raise HTTPException(400, "No active episode. Call /reset first.")152 153    s = _env.state()154    # Serialize via pydantic155    return JSONResponse(content=json.loads(s.model_dump_json()))156 157 158@app.get("/grade")159async def grade():160    """Grade the current episode using the deterministic grader."""161    if not _initialized:162        raise HTTPException(400, "No active episode. Call /reset first.")163 164    s = _env.state()165    if s.task_id not in GRADERS:166        raise HTTPException(400, f"No grader for task_id '{s.task_id}'")167 168    grader = GRADERS[s.task_id]()169    report = grader.grade(s)170    return JSONResponse(content=report.to_dict())171 172 173@app.get("/tasks")174async def list_tasks():175    """List available tasks with their configurations."""176    task_list = []177    for tid, cls in TASKS.items():178        task_list.append({179            "task_id": cls.task_id,180            "difficulty": cls.difficulty,181            "seed": cls.seed,182            "email_count": cls.email_count,183            "max_steps": cls.max_steps,184            "description": cls.description(),185            "grader": True,186            "grader_name": cls.grader_class().__name__187        })188    return JSONResponse(content=task_list)189 190 191@app.get("/openenv.yaml", response_class=HTMLResponse)192async def openenv_yaml():193    """Serve the openenv.yaml spec file."""194    yaml_path = Path(__file__).parent / "openenv.yaml"195    if yaml_path.exists():196        return HTMLResponse(content=yaml_path.read_text(), media_type="text/yaml")197    raise HTTPException(404, "openenv.yaml not found")198 199 200# ── Dashboard ─────────────────────────────────────────────────────────────────201 202DASHBOARD_DIR = Path(__file__).parent / "dashboard"203 204@app.get("/")205async def root_health():206    """Root health check for Hugging Face Spaces."""207    return {208        "status": "ok",209        "message": "SHUBHAMOS is running 🚀"210    }211 212 213@app.get("/dashboard", response_class=HTMLResponse)214async def dashboard():215    """Serve the minimal monitoring dashboard."""216    html_path = DASHBOARD_DIR / "index.html"217    if html_path.exists():218        return HTMLResponse(content=html_path.read_text())219    return HTMLResponse(content="""220    <html><body>221    <h1>SHUBHAMOS — Email Triage Environment</h1>222    <p>API is running. Dashboard not found.</p>223    <p>Endpoints: <a href="/docs">/docs</a> | <a href="/health">/health</a></p>224    </body></html>225    """)226 227 228# ── Serialization helper ──────────────────────────────────────────────────────229 230def _obs_to_json(obs) -> Dict[str, Any]:231    """Convert Observation to JSON-serializable dict."""232    return json.loads(obs.model_dump_json())233 234 235# ── Entry point ───────────────────────────────────────────────────────────────236 237if __name__ == "__main__":238    from .app import start_server239    start_server()240