CoolFace
Modelpublic

mohdbelal010/SecureAI-Gaurd

sourceHugging Faceapache-2.0updated 6mo agoView on Hugging Face
0likes
app.py194 linesDownload Raw Back to root
1"""2SecureAI-Guard: FastAPI environment server.3Exposes /reset, /step, /state and auxiliary endpoints.4"""5import logging6import os7import sys8 9import uvicorn10from fastapi import FastAPI, HTTPException11from fastapi.middleware.cors import CORSMiddleware12from fastapi.responses import FileResponse13from fastapi.staticfiles import StaticFiles14from pydantic import BaseModel15from typing import Optional, Dict, Any, List16 17from env.engine import SecureAIGuardEngine18from tasks.registry import TaskRegistry19from graders.security_grader import SecurityGrader20from schema.models import Action, StepResponse, Observation, State21 22logging.basicConfig(23    level=logging.INFO,24    format="%(asctime)s %(levelname)s %(name)s — %(message)s",25    handlers=[logging.StreamHandler(sys.stdout)],26)27logger = logging.getLogger("secureai-guard")28 29app = FastAPI(30    title="SecureAI-Guard API",31    description="Stateful POMDP for Autonomous Digital Defense — OpenEnv compliant",32    version="1.0.0",33)34 35app.add_middleware(36    CORSMiddleware,37    allow_origins=["*"],38    allow_credentials=True,39    allow_methods=["*"],40    allow_headers=["*"],41)42 43# ---------------------------------------------------------------------------44# Singletons45# ---------------------------------------------------------------------------46env = SecureAIGuardEngine()47task_registry = TaskRegistry()48grader = SecurityGrader()49_episode_rewards: List = []50_current_task_id: str = "basic_security"51 52 53# ---------------------------------------------------------------------------54# Request / Response models55# ---------------------------------------------------------------------------56class ResetRequest(BaseModel):57    task_id: Optional[str] = "basic_security"58    seed: Optional[int] = None59 60 61class StepRequest(BaseModel):62    action: Action63 64 65class GradeRequest(BaseModel):66    task_id: Optional[str] = "basic_security"67 68 69# ---------------------------------------------------------------------------70# Routes71# ---------------------------------------------------------------------------72@app.get("/")73async def root():74    return {75        "name": "SecureAI-Guard",76        "version": "1.0.0",77        "status": "running",78        "endpoints": ["/reset", "/step", "/state", "/tasks", "/grade", "/preference_data", "/health"],79    }80 81 82@app.get("/health")83async def health():84    return {"status": "healthy", "version": "1.0.0"}85 86 87@app.post("/reset")88async def reset(request: ResetRequest):89    """Reset environment. Returns first observation."""90    global _episode_rewards, _current_task_id91 92    _current_task_id = request.task_id or "basic_security"93    _episode_rewards = []94 95    try:96        observation = env.reset(seed=request.seed, task_id=_current_task_id)97        logger.info("Episode reset | task=%s seed=%s", _current_task_id, request.seed)98        return {99            "observation": observation.model_dump(),100            "state": env.get_state().model_dump(),101            "task_id": _current_task_id,102        }103    except Exception as exc:104        logger.exception("Reset failed")105        raise HTTPException(status_code=500, detail=str(exc))106 107 108@app.post("/step")109async def step(request: StepRequest):110    """Execute one environment step."""111    global _episode_rewards112 113    try:114        response: StepResponse = env.step(request.action)115        _episode_rewards.append(response.reward)116 117        result = response.model_dump()118 119        if response.done:120            grade = grader.grade_episode(response.state, _episode_rewards, _current_task_id)121            result["grade"] = grade.model_dump()122            logger.info(123                "Episode done | score=%.4f grade=%s steps=%d",124                grade.score, grade.grade, response.state.step_count,125            )126 127        return result128    except Exception as exc:129        logger.exception("Step failed")130        raise HTTPException(status_code=500, detail=str(exc))131 132 133@app.get("/state")134async def get_state():135    """Get current environment state."""136    return env.get_state().model_dump()137 138 139@app.get("/tasks")140async def list_tasks():141    """List all available tasks."""142    return {"tasks": [t.model_dump() for t in task_registry.list_tasks()]}143 144 145@app.get("/tasks/{task_id}")146async def get_task(task_id: str):147    """Get a specific task definition."""148    try:149        return task_registry.get_task(task_id).model_dump()150    except ValueError as exc:151        raise HTTPException(status_code=404, detail=str(exc))152 153 154@app.post("/grade")155async def grade_episode(request: GradeRequest):156    """Grade the current episode explicitly."""157    state = env.get_state()158    grade = grader.grade_episode(state, _episode_rewards, request.task_id or _current_task_id)159    return grade.model_dump()160 161 162@app.get("/preference_data")163async def get_preference_data():164    """Get logged DPO preference pairs."""165    return {"preference_pairs": env.get_preference_data()}166 167 168# ---------------------------------------------------------------------------169# Serve frontend170# ---------------------------------------------------------------------------171FRONTEND_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "frontend")172 173 174@app.get("/dashboard")175async def dashboard():176    return FileResponse(os.path.join(FRONTEND_DIR, "index.html"))177 178 179app.mount("/static", StaticFiles(directory=FRONTEND_DIR), name="frontend")180 181 182if __name__ == "__main__":183    import webbrowser184    import threading185 186    def open_browser():187        """Open the dashboard in the default browser after a short delay."""188        import time189        time.sleep(1.5)190        webbrowser.open("http://localhost:7860/dashboard")191 192    threading.Thread(target=open_browser, daemon=True).start()193    uvicorn.run("app:app", host="0.0.0.0", port=7860, log_level="info")194