mohdbelal010/SecureAI-Gaurd
0
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 