Disassembled websocket_server

This commit is contained in:
2025-06-17 13:20:44 +09:30
parent dc15108ea5
commit c92c76263a
18 changed files with 1335 additions and 0 deletions
+9
View File
@@ -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"
+107
View File
@@ -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', '<unknown>')
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()
View File
+48
View File
@@ -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}")
+6
View File
@@ -0,0 +1,6 @@
# websocket_server/events/base.py
# import threading
# # SocketIO will be initialized in main.py and shared here
socketio = None
+119
View File
@@ -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}")
+111
View File
@@ -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}")
+184
View File
@@ -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})
+88
View File
@@ -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)
+102
View File
@@ -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"])
+125
View File
@@ -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}")
+53
View File
@@ -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)
+26
View File
@@ -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)
+57
View File
@@ -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()
+8
View File
@@ -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()
+95
View File
@@ -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})
+107
View File
@@ -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}")
+90
View File
@@ -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()