215 lines
7.4 KiB
Python
215 lines
7.4 KiB
Python
#VNC Proxy server
|
|
import asyncio
|
|
import base64
|
|
import json
|
|
import logging
|
|
import uuid
|
|
from aiohttp import web
|
|
import requests
|
|
import socketio
|
|
import os
|
|
|
|
# Config
|
|
# Config from environment
|
|
HTTP_PORT = int(os.getenv("VNC_PROXY_HTTP_PORT", 6002))
|
|
VNC_PORT_LOCAL = int(os.getenv("VNC_PORT_LOCAL", 5900))
|
|
SOCKET_SERVER_URL = os.getenv("SOCKET_SERVER_URL", "http://172.17.0.1:6001")
|
|
CHECK_PERMISSION_URL = os.getenv("CHECK_PERMISSION_URL", "http://localhost:5000/api/check_permission")
|
|
ENGINEIO_LOGGER_LEVEL = os.getenv("ENGINEIO_LOGGER_LEVEL", "WARNING").upper()
|
|
SIO_LOGGER_LEVEL = os.getenv("SIO_LOGGER_LEVEL", "WARNING").upper()
|
|
|
|
logging.basicConfig(level=logging.DEBUG, format="%(asctime)s [%(levelname)s] %(message)s")
|
|
logger = logging.getLogger("middleware")
|
|
|
|
# Apply dynamic log levels from environment
|
|
sio_logger = logging.getLogger("socketio.client")
|
|
sio_logger.setLevel(getattr(logging, SIO_LOGGER_LEVEL, logging.WARNING))
|
|
|
|
engineio_logger = logging.getLogger("engineio.client")
|
|
engineio_logger.setLevel(getattr(logging, ENGINEIO_LOGGER_LEVEL, logging.WARNING))
|
|
|
|
|
|
# sio = socketio.AsyncServer(async_mode="aiohttp", cors_allowed_origins="*")
|
|
app = web.Application()
|
|
# sio.attach(app)
|
|
sio = socketio.AsyncClient(logger=sio_logger, engineio_logger=engineio_logger)
|
|
|
|
sessions = {} # vnc_request_id → WebSocket
|
|
routes = web.RouteTableDef()
|
|
|
|
|
|
@sio.event
|
|
async def connect():
|
|
logger.info("Connected to the server, requesting to join.")
|
|
|
|
@sio.event
|
|
async def disconnect(sid):
|
|
|
|
logger.warning("Worker disconnected")
|
|
|
|
@sio.on("vnc_frame_from_worker")
|
|
async def vnc_frame_from_worker(payload):
|
|
logger.debug(f"Got a frame from worker")
|
|
vnc_request_id = payload.get('vnc_request_id')
|
|
|
|
if not vnc_request_id:
|
|
logger.warning("Received VNC frame without vnc_request_id")
|
|
return
|
|
|
|
# Find the specific WebSocket for this vnc_request_id
|
|
ws = sessions.get(vnc_request_id)
|
|
|
|
if ws is None:
|
|
logger.debug(f"No active session found for vnc_request_id {vnc_request_id}")
|
|
return
|
|
|
|
if ws.closed:
|
|
logger.debug(f"WebSocket is closed for vnc_request_id {vnc_request_id}, removing session")
|
|
sessions.pop(vnc_request_id, None)
|
|
return
|
|
|
|
# WebSocket exists and is open, send the frame
|
|
try:
|
|
await ws.send_bytes(payload['data'])
|
|
except Exception as e:
|
|
logger.error(f"Failed to send VNC frame to client {vnc_request_id}: {e}")
|
|
# Clean up the failed session
|
|
sessions.pop(vnc_request_id, None)
|
|
|
|
def check_permission(resource_id, resource_type, action, user_token):
|
|
logger.debug(f"Checking permission for resource_id={resource_id}, resource_type={resource_type}, "
|
|
f"action={action}, user_token={user_token[:10]}...")
|
|
|
|
try:
|
|
payload = {
|
|
"user_token": user_token,
|
|
"resource_id": resource_id,
|
|
"resource_type": resource_type,
|
|
"action": action
|
|
}
|
|
response = requests.post(CHECK_PERMISSION_URL, json=payload, timeout=5)
|
|
response.raise_for_status()
|
|
allowed = response.json().get("allowed", False)
|
|
logger.debug(f"Permission check result: allowed={allowed}")
|
|
return allowed
|
|
except requests.RequestException as e:
|
|
logger.warning(f"Permission check failed: {e}")
|
|
return False
|
|
|
|
@routes.get("/vnc")
|
|
async def vnc_handler(request):
|
|
logger.debug("Received WebSocket connection request on /vnc")
|
|
|
|
token_b64 = request.query.get("token")
|
|
if not token_b64:
|
|
logger.error("Missing token in request")
|
|
return web.Response(status=400, text="Missing token")
|
|
|
|
try:
|
|
# Proper base64 padding
|
|
padded_token_b64 = token_b64 + '=' * (-len(token_b64) % 4)
|
|
token_json = base64.urlsafe_b64decode(padded_token_b64).decode("utf-8")
|
|
token_data = json.loads(token_json)
|
|
virtual_machine_id = token_data.get("virtual_machine_id")
|
|
user_token = token_data.get("user_token")
|
|
|
|
if not virtual_machine_id or not user_token:
|
|
raise ValueError("Missing required fields in token")
|
|
|
|
logger.debug(f"Parsed token: VM={virtual_machine_id}, user_token[10]={user_token[:10]}...")
|
|
|
|
except Exception as e:
|
|
logger.warning(f"Invalid token: {e}")
|
|
return web.Response(status=400, text="Invalid token")
|
|
|
|
# Delegate permission check (this function should internally verify the user_token)
|
|
allowed = check_permission(virtual_machine_id, "virtual_machine", "vnc_console", user_token)
|
|
if not allowed:
|
|
logger.warning(f"Permission denied for VNC access to VM {virtual_machine_id}")
|
|
return web.Response(status=403, text="Permission denied")
|
|
|
|
ws = web.WebSocketResponse(protocols=["binary"])
|
|
await ws.prepare(request)
|
|
|
|
req_id = uuid.uuid4().hex
|
|
sessions[req_id] = ws
|
|
logger.info(f"NoVNC client connected (req_id={req_id}, vm_id={virtual_machine_id})")
|
|
|
|
await sio.emit("start_vnc_stream", {
|
|
"vnc_request_id": req_id,
|
|
"user_token": user_token,
|
|
"virtual_machine_id": virtual_machine_id
|
|
})
|
|
|
|
try:
|
|
async for msg in ws:
|
|
if msg.type == web.WSMsgType.BINARY:
|
|
await sio.emit("vnc_frame_from_novnc", {
|
|
"vnc_request_id": req_id,
|
|
"data": msg.data
|
|
})
|
|
logger.debug(f"Got a frame from novnc")
|
|
elif msg.type == web.WSMsgType.ERROR:
|
|
logger.error(f"WebSocket error: {ws.exception()}")
|
|
break
|
|
elif msg.type == web.WSMsgType.CLOSE:
|
|
logger.info(f"NoVNC WebSocket closed normally for session {req_id}")
|
|
break
|
|
except Exception as e:
|
|
logger.exception(f"WebSocket communication error: {e}")
|
|
finally:
|
|
logger.info(f"NoVNC session {req_id} disconnected")
|
|
# Stop the VNC stream on the worker
|
|
try:
|
|
await sio.emit("stop_vnc_stream", {"vnc_request_id": req_id})
|
|
logger.debug(f"Sent stop_vnc_stream for session {req_id}")
|
|
except Exception as e:
|
|
logger.error(f"Failed to send stop_vnc_stream for session {req_id}: {e}")
|
|
|
|
# Clean up the session
|
|
sessions.pop(req_id, None)
|
|
|
|
# Close the WebSocket if not already closed
|
|
if not ws.closed:
|
|
await ws.close()
|
|
return ws
|
|
|
|
|
|
@routes.get("/healthz")
|
|
async def health_check(request):
|
|
"""
|
|
Healthcheck endpoint used for container liveness checks.
|
|
Returns HTTP 200 with a simple message.
|
|
"""
|
|
return web.Response(status=200, text="ok")
|
|
|
|
|
|
app.add_routes(routes)
|
|
|
|
async def connect_with_retry(max_retries=10):
|
|
delay = 1
|
|
for attempt in range(1, max_retries + 1):
|
|
try:
|
|
logger.info(f"Connecting to WebSocket server (attempt {attempt})...")
|
|
await sio.connect(SOCKET_SERVER_URL)
|
|
logger.info("WebSocket connection established")
|
|
return
|
|
except Exception as e:
|
|
logger.warning(f"WebSocket connection failed: {e}")
|
|
if attempt == max_retries:
|
|
logger.error("Max retries exceeded. Giving up.")
|
|
raise SystemExit(1)
|
|
await asyncio.sleep(delay)
|
|
delay *= 2
|
|
|
|
async def main():
|
|
runner = web.AppRunner(app)
|
|
await runner.setup()
|
|
await connect_with_retry()
|
|
site = web.TCPSite(runner, host="0.0.0.0", port=HTTP_PORT)
|
|
await site.start()
|
|
logger.info(f"NoVNC WebSocket at ws://0.0.0.0:{HTTP_PORT}/vnc")
|
|
await asyncio.Event().wait()
|
|
|
|
if __name__ == "__main__":
|
|
asyncio.run(main()) |