Spaces:
Sleeping
Sleeping
| 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 | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| def root(): | |
| return _HTML | |
| def health(): | |
| return {"status": "ok"} | |
| 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) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| 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()} | |
| def step(action: SentinelAction): | |
| obs, reward, done, info = env.step(action) | |
| return {"observation": obs.dict(), "reward": reward, "done": done, "info": info} | |
| def state(): | |
| return env._state.dict() | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # MULTI-AGENT HTTP ENDPOINTS | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| 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} | |
| 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} | |
| 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"} | |
| 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 | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| 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() |