sentinel-env / server /app.py
rudrapatel-1908's picture
Update server/app.py
d8ba59b verified
Raw
History Blame Contribute Delete
7.63 kB
import sys
import os
sys.path.insert(0, os.path.dirname(__file__))
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
import json
import uuid
import uvicorn
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import HTMLResponse
from pydantic import BaseModel
from typing import Optional, Dict
from models import (
SentinelAction, SentinelObservation, SentinelState,
MultiAgentAction, AttackerAgentAction, SentinelAgentAction,
)
from sentinel_env_environment import SentinelEnv, WarRoomEnv
app = FastAPI(
title="Sentinel Env",
description="Red vs Blue Multi-Agent Cybersecurity War Room β€” OpenEnv 2026",
)
# ── Single-agent env (legacy tasks) ──
env = SentinelEnv()
# ── Session-isolated war room instances (one per WebSocket session) ──
_war_room_sessions: Dict[str, WarRoomEnv] = {}
# ── Load HTML template ──
_TEMPLATE_PATH = os.path.join(os.path.dirname(__file__), "templates", "index.html")
with open(_TEMPLATE_PATH, "r", encoding="utf-8") as f:
_HTML = f.read()
class ResetRequest(BaseModel):
task_id: Optional[str] = "easy-lockdown"
class WarRoomResetRequest(BaseModel):
task_id: Optional[str] = "red-vs-blue"
# ─────────────────────────────────────────────
# UI & HEALTH
# ─────────────────────────────────────────────
@app.get("/", response_class=HTMLResponse)
def root():
return _HTML
@app.get("/health")
def health():
return {"status": "ok"}
@app.get("/metadata")
def metadata():
data = env.get_metadata()
data["current_task"] = env._state.current_task_id
data["step_count"] = env._state.step_count
data["active_sessions"] = len(_war_room_sessions)
return data
# ─────────────────────────────────────────────
# SINGLE-AGENT ENDPOINTS (legacy)
# ─────────────────────────────────────────────
@app.post("/reset")
def reset(request: Optional[ResetRequest] = None):
task_id = request.task_id if request and request.task_id else "easy-lockdown"
obs = env.reset(task_id=task_id)
return {"observation": obs.dict()}
@app.post("/step")
def step(action: SentinelAction):
obs, reward, done, info = env.step(action)
return {"observation": obs.dict(), "reward": reward, "done": done, "info": info}
@app.get("/state")
def state():
return env._state.dict()
# ─────────────────────────────────────────────
# MULTI-AGENT HTTP ENDPOINTS
# ─────────────────────────────────────────────
@app.post("/warroom/reset")
def warroom_reset(request: Optional[WarRoomResetRequest] = None):
session_id = str(uuid.uuid4())
wr = WarRoomEnv()
_war_room_sessions[session_id] = wr
task_id = request.task_id if request and request.task_id else "red-vs-blue"
obs = wr.reset(task_id=task_id)
return {"session_id": session_id, **obs}
@app.post("/warroom/step")
def warroom_step(actions: MultiAgentAction, session_id: Optional[str] = None):
# Use default session if none provided
if session_id and session_id in _war_room_sessions:
wr = _war_room_sessions[session_id]
elif _war_room_sessions:
wr = list(_war_room_sessions.values())[-1]
else:
wr = WarRoomEnv()
sid = str(uuid.uuid4())
_war_room_sessions[sid] = wr
wr.reset()
obs, rewards, done, info = wr.step(actions)
return {"observation": obs, "rewards": rewards, "done": done, "info": info}
@app.get("/warroom/state")
def warroom_state(session_id: Optional[str] = None):
if session_id and session_id in _war_room_sessions:
return _war_room_sessions[session_id].get_state().dict()
if _war_room_sessions:
return list(_war_room_sessions.values())[-1].get_state().dict()
return {"error": "No active war room session"}
@app.get("/warroom/log")
def warroom_log(session_id: Optional[str] = None):
if session_id and session_id in _war_room_sessions:
return {"war_log": _war_room_sessions[session_id]._war_log}
if _war_room_sessions:
return {"war_log": list(_war_room_sessions.values())[-1]._war_log}
return {"war_log": []}
# ─────────────────────────────────────────────
# WEBSOCKET β€” Session-Isolated Multi-Agent
# Per OpenEnv 2026 spec: /ws for persistent sessions
# ─────────────────────────────────────────────
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
await websocket.accept()
session_id = str(uuid.uuid4())
wr = WarRoomEnv()
_war_room_sessions[session_id] = wr
await websocket.send_json({
"type": "connected",
"session_id": session_id,
"message": "War Room session started. Send {type: 'reset'} to begin.",
})
try:
while True:
data = await websocket.receive_json()
msg_type = data.get("type", "")
if msg_type == "reset":
task_id = data.get("task_id", "red-vs-blue")
obs = wr.reset(task_id=task_id)
await websocket.send_json({
"type": "reset_ok",
"session_id": session_id,
**obs,
})
elif msg_type == "step":
try:
actions = MultiAgentAction(
attacker=AttackerAgentAction(**data["attacker"]),
scanner=SentinelAgentAction(**data["scanner"]),
remediator=SentinelAgentAction(**data["remediator"]),
)
obs, rewards, done, info = wr.step(actions)
await websocket.send_json({
"type": "step_ok",
"session_id": session_id,
"observation": obs,
"rewards": rewards,
"done": done,
"info": info,
})
if done:
await websocket.send_json({
"type": "episode_done",
"session_id": session_id,
"final_state": wr.get_state().dict(),
})
except Exception as e:
await websocket.send_json({"type": "error", "message": str(e)})
elif msg_type == "state":
await websocket.send_json({
"type": "state",
"state": wr.get_state().dict(),
})
elif msg_type == "log":
await websocket.send_json({
"type": "log",
"war_log": wr._war_log,
})
elif msg_type == "ping":
await websocket.send_json({"type": "pong"})
except WebSocketDisconnect:
_war_room_sessions.pop(session_id, None)
def main():
uvicorn.run("server.app:app", host="0.0.0.0", port=7860)
if __name__ == "__main__":
main()