590 lines
22 KiB
Python
Executable File
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)
|
|
|