Disassembled websocket_server
This commit is contained in:
@@ -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"
|
||||||
@@ -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()
|
||||||
@@ -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}")
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
# websocket_server/events/base.py
|
||||||
|
|
||||||
|
# import threading
|
||||||
|
|
||||||
|
# # SocketIO will be initialized in main.py and shared here
|
||||||
|
socketio = None
|
||||||
@@ -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}")
|
||||||
|
|
||||||
@@ -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}")
|
||||||
@@ -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})
|
||||||
@@ -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)
|
||||||
@@ -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"])
|
||||||
|
|
||||||
@@ -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}")
|
||||||
|
|
||||||
|
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -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})
|
||||||
@@ -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}")
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user