diff --git a/app/__init__.py b/app/__init__.py index 1804616..f88132d 100644 --- a/app/__init__.py +++ b/app/__init__.py @@ -34,8 +34,8 @@ app.config.update( WEBSOCKET_SERVER_URL="http://172.17.0.1:6001/api/create_task", # Celery settings (override via ENV) - CELERY_BROKER_URL=os.getenv("CELERY_BROKER_URL", "redis://172.17.0.1:6379/0"), - CELERY_RESULT_BACKEND=os.getenv("CELERY_RESULT_BACKEND", "redis://172.17.0.1:6379/1"), + CELERY_BROKER_URL=os.getenv("CELERY_BROKER_URL", "redis://172.17.0.1:6379/0"), # Redis Database 0 for celery broker tasks + CELERY_RESULT_BACKEND=os.getenv("CELERY_RESULT_BACKEND", "redis://172.17.0.1:6379/1"), # Redis Database 1 for celery results ) # --------------------------------------------------------------------------- # @@ -59,6 +59,8 @@ def _mask(value: str, visible: int = 4) -> str: app.config["CLOUDFLARE_API_TOKEN"] = os.getenv("CLOUDFLARE_API_TOKEN") app.config["CLOUDFLARE_ACCOUNT_ID"] = os.getenv("CLOUDFLARE_ACCOUNT_ID") app.config["CLOUDFLARE_ZONE_ID"] = os.getenv("CLOUDFLARE_ZONE_ID") +app.config["PING_HEARTBEAT_TIMEOUT_SECONDS"] = os.getenv("PING_HEARTBEAT_TIMEOUT_SECONDS",5) # How long before a Websocket PING\PONG is classed as a failure and triggers a worker offline event +app.config["REDIS_URL"] = os.getenv("REDIS_URL","redis://172.17.0.1:6379/2") # Redis Database 2 for operational tasks like websocket server and PING logging logger.info( "Cloudflare configuration set " diff --git a/app/celery_app.py b/app/celery_app.py index 51524aa..f9d4284 100644 --- a/app/celery_app.py +++ b/app/celery_app.py @@ -1,5 +1,5 @@ """ -celery_factory.py –– Re-export the canonical Celery instance. +celery_factory.py -- Re-export the canonical Celery instance. This file used to create a *second* Celery app. Since `app/__init__.py` now owns the one-and-only `celery_app`, we simply import and expose it @@ -11,7 +11,6 @@ after import. """ from logging import getLogger -from datetime import timedelta from app import celery_app as celery # ← single source-of-truth from app import app as flask_app # access to config & logger @@ -26,6 +25,10 @@ celery.conf.beat_schedule.update( "task": "tasks.log_heartbeat", "schedule": 30.0, }, + "reconcile_online_workers-30s": { + "task": "tasks.reconcile_online_workers", + "schedule": 2.0, + }, # Add more periodic jobs here if desired # "collect-system-metrics": { # "task": "tasks.collect_system_metrics", @@ -38,3 +41,4 @@ celery.conf.beat_schedule.update( __all__ = ["celery"] import app.tasks.process_workload_request import app.tasks.simple +import app.tasks.reconcile_online_workers \ No newline at end of file diff --git a/app/tasks/reconcile_online_workers.py b/app/tasks/reconcile_online_workers.py new file mode 100644 index 0000000..f785028 --- /dev/null +++ b/app/tasks/reconcile_online_workers.py @@ -0,0 +1,79 @@ +# app/tasks/reconcile_online_workers.py +""" +Celery task: reconcile worker status with Redis heartbeats. + +This task scans all workers marked as "online" in the database. +If a worker has not sent a heartbeat within the timeout window, +it will be marked as "offline" via API or direct update. + +This protects against undetected disconnects, e.g., when a WebSocket +server crashes or workers silently die without a disconnect event. +""" + +from __future__ import annotations +import time + +import redis +from app import celery_app as celery, db, logger +from app.models.models import WorkloadHost # Replace with your actual Worker model +from app import app +from datetime import datetime, timedelta + +API_URL = "http://localhost:5000/api/internal/workers" # Adjust to your environment +redisclient = redis.from_url(app.config["REDIS_URL"]) + + +@celery.task(name="tasks.reconcile_online_workers", bind=True) +def reconcile_online_workers(self) -> None: + """ + Check Redis for last heartbeat timestamps and reconcile online status of workers. + + This task ensures that any worker marked 'online' in the database + but missing a recent heartbeat is transitioned to 'offline'. + """ + logger.info("Starting worker heartbeat reconciliation task") + + now = int(time.time()) + + # Step 1: Get all workers marked as online + # online_workers = WorkloadHost.query.filter_by(_status="online",deleted=0).all() + + cutoff = datetime.utcnow() - timedelta(seconds=10) + + online_workers = WorkloadHost.query.filter( + WorkloadHost._status == "online", + WorkloadHost.deleted == 0, + WorkloadHost.updated_at < cutoff + ).all() + logger.info(f"Found {len(online_workers)} workers marked as online") + + stale_workers = [] + + worker: WorkloadHost # Strongly type so that VSCode IDE can autocomplete + + for worker in online_workers: + logger.debug(f"Checking worker {worker}") + redis_key = f"ws_liveness:{worker.id}" + last_seen = redisclient.get(redis_key) + + if not last_seen: + logger.warning(f"No heartbeat found for WorkloadHost {worker.id}") + is_stale = True + else: + try: + last_seen = int(last_seen) + is_stale = (now - last_seen) > app.config["PING_HEARTBEAT_TIMEOUT_SECONDS"] + except Exception: + logger.exception(f"Invalid heartbeat data for WorkloadHost {worker.id}") + is_stale = True + + if is_stale: + logger.warning(f"WorkloadHost {worker.id} is stale — marking offline") + stale_workers.append(worker.id) + + + # Set worker offline + worker.set_status("offline") + db.session.commit() + + logger.info(f"Finished reconciliation. Marked {len(stale_workers)} WorkloadHosts offline") diff --git a/websocket_server/config.py b/websocket_server/config.py index 87a0a94..e1ef95e 100644 --- a/websocket_server/config.py +++ b/websocket_server/config.py @@ -10,16 +10,23 @@ from sqlalchemy.orm import sessionmaker, declarative_base import pymysql # Redis configuration -REDIS_HOST = "172.17.0.1" -REDIS_PORT = 6379 +REDIS_URL = "redis://172.17.0.1:6379/2" + + +# Configurable intervals +PING_INTERVAL_SECONDS = 2 # How often to send PING requests +REDIS_PING_EXPIRY_SECONDS = 30 # Redis key expiry value +ASSIGN_INTERVAL_SECONDS = 30 +LIVENESS_CHECK_INTERVAL_SECONDS = 5 # How often to cross check connected_workers with ping results from Redis +PING_EXPIRY_SECONDS = 3 # How long between PING and PONG before we raise the alarm + # Database configuration DATABASE_URL = "mysql://root:password@172.17.0.1:3306/theapi" # Connection pool for Redis -redis_connection_pool = redis.ConnectionPool( - host=REDIS_HOST, - port=REDIS_PORT, +redis_connection_pool = redis.ConnectionPool.from_url( + REDIS_URL, decode_responses=True, max_connections=10, socket_timeout=5 diff --git a/websocket_server/events/messages.py b/websocket_server/events/messages.py index 5d0f64b..0279d0d 100644 --- a/websocket_server/events/messages.py +++ b/websocket_server/events/messages.py @@ -1,14 +1,14 @@ # websocket_server/events/messages.py +import time 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 +from websocket_server.shared_state import connected_workers, connected_sids_lock, connected_sids, worker_lock, ping_tracker +from websocket_server.config import REDIS_PING_EXPIRY_SECONDS logger = logging.getLogger("websocket_server") @@ -115,3 +115,26 @@ def register_socketio_handlers(socketio): except Exception as e: logger.exception(f"Failed to send message to {worker_id}: {e}") emit("error", {"message": f"Failed to send message to {worker_id}"}) + + @socketio.on("worker_pong") + def handle_worker_pong(data): + worker_id = data.get("worker_id") + pong_id = data.get("ping_id") + redis = get_redis_client() + + if not worker_id or not pong_id: + logger.warning("Malformed pong received") + return + + ping_info = ping_tracker.get(worker_id) + if not ping_info or ping_info["ping_id"] != pong_id: + logger.warning(f"[{worker_id}] Received unmatched pong") + return + + latency_ms = int((time.time() - ping_info["sent_time"]) * 1000) + logger.info(f"[{worker_id}] Pong received. Latency: {latency_ms}ms") + + redis.set(f"ws_liveness:{worker_id}", int(time.time()), ex=REDIS_PING_EXPIRY_SECONDS) + redis.set(f"ws_latency:{worker_id}", latency_ms, ex=REDIS_PING_EXPIRY_SECONDS) + + del ping_tracker[worker_id] diff --git a/websocket_server/shared_state.py b/websocket_server/shared_state.py index 5c7a445..21bfc39 100644 --- a/websocket_server/shared_state.py +++ b/websocket_server/shared_state.py @@ -6,3 +6,4 @@ connected_workers = {} worker_lock = threading.Lock() connected_sids = set() connected_sids_lock = threading.Lock() +ping_tracker = {} # Maps worker_id → {"ping_id": str, "sent_time": float} diff --git a/websocket_server/worker_manager.py b/websocket_server/worker_manager.py index 7dde48a..8838202 100644 --- a/websocket_server/worker_manager.py +++ b/websocket_server/worker_manager.py @@ -1,12 +1,15 @@ # websocket_server/worker_manager.py -import threading import requests +import uuid +import time +import threading import logging -# from websocket_server.redis_utils import redis_subscribe, thread_stop_flags, worker_dispatch_flags +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.config import get_redis_client -from websocket_server.shared_state import connected_workers, worker_lock +from websocket_server.events import base + logger = logging.getLogger("websocket_server") @@ -16,6 +19,13 @@ 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. @@ -65,15 +75,48 @@ def worker_dispatch_flag_check(): def all_worker_watchdog(): """ - Periodically iterate over connected workers and ensure they are polled for task assignments. - Prevents idle starvation or missed Redis events. + 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: - for wid in list(connected_workers.keys()): - logger.info(f"[{wid}] Watchdog triggering assign check") - assign_task_to_worker(wid) - threading.Event().wait(30) + 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): diff --git a/worker/workerClient.py b/worker/workerClient.py index d7af66c..59eacba 100644 --- a/worker/workerClient.py +++ b/worker/workerClient.py @@ -73,6 +73,20 @@ class WorkerClient: self.sio.on("stop_vnc_stream_on_worker", self.stop_vnc_stream) self.sio.on("vnc_frame_from_novnc", self.vnc_frame_from_novnc) # Add this line + self.sio.on("worker_ping", self.handle_worker_ping) + + async def handle_worker_ping(self, data): + """ + Respond to a liveness ping from the WebSocket server. + Emits a 'worker_pong' event with the same ping_id and this worker's ID. + """ + try: + ping_id = data["ping_id"] + worker_id = self.worker_id + await self.sio.emit("worker_pong", {"ping_id": ping_id,"worker_id": worker_id}) + logger.debug(f"[{worker_id}] Responded to ping {ping_id}") + except Exception as e: + logger.error(f"Error handling worker_ping: {e}") async def on_connect(self): logger.info("Connected to the server, requesting to join.") @@ -224,10 +238,6 @@ class WorkerClient: await self.sio.emit("docker_event", payload) logger.debug(f"Sent {source} event: {event_type}\n{payload}") - # async def debug_all_events(self, event, data): - # """Debug all incoming data.""" - # logger.debug(f"Event: {event} | Data: {json.dumps(data, indent=2)}") - async def process_event_queue(self): """Process events from the queue and send them to the server.""" if not self.event_queue: @@ -265,6 +275,8 @@ class WorkerClient: """Stop the worker client.""" self.running = False + + async def logs_retrieve(self,data): logger.info(f"Handling container-log {data}") @@ -287,8 +299,6 @@ class WorkerClient: await self.sio.emit("container_log_response", log_data) logger.debug(f"Sent container-log-response for {container_name}") - - async def logs_stream_start(self, data): request_id = data["request_id"] container_name = data["container_name"]