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()