diff --git a/websocket_server/__init__.py b/websocket_server/__init__.py new file mode 100644 index 0000000..097d2ae --- /dev/null +++ b/websocket_server/__init__.py @@ -0,0 +1,9 @@ +""" +websocket_server package initializer. + +This package provides a WebSocket server for managing task dispatch, +container and virtual machine event streaming, terminal sessions, +VNC proxying, and real-time communication between workers and users. +""" + +__version__ = "0.1.0" diff --git a/websocket_server/config.py b/websocket_server/config.py new file mode 100644 index 0000000..b0382a1 --- /dev/null +++ b/websocket_server/config.py @@ -0,0 +1,107 @@ +# websocket_server/config.py + +import os +import redis +import logging +from logging import StreamHandler, FileHandler +from colorlog import ColoredFormatter +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker, declarative_base +import pymysql + +# Redis configuration +REDIS_HOST = "localhost" +REDIS_PORT = 6379 + +# Database configuration +DATABASE_URL = "mysql://root:password@172.17.0.1:3306/defaultdb" + +# Connection pool for Redis +redis_connection_pool = redis.ConnectionPool( + host=REDIS_HOST, + port=REDIS_PORT, + decode_responses=True, + max_connections=10, + socket_timeout=5 +) + +def get_redis_client(): + return redis.StrictRedis(connection_pool=redis_connection_pool) + +# SQLAlchemy setup +pymysql.install_as_MySQLdb() +engine = create_engine( + DATABASE_URL, + isolation_level='SERIALIZABLE', + pool_size=5, + max_overflow=10, + pool_timeout=30, + pool_recycle=3600, + pool_pre_ping=True +) + +SessionLocal = sessionmaker( + autocommit=False, + autoflush=False, + bind=engine, + expire_on_commit=False +) + +Base = declarative_base() + +# Logging setup +def setup_logging(): + class FunctionNameFilter(logging.Filter): + def filter(self, record): + record.funcName = getattr(record, 'funcName', '') + return True + + log_format = ( + "%(log_color)s%(asctime)s - %(levelname)s - %(funcName)s - %(message)s" + ) + date_format = "%Y-%m-%d %H:%M:%S" + + formatter = ColoredFormatter( + log_format, + datefmt=date_format, + log_colors={ + "DEBUG": "cyan", + "INFO": "green", + "WARNING": "yellow", + "ERROR": "red", + "CRITICAL": "bold_red", + }, + ) + + logger = logging.getLogger("websocket_server") + logger.setLevel(logging.DEBUG) + logger.addFilter(FunctionNameFilter()) + + stream_handler = StreamHandler() + stream_handler.setFormatter(formatter) + logger.addHandler(stream_handler) + + # Ensure logs directory exists + log_dir = "logs" + os.makedirs(log_dir, exist_ok=True) + + file_handler = FileHandler("logs/app.log") + file_handler.setFormatter(formatter) + file_handler.setLevel(logging.WARNING) + logger.addHandler(file_handler) + + return logger + +# Context manager for DB session +from contextlib import contextmanager +@contextmanager +def get_db_session(): + session = SessionLocal() + try: + yield session + session.commit() + except Exception: + session.rollback() + raise + finally: + session.close() diff --git a/websocket_server/events/__init__.py b/websocket_server/events/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/websocket_server/events/ack.py b/websocket_server/events/ack.py new file mode 100644 index 0000000..27a14ea --- /dev/null +++ b/websocket_server/events/ack.py @@ -0,0 +1,48 @@ +# websocket_server/events/ack.py + +from flask_socketio import SocketIO +import logging +from datetime import datetime +import json + +from websocket_server.events import base +from websocket_server.config import get_db_session, get_redis_client +from websocket_server.models import Task + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + +def register_socketio_handlers(socketio): + + @socketio.on("ack") + def handle_ack(data): + """ + Worker acknowledges task completion. + Updates database with result and execution timing. + """ + worker_id = data.get("worker_id") + task_id = data.get("task_id") + result = data.get("result", {}) + + logger.debug(f"[{worker_id}] Acknowledged task {task_id}, result: {result}") + + try: + with get_db_session() as session: + task = session.query(Task).filter(Task.id == task_id).first() + if not task: + logger.warning(f"[{worker_id}] Task {task_id} not found in DB") + return + + task.status = "acknowledged" + task.success = 1 if result.get("success") else 0 + task.response = json.dumps(result.get("response", "")) + task.finish_time = datetime.utcnow() + task.execution_time = (task.finish_time - task.start_time).total_seconds() + + logger.info(f"[{worker_id}] Task {task_id} acknowledged. Exec time: {task.execution_time:.2f}s") + + redis_client = get_redis_client() + redis_client.set(f"worker_status_{worker_id}", "idle") + + except Exception as e: + logger.exception(f"[{worker_id}] Failed to acknowledge task {task_id}: {e}") diff --git a/websocket_server/events/base.py b/websocket_server/events/base.py new file mode 100644 index 0000000..83ce2ee --- /dev/null +++ b/websocket_server/events/base.py @@ -0,0 +1,6 @@ +# websocket_server/events/base.py + +# import threading + +# # SocketIO will be initialized in main.py and shared here +socketio = None diff --git a/websocket_server/events/docker.py b/websocket_server/events/docker.py new file mode 100644 index 0000000..8e081ed --- /dev/null +++ b/websocket_server/events/docker.py @@ -0,0 +1,119 @@ +# websocket_server/events/docker.py + +from flask_socketio import SocketIO +from datetime import datetime +import logging +import requests +import json + +from websocket_server.events import base + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + +api_server_url = "http://127.0.0.1:5000/api" + +def register_socketio_handlers(socketio): + + + @socketio.on("docker_event") + def handle_docker_event(data): + """ + Handle Docker events received from the worker. + Extract relevant data and log it to the console. + """ + try: + # Extract relevant data from the event + worker_id = data.get("worker_id") + event_type = data.get("type") + details = data.get("details", {}) + + container_id = details.get("container_id") + system_container_id = details.get("system_container_id") + container_name = details.get("container_name") + unix_timestamp = details.get("timestamp") + timestamp=datetime.utcfromtimestamp(unix_timestamp).isoformat() + status = details.get("details", {}).get("status") + image = details.get("details", {}).get("from") + action = details.get("details", {}).get("Action") + + # Log the extracted data + logger.info(f"Received Docker event from worker {worker_id}:") + logger.info(f" Event Type: {event_type}") + logger.info(f" Docker Container ID: {container_id}") + logger.info(f" System Container ID: {system_container_id}") + logger.info(f" Container Name: {container_name}") + logger.info(f" Timestamp: {timestamp}") + logger.info(f" Status: {status}") + logger.info(f" Image: {image}") + logger.info(f" Action: {action}") + + # Optionally, log the full details for debugging purposes + logger.debug(f"Full event details: {data}") + + if event_type=="docker_start": + #Send the update to the API Server + # system_container_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/containers/status_update/{system_container_id}" + logger.debug(f"Sending task update payload to API server for Container ID {system_container_id}") + + payload = { + "new_status": "running", + "system_container_id": system_container_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + if event_type=="docker_destroy": + #Send the update to the API Server + # system_container_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/containers/status_update/{system_container_id}" + logger.debug(f"Sending task update payload to API server for Container ID {system_container_id}") + + payload = { + "new_status": "deleted", + "system_container_id": system_container_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + if event_type=="docker_die": + #Send the update to the API Server + # system_container_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/containers/status_update/{system_container_id}" + logger.debug(f"Sending task update payload to API server for Container ID {system_container_id}") + + payload = { + "new_status": "dead", + "system_container_id": system_container_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + except Exception as e: + logger.error(f"Error processing Docker event: {e}") + diff --git a/websocket_server/events/libvirt.py b/websocket_server/events/libvirt.py new file mode 100644 index 0000000..8dc61d3 --- /dev/null +++ b/websocket_server/events/libvirt.py @@ -0,0 +1,111 @@ +# websocket_server/events/libvirt.py + +from flask_socketio import SocketIO +from datetime import datetime +import logging +import requests +import json + +from websocket_server.events import base + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + +api_server_url = "http://127.0.0.1:5000/api" + +def register_socketio_handlers(socketio): + + @socketio.on("libvirt_event") + def handle_libvirt_event(data): + """ + Handle libvirt events received from the worker. + Extract relevant data and log it to the console. + """ + try: + # Extract relevant data from the event + worker_id = data.get("worker_id") + event_type = data.get("type") + details = data.get("details", {}) + + libvirt_vm_id = details.get("libvirt_vm_id") + system_vm_id = details.get("system_vm_id") + unix_timestamp = details.get("timestamp") + timestamp=datetime.utcfromtimestamp(unix_timestamp).isoformat() + + + # Log the extracted data + logger.info(f"Received libvirt event from worker {worker_id}:") + logger.info(f" Event Type: {event_type}") + logger.info(f" Libvirt VM ID: {libvirt_vm_id}") + logger.info(f" System VM ID: {system_vm_id}") + logger.info(f" Timestamp: {timestamp}") + + # Optionally, log the full details for debugging purposes + logger.debug(f"Full event details: {data}") + + if event_type=="libvirt_started": + #Send the update to the API Server + # system_vm_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/virtual_machines/status_update/{system_vm_id}" + logger.debug(f"Sending task update payload to API server for VM ID {system_vm_id}") + + payload = { + "new_status": "running", + "system_vm_id": system_vm_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + if event_type=="libvirt_stopped": + #Send the update to the API Server + # system_vm_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/virtual_machines/status_update/{system_vm_id}" + logger.debug(f"Sending task update payload to API server for VM ID {system_vm_id}") + + payload = { + "new_status": "stopped", + "system_vm_id": system_vm_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + + if event_type=="libvirt_undefined": + #Send the update to the API Server + # system_vm_id + # worker_id + # state="started" + # timestamp + action_URL=f"workloads/virtual_machines/status_update/{system_vm_id}" + logger.debug(f"Sending task update payload to API server for VM ID {system_vm_id}") + + payload = { + "new_status": "deleted", + "system_vm_id": system_vm_id, + "worker_id": worker_id, + "timestamp": timestamp + } + + headers = {"Content-Type": "application/json"} + websocket_server_response = requests.put(f"{api_server_url}/{action_URL}", data=json.dumps(payload), headers=headers) + logger.debug( websocket_server_response) + logger.info("Update complete") + + except Exception as e: + logger.error(f"Error processing libvirt event: {e}") diff --git a/websocket_server/events/logs.py b/websocket_server/events/logs.py new file mode 100644 index 0000000..030a4fa --- /dev/null +++ b/websocket_server/events/logs.py @@ -0,0 +1,184 @@ +# websocket_server/events/logs.py + +from flask import request +from flask_socketio import emit +import logging +import json +import requests + +from websocket_server.config import get_redis_client +from websocket_server.events import base +from websocket_server.shared_state import connected_sids, connected_sids_lock, connected_workers + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + + +api_server_url = "http://127.0.0.1:5000/api" +def register_socketio_handlers(socketio): + + @socketio.on("user_request_container_logs") + def handle_user_log_request(data): + container_id = data.get("container_id") + user_id = data.get("user_id") + lines = data.get("lines", 100) + request_id = data.get("request_id") + + logger.info(f"[{request_id}] User {user_id} requested logs for container {container_id}") + + try: + container_response = requests.get(f"{api_server_url}/workloads/containers/{container_id}") + if container_response.status_code != 200: + emit("user_log_response", { + "success": False, + "logs": "Container not found", + "request_id": request_id + }, to=request.sid) + return + + container_info = container_response.json() + worker_id = container_info["workload_host_id"] + + if worker_id not in connected_workers: + emit("user_log_response", { + "success": False, + "logs": "Worker not connected", + "request_id": request_id + }, to=request.sid) + return + + redis_client = get_redis_client() + request_context = { + "user_sid": request.sid, + "request_id": request_id, + "container_id": container_id, + "user_id": user_id, + "worker_id": worker_id + } + redis_client.setex(f"log_request_context:{request_id}", 60, json.dumps(request_context)) + + socketio.emit("container-log", { + "container_name": container_info["id"], + "lines": lines, + "request_id": request_id, + "worker_id": worker_id + }, to=connected_workers[worker_id]) + + logger.info(f"[{request_id}] Dispatched log request to worker {worker_id}") + + except Exception as e: + logger.exception(f"[{request_id}] Error processing log request") + emit("user_log_response", { + "success": False, + "logs": "Internal error", + "request_id": request_id + }, to=request.sid) + + + @socketio.on("container_log_response") + def handle_worker_log_response(data): + request_id = data.get("request_id") + redis_client = get_redis_client() + context_data = redis_client.get(f"log_request_context:{request_id}") + + if not context_data: + logger.warning(f"[{request_id}] No context found") + return + + try: + context = json.loads(context_data) + user_sid = context.get("user_sid") + socketio.emit("user_log_response", { + "container_id": context["container_id"], + "logs": data.get("logs"), + "success": data.get("success", False), + "request_id": request_id + }, to=user_sid) + redis_client.delete(f"log_request_context:{request_id}") + logger.info(f"[{request_id}] Delivered one-shot logs to user") + except Exception as e: + logger.exception(f"[{request_id}] Error handling worker response") + + + @socketio.on("user_start_log_stream") + def handle_user_stream_start(data): + container_id = data.get("container_id") + user_id = data.get("user_id") + request_id = data.get("request_id") + + try: + logger.info(f"[{request_id}] Starting log stream for container {container_id}") + container_response = requests.get(f"{api_server_url}/workloads/containers/{container_id}") + container_response.raise_for_status() + container_info = container_response.json() + + worker_id = container_info["workload_host_id"] + if worker_id not in connected_workers: + emit("user_log_response", { + "success": False, + "logs": "Worker not connected", + "request_id": request_id + }, to=request.sid) + return + + redis_client = get_redis_client() + context_payload = { + "user_sid": request.sid, + "container_id": container_id, + "user_id": user_id, + "worker_id": worker_id + } + redis_client.setex(f"log_request_context:{request_id}", 600, json.dumps(context_payload)) + + socketio.emit("start_container_log_stream", { + "container_name": container_info["id"], + "request_id": request_id, + "worker_id": worker_id + }, to=connected_workers[worker_id]) + + except Exception as e: + logger.exception(f"[{request_id}] Failed to start log stream") + + + @socketio.on("user_stop_log_stream") + def handle_user_stream_stop(data): + request_id = data.get("request_id") + redis_client = get_redis_client() + context_data = redis_client.get(f"log_request_context:{request_id}") + if context_data: + context = json.loads(context_data) + worker_id = context.get("worker_id") + else: + logger.warning(f"[{request_id}] No Redis context, guessing worker_id from payload") + worker_id = data.get("worker_id") + + if worker_id in connected_workers: + socketio.emit("stop_container_log_stream", { + "request_id": request_id + }, to=connected_workers[worker_id]) + redis_client.delete(f"log_request_context:{request_id}") + + + @socketio.on("container-log-stream") + def handle_streamed_logs(data): + request_id = data.get("request_id") + logs = data.get("logs") + worker_id = data.get("worker_id") + + redis_client = get_redis_client() + context_data = redis_client.get(f"log_request_context:{request_id}") + if not context_data: + logger.warning(f"[{request_id}] No context found for streaming") + return + + context = json.loads(context_data) + user_sid = context.get("user_sid") + with base.connected_sids_lock: + if user_sid in base.connected_sids: + socketio.emit("user_log_stream_update", { + "logs": logs, + "request_id": request_id + }, to=user_sid) + else: + logger.warning(f"[{request_id}] User SID disconnected, stopping stream") + handle_user_stream_stop({"request_id": request_id, "worker_id": worker_id}) diff --git a/websocket_server/events/messages.py b/websocket_server/events/messages.py new file mode 100644 index 0000000..c5de556 --- /dev/null +++ b/websocket_server/events/messages.py @@ -0,0 +1,88 @@ +# websocket_server/events/messages.py + +from flask import request +from flask_socketio import emit +import logging +import json + +from websocket_server.config import get_redis_client +from websocket_server.events import base +from api_client.client import IaaSClient +from websocket_server.shared_state import connected_workers, connected_sids_lock, connected_sids, worker_lock + + +logger = logging.getLogger("websocket_server") + +def register_socketio_handlers(socketio): + @socketio.on("connect") + def handle_connect(): + with connected_sids_lock: + connected_sids.add(request.sid) + logger.info(f"Client connected from {request.remote_addr} SID={request.sid}") + emit("welcome", {"message": f"Connected to server. Your SID is {request.sid}"}) + + + @socketio.on("disconnect") + def handle_disconnect(): + from websocket_server.worker_manager import notify_worker_disconnect + + sid = request.sid + logger.info(f"SID={sid} disconnected") + + with connected_sids_lock: + connected_sids.discard(sid) + + # Check if it's a worker + worker_id = None + with worker_lock: + for w_id, s_id in connected_workers.items(): + if s_id == sid: + worker_id = w_id + break + if worker_id: + del connected_workers[worker_id] + logger.warning(f"Worker {worker_id} disconnected") + notify_worker_disconnect(worker_id) + + + @socketio.on("join_request") + def handle_join_request(data): + from websocket_server.worker_manager import notify_worker_online, start_worker_thread + from websocket_server.task_assigner import assign_task_to_worker + + worker_id = data.get("worker_id") + worker_secret = data.get("worker_secret") + + logger.info(f"Join request received from {worker_id}") + + try: + api = IaaSClient("xxx") + db_worker = api.get_workload_host(worker_id) + + if not db_worker: + logger.error(f"{worker_id} not found in DB") + emit("join_reject", {"message": "Join request rejected"}, to=request.sid) + return + + if db_worker["secret_key"] != worker_secret: + logger.error(f"Invalid secret for worker {worker_id}") + emit("join_reject", {"message": "Join request rejected"}, to=request.sid) + return + + with worker_lock: + connected_workers[worker_id] = request.sid + + redis_client = get_redis_client() + redis_client.set(f"worker_status_{worker_id}", "idle") + + start_worker_thread(worker_id) + notify_worker_online(worker_id) + + logger.info(f"Worker {worker_id} accepted") + emit("join_accept", {"message": f"Worker {worker_id} accepted"}, to=request.sid) + emit("join_broadcast", {"message": f"Worker {worker_id} subscribed to task queue."}) + assign_task_to_worker(worker_id) + + except Exception as e: + logger.exception(f"Error processing join for {worker_id}: {e}") + emit("join_reject", {"message": "Internal error during join"}, to=request.sid) diff --git a/websocket_server/events/terminal.py b/websocket_server/events/terminal.py new file mode 100644 index 0000000..fd21b9b --- /dev/null +++ b/websocket_server/events/terminal.py @@ -0,0 +1,102 @@ +# websocket_server/events/terminal.py + +from flask_socketio import emit +import json +import logging +import requests + +from websocket_server.config import get_redis_client +from websocket_server.events import base +from websocket_server.shared_state import connected_workers, connected_sids_lock, connected_sids + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + +api_server_url = "http://127.0.0.1:5000/api" +def register_socketio_handlers(socketio): + + @socketio.on("start_terminal_session") + def handle_start_terminal_session(data): + container_id = data["container_id"] + request_id = data["request_id"] + user_sid = request.sid + + logger.info(f"[start_terminal_session] Request ID: {request_id}, Container ID: {container_id}, User SID: {user_sid}") + + try: + # Lookup container details + response = requests.get(f"{api_server_url}/workloads/containers/{container_id}") + response.raise_for_status() + info = response.json() + worker_id = info["workload_host_id"] + + context = { + "user_sid": user_sid, + "worker_id": worker_id, + "container_id": container_id + } + + redis_client = get_redis_client() + redis_key = f"terminal_context:{request_id}" + redis_client.setex(redis_key, 600, json.dumps(context)) + logger.info(f"[start_terminal_session] Stored context in Redis under key {redis_key}") + + if worker_id not in connected_workers: + logger.error(f"[start_terminal_session] Worker ID {worker_id} not connected.") + return + + socketio.emit("start_terminal", { + "container_name": container_id, + "request_id": request_id + }, to=connected_workers[worker_id]) + logger.info(f"[start_terminal_session] Emitted 'start_terminal' to worker {worker_id}") + + except Exception as e: + logger.exception(f"[start_terminal_session] Exception occurred: {e}") + + @socketio.on("stop_terminal_session") + def handle_stop_terminal_session(data): + request_id = data["request_id"] + redis_key = f"terminal_context:{request_id}" + + try: + redis_client = get_redis_client() + context_raw = redis_client.get(redis_key) + context = json.loads(context_raw or '{}') + + if not context: + logger.warning(f"[stop_terminal_session] No context found for request ID {request_id}") + return + + worker_id = context["worker_id"] + socketio.emit("stop_terminal", {"request_id": request_id}, to=connected_workers[worker_id]) + logger.info(f"[stop_terminal_session] Emitted 'stop_terminal' to worker {worker_id}") + + redis_client.delete(redis_key) + logger.info(f"[stop_terminal_session] Deleted Redis key {redis_key}") + + except Exception as e: + logger.exception(f"[stop_terminal_session] Exception occurred: {e}") + + + @socketio.on("terminal_input") + def handle_terminal_input(data): + request_id = data["request_id"] + context = json.loads(get_redis_client().get(f"terminal_context:{request_id}") or '{}') + if context: + socketio.emit("terminal_input", { + "request_id": request_id, + "data": data["data"] + }, to=connected_workers[context["worker_id"]]) + + + @socketio.on("terminal_output") + def handle_terminal_output(data): + request_id = data["request_id"] + context = json.loads(get_redis_client().get(f"terminal_context:{request_id}") or '{}') + if context: + socketio.emit("terminal_output", { + "request_id": request_id, + "data": data["data"] + }, to=context["user_sid"]) + diff --git a/websocket_server/events/vnc.py b/websocket_server/events/vnc.py new file mode 100644 index 0000000..73e9ca6 --- /dev/null +++ b/websocket_server/events/vnc.py @@ -0,0 +1,125 @@ +# websocket_server/events/vnc.py + +from flask_socketio import emit +import logging +import requests + +from websocket_server.events import base +from websocket_server.shared_state import connected_workers, connected_sids_lock, connected_sids + + +logger = logging.getLogger("websocket_server") +socketio = base.socketio + +api_server_url = "http://127.0.0.1:5000/api" + +# In-memory request map: vnc_request_id -> {worker_id, vnc_proxy_sid, vm_id} +vnc_request_worker_map = {} +def register_socketio_handlers(socketio): + + + + + @socketio.on("vnc_frame_from_worker") + def vnc_frame_from_worker(payload): + vnc_request_id = payload.get("vnc_request_id") + if not vnc_request_id: + logger.warning("Received VNC frame without request ID") + return + + context = vnc_request_worker_map.get(vnc_request_id) + if not context: + logger.warning(f"No mapping found for VNC request {vnc_request_id}") + return + + proxy_sid = context["vnc_proxy_sid"] + emit("vnc_frame_from_worker", payload, to=proxy_sid) + logger.debug(f"Forwarded VNC frame for {vnc_request_id} to VNC proxy SID {proxy_sid}") + + + @socketio.on("vnc_frame_from_novnc") + def vnc_frame_from_novnc(payload): + logger.debug(f"Got a frame from novnc") + #Send the VNC frame back to the VNC Proxy that matches the requestID + # logger.debug(f"Scored a VNC frame from novnc {data}") + emit("vnc_frame_from_novnc",payload,broadcast=True) + + + @socketio.on("start_vnc_stream") + def start_vnc_stream(data): + logger.info(f"Got VNC start stream request {data}") + vnc_request_id = data.get("vnc_request_id") + virtual_machine_id = data.get("virtual_machine_id") + user_token = data.get("user_token") + vnc_proxy_sid = request.sid # Who sent the request + logger.debug(f"User token: {user_token}") + if not vnc_request_id or not virtual_machine_id or not user_token: + logger.error("Missing required data in VNC stream request") + return + + try: + vm_url = f"{api_server_url}/workloads/virtual_machines/{virtual_machine_id}" + response = requests.get(vm_url, timeout=5) + response.raise_for_status() + vm_info = response.json() + logger.debug(f"Fetched VM info: {vm_info}") + except Exception as e: + logger.error(f"Failed to fetch VM info for {virtual_machine_id}: {e}") + return + + workload_host_id = vm_info.get("workload_host_id") + if not workload_host_id: + logger.error(f"No workload_host_id found for VM {virtual_machine_id}") + return + + socket_id = connected_workers.get(workload_host_id) + if not socket_id: + logger.warning(f"Worker {workload_host_id} not connected") + return + + # Save mapping with proxy SID + vnc_request_worker_map[vnc_request_id] = { + "worker_id": workload_host_id, + "vnc_proxy_sid": vnc_proxy_sid, + "virtual_machine_id": virtual_machine_id + } + + emit("start_vnc_stream_on_worker", { + "vnc_request_id": vnc_request_id, + "virtual_machine_id": virtual_machine_id, + "workload_host_id": workload_host_id, + "vnc_port": 5900, + "vnc_ip_address": "127.0.0.1" + }, to=socket_id) + + logger.info(f"Dispatched VNC stream request {vnc_request_id} to worker {workload_host_id}, from proxy {vnc_proxy_sid}") + + @socketio.on("stop_vnc_stream") + def stop_vnc_stream(data): + logger.info(f"Got VNC stop stream request {data}") + vnc_request_id = data.get("vnc_request_id") + + if not vnc_request_id: + logger.error("Missing vnc_request_id in stop request") + return + + context = vnc_request_worker_map.get(vnc_request_id) + if not context: + logger.warning(f"No mapping found for stop request {vnc_request_id}") + return + + worker_id = context["worker_id"] + socket_id = connected_workers.get(worker_id) + if socket_id: + emit("stop_vnc_stream_on_worker", { + "vnc_request_id": vnc_request_id, + }, to=socket_id) + logger.info(f"Sent stop_vnc_stream to worker {worker_id}") + else: + logger.warning(f"Worker {worker_id} not connected, skip sending stop signal") + + # Clean up mapping + del vnc_request_worker_map[vnc_request_id] + logger.info(f"Cleaned up vnc_request_worker_map entry for {vnc_request_id}") + + diff --git a/websocket_server/main.py b/websocket_server/main.py new file mode 100644 index 0000000..72717b7 --- /dev/null +++ b/websocket_server/main.py @@ -0,0 +1,53 @@ +# websocket_server/main.py + +from flask import Flask +from flask_socketio import SocketIO +import threading +import logging + +from websocket_server.config import setup_logging +from websocket_server.worker_manager import worker_dispatch_flag_check, all_worker_watchdog +from websocket_server.events import base +# Setup logging +logger = setup_logging() +logger.info("Starting WebSocket server...") + +# Flask app and SocketIO setup +app = Flask(__name__) +base.socketio = SocketIO( + app, + cors_allowed_origins="*", + ping_timeout=5, + ping_interval=2 +) + +# Import all socket event handlers +from websocket_server.events import messages, ack, docker, libvirt, logs, terminal, vnc +messages.register_socketio_handlers(base.socketio) +ack.register_socketio_handlers(base.socketio) +docker.register_socketio_handlers(base.socketio) +libvirt.register_socketio_handlers(base.socketio) +logs.register_socketio_handlers(base.socketio) +terminal.register_socketio_handlers(base.socketio) +vnc.register_socketio_handlers(base.socketio) + + + +# Attach routes if needed +from websocket_server import task_api +task_api.register_routes(app) + + +def start_background_threads(): + logger.info("Starting background worker threads...") + + dispatch_thread = threading.Thread(target=worker_dispatch_flag_check, daemon=True) + dispatch_thread.start() + + watchdog_thread = threading.Thread(target=all_worker_watchdog, daemon=True) + watchdog_thread.start() + + +if __name__ == "__main__": + start_background_threads() + base.socketio.run(app, host="0.0.0.0", port=6001) diff --git a/websocket_server/models.py b/websocket_server/models.py new file mode 100644 index 0000000..7a00f9e --- /dev/null +++ b/websocket_server/models.py @@ -0,0 +1,26 @@ +# websocket_server/models.py + +from sqlalchemy import Column, Integer, String, Text, DateTime, Float +from websocket_server.config import Base, engine +from datetime import datetime + +class Task(Base): + __tablename__ = "tasks" + + id = Column(Integer, primary_key=True) + worker_id = Column(String(64), index=True) + job_details = Column(Text) + response = Column(Text) + status = Column(Text) + success = Column(Text) + creation_time = Column(DateTime, default=datetime.utcnow) + start_time = Column(DateTime, nullable=True) + finish_time = Column(DateTime, nullable=True) + wait_time = Column(Float, nullable=True) + execution_time = Column(Float, nullable=True) + task_type = Column(String(50), nullable=False) + depends_on = Column(Integer, nullable=True) # ID of the task this task depends on + not_before = Column(DateTime, nullable=True) # Earliest time this task can run + +# Create tables +Base.metadata.create_all(bind=engine) diff --git a/websocket_server/redis_utils.py b/websocket_server/redis_utils.py new file mode 100644 index 0000000..c6cd6fd --- /dev/null +++ b/websocket_server/redis_utils.py @@ -0,0 +1,57 @@ +# websocket_server/redis_utils.py + +import threading +import traceback +import time +from websocket_server.config import get_redis_client +from websocket_server.task_assigner import assign_task_to_worker +from websocket_server.worker_manager import notify_worker_disconnect +import logging + +logger = logging.getLogger("websocket_server") + +# Shared state +worker_dispatch_flags = {} +thread_stop_flags = {} + +def redis_subscribe(worker_id): + """ + Redis Pub/Sub monitor for a specific worker. + Subscribes to worker status and queue keys and sets dispatch flags accordingly. + """ + redis_client = get_redis_client() + redis_client.config_set("notify-keyspace-events", "KA") + + pubsub = redis_client.pubsub() + channels = [ + f"__keyspace@0__:worker_status_{worker_id}", + f"__keyspace@0__:worker_queue_{worker_id}" + ] + + try: + pubsub.psubscribe(channels) + logger.info(f"Subscribed to Redis keyspace events: {channels}") + + while not thread_stop_flags.get(worker_id, False): + message = pubsub.get_message(timeout=1.0) + if message and message["type"] in ["pmessage", "message"]: + current_status = redis_client.get(f"worker_status_{worker_id}") + logger.debug(f"[{worker_id}] Current Redis status: {current_status}") + + if current_status and current_status.lower() == "idle": + if redis_client.get(f"worker_queue_{worker_id}").lower() == "true": + logger.info(f"[{worker_id}] Queue is true, setting dispatch flag") + worker_dispatch_flags[worker_id] = True + else: + logger.info(f"[{worker_id}] Queue false, no dispatch needed") + else: + logger.info(f"[{worker_id}] Status is not idle, skipping") + + logger.warning(f"[{worker_id}] Redis pub/sub loop exited, disconnecting") + notify_worker_disconnect(worker_id) + + except Exception as e: + error_details = traceback.format_exc() + logger.error(f"[{worker_id}] Redis pub/sub error: {e}\n{error_details}") + finally: + pubsub.close() diff --git a/websocket_server/shared_state.py b/websocket_server/shared_state.py new file mode 100644 index 0000000..5c7a445 --- /dev/null +++ b/websocket_server/shared_state.py @@ -0,0 +1,8 @@ +# websocket_server/shared_state.py + +import threading + +connected_workers = {} +worker_lock = threading.Lock() +connected_sids = set() +connected_sids_lock = threading.Lock() diff --git a/websocket_server/task_api.py b/websocket_server/task_api.py new file mode 100644 index 0000000..b72ff34 --- /dev/null +++ b/websocket_server/task_api.py @@ -0,0 +1,95 @@ +# websocket_server/task_api.py + +from flask import request, jsonify +from datetime import datetime +import json +import logging + +from websocket_server.models import Task +from websocket_server.config import get_db_session, get_redis_client + +logger = logging.getLogger("websocket_server") +redis_client = get_redis_client() + +# Import dispatch logic +from websocket_server.task_assigner import assign_task_to_worker +from websocket_server.shared_state import connected_workers, worker_lock + + +def register_routes(app): + """ + Register all Flask routes to the given app. + """ + + @app.route("/api/create_task", methods=["POST"]) + def create_task(): + """ + Create a new task and insert it into the database. + """ + data = request.get_json() + worker_id = data.get("worker_id") + job_details = json.dumps(data.get("job_details")) + task_type = data.get("task_type", "default") + depends_on = data.get("depends_on") + not_before = data.get("not_before") + + logger.info(f"Creating new task for worker {worker_id}") + + if not worker_id or not job_details: + return jsonify({"error": "worker_id and job_details are required"}), 400 + + try: + not_before_dt = None + if not_before: + not_before_dt = datetime.fromisoformat(not_before) + + with get_db_session() as session: + task = Task( + worker_id=worker_id, + job_details=job_details, + task_type=task_type, + status="pending", + success=None, + creation_time=datetime.utcnow(), + depends_on=depends_on, + not_before=not_before_dt + ) + session.add(task) + session.commit() + + logger.info(f"Created task {task.id} for worker {worker_id}") + redis_client.set(f"worker_queue_{worker_id}", "True") + + return jsonify({"message": "Task created", "task_id": task.id}), 201 + + except Exception as e: + logger.error(f"Error creating task: {e}") + return jsonify({"error": "Internal error creating task"}), 500 + + @app.route("/api/assign_task", methods=["POST"]) + def manually_assign_task(): + """ + Force task assignment for a specific worker via HTTP trigger. + """ + data = request.get_json() + worker_id = data.get("worker_id") + + if not worker_id: + return jsonify({"error": "Missing worker_id"}), 400 + + logger.info(f"Manual trigger for assign_task to worker {worker_id}") + assign_task_to_worker(worker_id) + return jsonify({"message": f"Dispatched task assignment to worker {worker_id}."}) + + @app.route("/api/connected_clients", methods=["GET"]) + def get_connected_clients(): + """ + Returns currently connected worker IDs and socket session IDs. + """ + logger.debug("Listing connected clients") + with worker_lock: + result = [ + {"worker_id": wid, "socket_id": sid, "ip": "unknown"} + for wid, sid in connected_workers.items() + ] + return jsonify({"connected_clients": result}) diff --git a/websocket_server/task_assigner.py b/websocket_server/task_assigner.py new file mode 100644 index 0000000..437b782 --- /dev/null +++ b/websocket_server/task_assigner.py @@ -0,0 +1,107 @@ +# websocket_server/task_assigner.py + +import time +from datetime import datetime +from sqlalchemy.exc import OperationalError +from sqlalchemy import or_ +from websocket_server.config import get_db_session, get_redis_client +from websocket_server.models import Task +import logging +from websocket_server.shared_state import connected_workers, worker_lock +from websocket_server.events import base + +logger = logging.getLogger("websocket_server") + +def assign_task_to_worker(worker_id): + """ + Attempt to assign a pending task to a connected worker. + Ensures locking, dependency satisfaction, and Redis status tracking. + """ + redis_client = get_redis_client() + max_lock_retries = 5 + lock_retry_wait = 0.2 + + logger.info(f"[{worker_id}] Starting task assignment process") + + try: + redis_client.set(f"worker_status_{worker_id}", "busy") + logger.debug(f"[{worker_id}] Marked as 'busy' in Redis") + + for attempt in range(max_lock_retries): + logger.debug(f"[{worker_id}] Lock attempt {attempt + 1}/{max_lock_retries}") + with get_db_session() as session: + now = datetime.utcnow() + candidates = session.query(Task).filter( + Task.worker_id == worker_id, + Task.status == 'pending', + or_(Task.not_before == None, Task.not_before <= now) + ).all() + + ready = [] + for task in candidates: + if not task.depends_on: + ready.append(task) + else: + dep = session.query(Task).filter( + Task.id == task.depends_on, + Task.status == 'acknowledged' + ).first() + if dep: + ready.append(task) + + if not ready: + logger.info(f"[{worker_id}] No ready tasks with satisfied dependencies") + redis_client.set(f"worker_queue_{worker_id}", "False") + redis_client.set(f"worker_status_{worker_id}", "idle") + return + + # New isolated session for locking and updating + with get_db_session() as session: + selected = session.query(Task).filter( + Task.id.in_([t.id for t in ready]), + Task.status == 'pending' + ).with_for_update(skip_locked=True).limit(1).one_or_none() + + if selected: + logger.info(f"[{worker_id}] Task selected: {selected.id}") + break + + logger.warning(f"[{worker_id}] Lock contention, retrying...") + time.sleep(lock_retry_wait) + lock_retry_wait *= 2 + else: + logger.error(f"[{worker_id}] Failed to acquire task after {max_lock_retries} attempts") + redis_client.set(f"worker_status_{worker_id}", "idle") + redis_client.set(f"worker_queue_{worker_id}", "False") + return + + with worker_lock: + if worker_id in connected_workers: + logger.debug(f"[{worker_id}] Dispatching task {selected.id}") + base.socketio.emit("task", { + "task_id": selected.id, + "worker_id": worker_id, + "job_details": selected.job_details, + "type": selected.task_type + }, to=connected_workers[worker_id]) + + with get_db_session() as session: + task = session.query(Task).filter(Task.id == selected.id).first() + task.start_time = datetime.utcnow() + task.wait_time = (task.start_time - task.creation_time).total_seconds() + task.status = "in-progress" + session.add(task) + session.commit() + + remaining = session.query(Task).filter( + Task.worker_id == worker_id, + Task.status == 'pending' + ).count() + + redis_client.set(f"worker_queue_{worker_id}", "True" if remaining > 0 else "False") + logger.info(f"[{worker_id}] Task {selected.id} marked in-progress") + else: + logger.warning(f"[{worker_id}] Worker disconnected before dispatch") + + except Exception as e: + logger.exception(f"[{worker_id}] Error during task assignment: {e}") diff --git a/websocket_server/worker_manager.py b/websocket_server/worker_manager.py new file mode 100644 index 0000000..400f919 --- /dev/null +++ b/websocket_server/worker_manager.py @@ -0,0 +1,90 @@ +# websocket_server/worker_manager.py + +import threading +import requests +import logging +# from websocket_server.redis_utils import redis_subscribe, thread_stop_flags, worker_dispatch_flags +from websocket_server.task_assigner import assign_task_to_worker +from websocket_server.config import get_redis_client +from websocket_server.shared_state import connected_workers, worker_lock + +logger = logging.getLogger("websocket_server") + +# API server endpoint +api_server_url = "http://127.0.0.1:5000/api" + +# Tracks threads for each worker +worker_threads = {} + +def notify_worker_online(worker_id): + """ + Inform the API server that a worker has connected. + """ + logger.debug(f"[{worker_id}] Notifying API server: online") + payload = {"status": "online"} + headers = {"Content-Type": "application/json"} + + try: + response = requests.put(f"{api_server_url}/workload_hosts/{worker_id}", json=payload, headers=headers) + logger.debug(response.text) + logger.info(f"[{worker_id}] API server acknowledged online state") + except Exception as e: + logger.error(f"[{worker_id}] Failed to notify API server of online status: {e}") + + +def notify_worker_disconnect(worker_id): + """ + Inform the API server that a worker has disconnected. + """ + logger.debug(f"[{worker_id}] Notifying API server: offline") + payload = {"status": "offline"} + headers = {"Content-Type": "application/json"} + + try: + response = requests.put(f"{api_server_url}/workload_hosts/{worker_id}", json=payload, headers=headers) + logger.debug(response.text) + logger.info(f"[{worker_id}] API server acknowledged offline state") + except Exception as e: + logger.error(f"[{worker_id}] Failed to notify API server of disconnect: {e}") + + +def worker_dispatch_flag_check(): + from websocket_server.redis_utils import redis_subscribe, thread_stop_flags, worker_dispatch_flags + + """ + Continuously checks for workers with dispatch flags and triggers task assignment. + """ + logger.info("Worker dispatch flag thread started") + while True: + for wid in list(worker_dispatch_flags.keys()): + logger.info(f"[{wid}] Dispatch flag set. Triggering task assign.") + del worker_dispatch_flags[wid] + assign_task_to_worker(wid) + threading.Event().wait(0.01) + + +def all_worker_watchdog(): + """ + Periodically iterate over connected workers and ensure they are polled for task assignments. + Prevents idle starvation or missed Redis events. + """ + logger.info("Watchdog thread started") + while True: + for wid in list(connected_workers.keys()): + logger.info(f"[{wid}] Watchdog triggering assign check") + assign_task_to_worker(wid) + threading.Event().wait(30) + + +def start_worker_thread(worker_id): + from websocket_server.redis_utils import redis_subscribe, thread_stop_flags, worker_dispatch_flags + + """ + Launch Redis subscription thread for the given worker. + """ + logger.info(f"[{worker_id}] Starting Redis listener thread") + + thread_stop_flags[worker_id] = False + sub_thread = threading.Thread(target=redis_subscribe, args=(worker_id,), daemon=True) + worker_threads[worker_id] = sub_thread + sub_thread.start()