import os import time from collections import defaultdict from multiprocessing import Event, Pipe, Semaphore from multiprocessing.connection import Connection from typing import Any, Dict, Iterable, List, Optional, Tuple, Type from logger import get_logger from utils.key import Key from utils.slack import SlackReporter from . import get_step_depends_on from .base import BaseStep __all__ = ["Supervisor"] class Supervisor: def __init__(self, key: Key, max_retries: int = 0, max_processes: Optional[int] = os.cpu_count()): if not max_processes: max_processes = 1 self.__pool = Semaphore(max_processes) self.__stop_event = Event() self.__sleep_timeout = 0.001 self.__max_restarts = max_retries self.key = key self.logger = get_logger(str(key)) self.slack_reporter = SlackReporter(self.logger, self.key) # processes self.__waiting: List[Type[BaseStep]] = [] self.__running: List[BaseStep] = [] self.__done: List[BaseStep] = [] self.__restart_counter: Dict[Type[BaseStep], int] = defaultdict(int) self.__pipes: Dict[str, Connection] = {} def _init_step(self, step_cls: Type[BaseStep]) -> BaseStep: parent_conn, child_conn = Pipe() self.__pipes[step_cls.step_name] = parent_conn return step_cls(key=self.key, conn=child_conn) def _is_step_ready_to_start(self, step_cls: Type[BaseStep]) -> bool: return step_cls in self.__waiting and not bool( get_step_depends_on(step_cls) - set(step.__class__ for step in self.__done) ) def _get_result(self, step_name: str) -> Tuple[Dict[str, Any], Dict[str, Any]]: conn = self.__pipes[step_name] result = conn.recv() conn.close() return result def _check_running(self): for step in self.__running: if self.__stop_event.is_set(): break # process exited with error if step.exitcode: self._step_failed(step) continue # process done if step.is_done(): self._step_done(step) def _check_waiting(self): for step_cls in self.__waiting: if self.__stop_event.is_set(): break if self._is_step_ready_to_start(step_cls) and self.__pool.acquire(block=False): self._start_step(step_cls) def _check_all_done(self) -> bool: return not self.__waiting and not self.__running def _start_step(self, step_cls: Type[BaseStep]): self.logger.info(f"'{step_cls.step_name}' starting") self.slack_reporter.notify_step_started(step_cls.step_name) step = self._init_step(step_cls) step.start() self.__waiting.remove(step_cls) self.__running.append(step) def _step_done(self, step: BaseStep): self.logger.info(f"'{step.__class__.step_name}' is done") result, perf_report = self._get_result(step.__class__.step_name) self.slack_reporter.notify_step_done(step.__class__.step_name, perf_report) step.join() step.close() self.__pool.release() self.__running.remove(step) self.__done.append(step) def _step_failed(self, step: BaseStep): self.logger.error(f"'{step.__class__.step_name}' exited with code {step.exitcode}") self.slack_reporter.notify_step_failed(step.__class__.step_name) if self.__restart_counter[step.__class__] >= self.__max_restarts: self.logger.info("Max restarts reached. Exiting...") self.__stop_event.set() return self.logger.info("Restarting process") step.close() self.__running.remove(step) self.__waiting.append(step.__class__) self.__restart_counter[step.__class__] += 1 def _perf_report(self): for step in self.__done: self.logger.info(step.perf_counter.report) def _run(self, steps_cls: Iterable[Type[BaseStep]]): self.__waiting = list(sorted(steps_cls, key=lambda step_cls: step_cls.priority)) while True: self._check_running() self._check_waiting() if self._check_all_done(): break if self.__stop_event.is_set(): self.logger.info("Got stop event, terminating processing steps") for step in self.__running: step.terminate() step.join() self.slack_reporter.post_header("ETL process failed") exit(1) time.sleep(self.__sleep_timeout) self._perf_report() self.logger.info(f"ETL process with key {self.key} finished") self.slack_reporter.post_header("ETL processing finished") def run(self, steps_cls: Iterable[Type[BaseStep]]): self.logger.info(f"Starting ETL process with key {self.key}") self.slack_reporter.post_header("Starting ETL process") try: self._run(steps_cls) except Exception: self.slack_reporter.post_header("ETL process failed") raise