Files
3cloud-backend/websocket_server.py
T
2025-03-03 10:59:04 +10:30

590 lines
22 KiB
Python
Executable File

from datetime import datetime
import json
import os
import threading
import time
import traceback
from flask import Flask, request, jsonify, render_template
from flask_socketio import SocketIO, emit
from sqlalchemy import create_engine, Column, String, Text, DateTime, Float, Integer
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker
import logging
import redis
from colorlog import ColoredFormatter
from sqlalchemy.exc import OperationalError
from contextlib import contextmanager
from datetime import datetime
import pymysql
api_server_url="http://127.0.0.1:5000/api/"
# Custom filter to include the function name
class FunctionNameFilter(logging.Filter):
def filter(self, record):
record.funcName = record.funcName if hasattr(record, 'funcName') else '<unknown>'
return True
# Define the colorized log format
log_format = (
"%(log_color)s%(asctime)s - %(levelname)s - %(funcName)s - %(message)s"
)
date_format = "%Y-%m-%d %H:%M:%S"
# Configure the formatter with colors
formatter = ColoredFormatter(
log_format,
datefmt=date_format,
log_colors={
"DEBUG": "cyan",
"INFO": "green",
"WARNING": "yellow",
"ERROR": "red",
"CRITICAL": "bold_red",
},
)
# Configure the handler
handler = logging.StreamHandler()
handler.setFormatter(formatter)
# Configure the logger
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
logger.addFilter(FunctionNameFilter())
logger.addHandler(handler)
# Ensure the logs directory exists
log_dir = "logs"
if not os.path.exists(log_dir):
os.makedirs(log_dir)
# Configure the file handler
file_handler = logging.FileHandler("logs/app.log")
file_handler.setFormatter(formatter)
file_handler.setLevel(logging.WARNING)
logger.addHandler(file_handler)
# Flask and SocketIO setup
app = Flask(__name__)
socketio = SocketIO(
app,
cors_allowed_origins="*",
ping_timeout=5,
ping_interval=2
)
pymysql.install_as_MySQLdb()
# Database setup with improved connection handling
DATABASE_URL = "mysql://root:password@172.17.0.1:3306/defaultdb"
# DATABASE_URL = "cockroachdb://root@172.17.0.1:26257/defaultdb"
engine = create_engine(
DATABASE_URL,
isolation_level='SERIALIZABLE',
pool_size=5,
max_overflow=10,
pool_timeout=30,
pool_recycle=3600, # Recycle connections after an hour
pool_pre_ping=True # Enable connection health checks
)
SessionLocal = sessionmaker(
autocommit=False,
autoflush=False,
bind=engine,
expire_on_commit=False
)
Base = declarative_base()
# Database Models
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)
Base.metadata.create_all(bind=engine)
# Redis connection pool
redis_connection_pool = redis.ConnectionPool(
host="localhost",
port=6379,
decode_responses=True,
max_connections=10, # Limit maximum connections
socket_timeout=5 # Add timeout for operations
)
worker_threads = {}
thread_stop_flags = {}
connected_workers = {}
worker_dispatch_flags={}
worker_lock = threading.Lock()
# Context manager for database sessions
@contextmanager
def get_db_session():
session = SessionLocal()
try:
yield session
session.commit()
except Exception:
session.rollback()
raise
finally:
session.close()
# Function to get Redis client
def get_redis_client():
return redis.StrictRedis(connection_pool=redis_connection_pool)
def all_worker_watchdog():
"""Force checks outstanding jobs for all connected clients every x seconds
This is designed to avoid any job's being 'missed'
"""
while 1:
logger.info(f"Watchdog check loop running")
for wid, value in list(connected_workers.items()):
logger.info(f"Heartbeat check for {wid}")
# time.sleep(2)
assign_task_to_worker(wid)
# logger.info(f"Heartbeat check loop complete, sleeping")
time.sleep(30)
def worker_dispatch_flag_check():
logger.info(f"WorkerID dispatch loop starting")
while 1:
#We use this 'list' function to make a copy of the dict so that
# if we delte a vlaue we arent deleteing the vlaue form the list we are iterating over which causes an error
for wid, value in list(worker_dispatch_flags.items()):
logger.info(f"Found dispatch flag for {wid}")
del worker_dispatch_flags[wid]
assign_task_to_worker(wid)
time.sleep(.01)
def assign_task_to_worker(worker_id):
"""Assign a new task to a worker with improved lock handling"""
redis_client = get_redis_client()
max_lock_retries = 5
lock_retry_wait = 0.2
logger.info(f"WorkerID:{worker_id} Starting task assignment process.")
try:
redis_client.set(f"worker_status_{worker_id}", "busy")
logger.info(f"WorkerID:{worker_id} Updated status to 'busy' in Redis.")
for lock_attempt in range(max_lock_retries):
logger.info(f"WorkerID:{worker_id} Attempting to query tasks for worker. Attempt {lock_attempt + 1}/{max_lock_retries}.")
with get_db_session() as session:
task_count = session.query(Task).filter(
Task.worker_id == worker_id,
Task.status == 'pending'
).count()
logger.debug(f"WorkerID:{worker_id} Pending task count: {task_count}.")
if task_count == 0:
logger.warning(f"WorkerID:{worker_id} No pending tasks found. Setting status to 'idle'.")
redis_client.set(f"worker_queue_{worker_id}", "False")
redis_client.set(f"worker_status_{worker_id}", "idle")
session.close()
return
try:
logger.info(f"WorkerID:{worker_id} Attempting to acquire lock for task retrieval.")
task = session.query(Task).filter(
Task.worker_id == worker_id,
Task.status == 'pending'
).with_for_update(skip_locked=True).limit(1).one_or_none()
if task:
logger.info(f"WorkerID:{worker_id} Successfully retrieved task with ID: {task.id}.")
break
except OperationalError as e:
logger.warning(
f"WorkerID:{worker_id} Lock acquisition failed on attempt {lock_attempt + 1}/{max_lock_retries}. Retrying after {lock_retry_wait:.2f}s. Error: {e}"
)
time.sleep(lock_retry_wait)
lock_retry_wait *= 2
continue
if not task:
logger.error(f"WorkerID:{worker_id} Unable to acquire task lock after {max_lock_retries} attempts. Setting status to 'idle'.")
redis_client.set(f"worker_queue_{worker_id}", "False")
redis_client.set(f"worker_status_{worker_id}", "idle")
return
logger.info(f"WorkerID:{worker_id} Task details: ID={task.id}, JobDetails={task.job_details}.")
with worker_lock:
logger.debug(f"WorkerID:{worker_id} Acquired worker lock.")
if worker_id in connected_workers:
logger.info(f"WorkerID:{worker_id} Worker is connected. Dispatching task ID {task.id}.")
socketio.emit(
"task",
{
"task_id": task.id,
"worker_id": worker_id,
"job_details": task.job_details,
"type": task.task_type
},
to=connected_workers[worker_id],
)
current_time = datetime.utcnow()
task.start_time = current_time
task.wait_time = (current_time - task.creation_time).total_seconds()
task.status = "in-progress"
session.add(task)
session.commit()
session.close()
logger.info(f"WorkerID:{worker_id} Updated task ID {task.id} to 'in-progress' with start time {current_time}.")
if task_count <= 1:
redis_client.set(f"worker_queue_{worker_id}", "false")
logger.info(f"WorkerID:{worker_id} Set worker queue to 'false' as task count is {task_count}.")
return
else:
logger.warning(f"WorkerID:{worker_id} Worker is not connected. Task ID {task.id} not dispatched.")
session.close()
return None
except Exception as e:
logger.error(f"WorkerID:{worker_id} Error during task assignment: {e}")
raise
def redis_subscribe(worker_id):
"""Redis Pub/Sub monitor with improved connection handling
If a change to worker_status_ or worker_queue_ is detected, raise a flag in a global variable which will trigger the worker manager thread to take action
"""
redis_client = get_redis_client()
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 channels: {channels}")
while not thread_stop_flags.get(worker_id, False):
message = pubsub.get_message(timeout=1.0)
if message:
if message["type"] in ["pmessage", "message"]:
current_status = redis_client.get(f"worker_status_{worker_id}")
logger.debug(f"Current status of {worker_id} = {current_status}")
if current_status.lower() == "idle":
logger.info(f"Worker {worker_id} is idle. Checking queue")
if redis_client.get(f"worker_queue_{worker_id}").lower() == "true":
logger.info(f"Worker {worker_id} queue is true in redis, setting flag.")
worker_dispatch_flags[worker_id]=True
else:
logger.info(f"Worker {worker_id} queue flag not set, skipping")
else:
logger.info(f"Worker {worker_id} is not idle, skipping.")
else:
logger.debug(f"Message ignored {message}")
logger.error("Loop exited, should never happen")
except Exception as e:
error_details = traceback.format_exc()
logger.error(f"Error in Pub/Sub for worker {worker_id}: {e}\n{error_details}")
finally:
pubsub.close()
@app.route("/")
def index():
return render_template("index.html")
@socketio.on("connect")
def handle_connect():
client_ip = request.remote_addr
logger.info(f"New client connected from IP: {client_ip}")
emit("welcome", {"message": "Connected to the API server!"}, broadcast=True)
@socketio.on("disconnect")
def handle_disconnect():
worker_id = None
try:
with worker_lock:
for w_id, sid in connected_workers.items():
if sid == request.sid:
worker_id = w_id
break
if worker_id:
del connected_workers[worker_id]
logger.info(f"Worker {worker_id} disconnected. Session ID: {request.sid}")
if worker_id in worker_threads:
logger.info("Worker found in worker threads")
thread_stop_flags[worker_id] = True
worker_threads[worker_id].join(timeout=5.0) # Add timeout to prevent hanging
del worker_threads[worker_id]
del thread_stop_flags[worker_id]
logger.info(f"Stopped thread for worker {worker_id}.")
else:
logger.warning(f"Disconnected client not found in connected workers.")
except Exception as e:
logger.error(f"Error during disconnect handling: {e}")
@socketio.on("join_request")
def handle_join_request(data):
worker_id = data["worker_id"]
worker_secret = data["worker_secret"]
logger.info(f"Join request recieved from {worker_id}")
try:
with worker_lock:
connected_workers[worker_id] = request.sid
# TODO - Check worker ID vs worker secret in the DB.. If they dont match kick it to the curb
# We will make a call to the API server to cross check workerId, secret and region ID. They must all match
# if not worker_secret==workerDBITEM.secret:
# socketio.emit("join_reject", {"message": f"Worker {worker_id} join request rejected."},to=request.sid)
# return
from app.api_client.client import IaaSClient
api_client=IaaSClient("xxx")
db_worker=api_client.get_workload_host(worker_id)
if not db_worker:
logger.error(f"{worker_id} not found in DB")
socketio.emit("join_reject", {"message": f"Worker {worker_id} join request rejected."},to=request.sid)
# raise Exception("Worker supplied invalid worker ID")
return
logger.debug(f"Worker claims to be {worker_id}, should have secret {db_worker['secret_key']} supplied secret {worker_secret}")
if not worker_secret==db_worker['secret_key']:
logger.error(f"{worker_id} Secret does not match")
# Dont specify any detailin the join reject, for security reasons if the ID is wrong or
# if the password is wrong we return a generic "join rejected message"
socketio.emit("join_reject", {"message": f"Worker {worker_id} join request rejected."},to=request.sid)
# raise Exception("Worker supplied invalid worker ID")
return
logger.info(f"Worker {worker_id} joined. Session ID: {request.sid}")
redis_client = get_redis_client()
redis_client.set(f"worker_status_{worker_id}", "idle")
thread_stop_flags[worker_id] = False
subscribe_thread = threading.Thread(target=redis_subscribe, args=(worker_id,), daemon=True)
worker_threads[worker_id] = subscribe_thread
subscribe_thread.start()
logger.info(f"Thread started for worker {worker_id}.")
socketio.emit("join_broadcast", {"message": f"Worker {worker_id} subscribed to task queue."})
socketio.emit("join_accept", {"message": f"Worker {worker_id} join request accepted."},to=request.sid)
assign_task_to_worker(worker_id)
except Exception as e:
logger.error(f"Error during worker join: {e}")
raise
@socketio.on("ack")
def handle_ack(data):
logger.debug(f"ack data {data}")
worker_id = data["worker_id"]
task_id = data["task_id"]
result=data["result"]
try:
with get_db_session() as session:
task = session.query(Task).filter(Task.id == task_id).first()
if task:
task.status = "acknowledged"
task.success = 1 if result["success"]==True else 0
task.response=str(result["response"])
task.finish_time = datetime.utcnow()
logger.debug(f"({task.finish_time} - {task.start_time}")
task.execution_time = (task.finish_time - task.start_time).total_seconds()
logger.info(
f"Task {task_id} finished: finish_time={task.finish_time}, "
f"execution_time={task.execution_time}."
)
redis_client = get_redis_client()
redis_client.set(f"worker_status_{worker_id}", "idle")
logger.info(f"Worker {worker_id} status updated to 'idle'.")
# Dont update container Status, handle this in docker_event or vm_event or somethingelse_event
# if task.task_type=="container-create":
# # Lets ack the container thats probably just been created
# if result["success"] == True:
# new_status="running"
# # TODO - This is wrong, dont update the status to runnign, let the container monitor handle this. If the container transitions to running we'll be notified
# # only update the taskj status in the DB
# else:
# new_status="create-failed"
# # TODO - Take action here to retry
# payload = {
# "new_status": new_status
# }
# logger.debug(f"Sending task update payload to API server {payload}")
# headers = {"Content-Type": "application/json"}
# websocket_server_response = requests.put(f"{api_server_url}/workloads/containers/status_update/{container_id}", data=json.dumps(payload), headers=headers)
# logger.debug( websocket_server_response)
# logger.info("Update complete")
except Exception as e:
logger.error(f"Error processing acknowledgment for worker {worker_id}, task {task_id}: {e}")
raise
@socketio.on("send_message")
def handle_send_message(data):
worker_id = data["worker_id"]
message = data["message"]
try:
with worker_lock:
if worker_id in connected_workers:
socket_id = connected_workers[worker_id]
emit("message", {"message": message}, to=socket_id)
logger.info(f"Message sent to worker {worker_id}: {message}")
else:
logger.warning(f"Worker {worker_id} not found!")
except Exception as e:
logger.error(f"Error sending message: {e}")
@socketio.on("docker_event")
def handle_docker_event(data):
# TODO - Somehow we need to validate the authenticity of the worker ID, can we crosscheck the workerID suppied in the payload vs the socket id?
event_type = data["event_type"],
container_id = data["container_id"],
container_name = data["container_name"],
worker_id = data["worker_id"],
timestamp = data[""]
attributes = data["attributes"],
# TODO - Figure out what has happened and make an API request back tot he API server updating the releveant resource
# Pay attention to container deleted events
# Make sure the API server cross checks the worker ID this event is coming from and the resource ID it's updating
@app.route("/api/assign_task", methods=["POST"])
def manually_assign_task():
"""Manually trigger task assignment to a worker."""
worker_id = request.json.get("worker_id")
redis_client = redis.StrictRedis(host="localhost", port=6379, decode_responses=True)
assign_task_to_worker(worker_id, redis_client)
return jsonify({"message": f"Task assignment process initiated for worker {worker_id}."})
@app.route("/api/connected_clients", methods=["GET"])
def get_connected_clients():
"""API to get a list of connected clients."""
logger.debug(connected_workers.items())
with worker_lock:
connected_clients = [
{
"worker_id": worker_id,
"socket_id": socket_id,
"ip": "none",
}
for worker_id, socket_id in connected_workers.items()
]
return jsonify({"connected_clients": connected_clients})
@app.route("/api/create_task", methods=["POST"])
def create_task():
"""API to create a new task and save it to 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")
logger.info(f"Creating task {data}")
if not worker_id or not job_details:
return jsonify({"error": "worker_id and job_details are required"}), 400
try:
with get_db_session() as session:
new_task = Task(
worker_id=worker_id,
job_details=job_details,
status="pending",
success=None,
task_type=task_type,
creation_time=datetime.utcnow()
)
session.add(new_task)
session.commit()
logger.info(f"New task created: ID={new_task.id}, Worker={worker_id}, Type={task_type}")
redis_client = get_redis_client()
redis_client.set(f"worker_queue_{worker_id}", "True") # Flag the worker for a new task
return jsonify({"message": "Task created successfully", "task_id": new_task.id}), 201
except Exception as e:
logger.error(f"Error creating task: {e}")
return jsonify({"error": "Failed to create task"}), 500
# # # Example request structure
# payload_text = {
# "worker_id": "234",
# "task_type": "container",
# "job_details": {
# "tenancyID": "tenant123a",
# "containers": [
# {
# "container_id": "af47632e-43b7-4474-82fa-475a24fe88d5",
# "docker_image": "nginx",
# "cpu_shares": 1,
# "mem_limit": 128,
# "container_name": "web-server1",
# "command": ["nginx", "-g", "daemon off;"],
# "working_dir": "/usr/share/nginx/html"
# }
# ]
# }
# }
if __name__ == "__main__":
logger.info("Starting server...")
#Start the worker dispatch thread
dispatch_thread = threading.Thread(target=worker_dispatch_flag_check, args=(), daemon=True)
dispatch_thread.start()
# TODO Move the dispatch thread concept into an 'event'
# https://gpttutorpro.com/how-to-create-and-handle-events-in-python/
#Start the worker dispatch thread
all_worker_watchdog_thread = threading.Thread(target=all_worker_watchdog, args=(), daemon=True)
all_worker_watchdog_thread.start()
socketio.run(app, host="0.0.0.0", port=6000)