The only thing this doesnt handle is when a client goes into offline due to a transient link failure but it comes back online. We want to avoid updating the DB every time a pong is RX'd.. Perhaps we do a seocndary corss checks with offline hosts?
134 lines
4.6 KiB
Python
134 lines
4.6 KiB
Python
# websocket_server/worker_manager.py
|
|
|
|
import requests
|
|
import uuid
|
|
import time
|
|
import threading
|
|
import logging
|
|
from websocket_server.shared_state import connected_workers, worker_lock, ping_tracker
|
|
from websocket_server.config import get_redis_client,PING_INTERVAL_SECONDS, ASSIGN_INTERVAL_SECONDS, LIVENESS_CHECK_INTERVAL_SECONDS, PING_EXPIRY_SECONDS
|
|
from websocket_server.task_assigner import assign_task_to_worker
|
|
from websocket_server.events import base
|
|
|
|
|
|
logger = logging.getLogger("websocket_server")
|
|
|
|
# API server endpoint
|
|
api_server_url = "http://172.17.0.1:5000/api"
|
|
|
|
# Tracks threads for each worker
|
|
worker_threads = {}
|
|
|
|
redis = get_redis_client()
|
|
|
|
# Last timestamps
|
|
last_ping_time = 0
|
|
last_assign_time = 0
|
|
|
|
|
|
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():
|
|
"""
|
|
Watchdog loop that:
|
|
- Sends a liveness ping every `PING_INTERVAL_SECONDS`
|
|
- Triggers task reassignment every `ASSIGN_INTERVAL_SECONDS`
|
|
"""
|
|
global last_ping_time, last_assign_time
|
|
logger.info("Watchdog thread started")
|
|
|
|
redis = get_redis_client()
|
|
last_liveness_check = 0
|
|
|
|
while True:
|
|
now = time.time()
|
|
worker_ids = list(connected_workers.keys())
|
|
|
|
# 1. Send custom ping to all workers
|
|
if now - last_ping_time >= PING_INTERVAL_SECONDS:
|
|
if worker_ids:
|
|
logger.debug("Sending liveness pings to all workers")
|
|
for wid in worker_ids:
|
|
ping_id = str(uuid.uuid4())
|
|
ping_tracker[wid] = {"ping_id": ping_id, "sent_time": now}
|
|
sid = connected_workers.get(wid)
|
|
if sid:
|
|
try:
|
|
base.socketio.emit("worker_ping", {"ping_id": ping_id}, to=sid)
|
|
logger.info(f"[{wid}] Sent ping ID: {ping_id} to {wid}")
|
|
except Exception as e:
|
|
logger.warning(f"[{wid}] Ping send failed: {e}")
|
|
last_ping_time = now
|
|
|
|
# 2. Trigger task assignment if due
|
|
if now - last_assign_time >= ASSIGN_INTERVAL_SECONDS:
|
|
if worker_ids:
|
|
logger.debug("Checking for task assignment across all workers")
|
|
for wid in worker_ids:
|
|
logger.info(f"[{wid}] Triggering assign check")
|
|
assign_task_to_worker(wid)
|
|
last_assign_time = now
|
|
|
|
|
|
|
|
time.sleep(1)
|
|
|
|
|
|
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()
|