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