import socket import threading import traceback import redis import time import json import logging import signal import random import string import os import sys import contextlib from .metrics import put_queue_wait_metric # pylint: disable=relative-beyond-top-level # pylint: disable=relative-beyond-top-level from .signal_handling import SignalProtector # pylint: disable=relative-beyond-top-level from .perf_recorder import JSONPerfRecorder class NoDefault: pass NO_DEFAULT = NoDefault() def int_upon_1000(x): return int(x) / 1000.0 def getenv(varname, default=NO_DEFAULT, cast=str): val = os.getenv(varname, None) if val is None: if default is NO_DEFAULT: raise RuntimeError(f"Missing environment variable {varname}") return default return cast(val) class RedisClient: def __init__( self, queue_name, queue_time_name, max_runtime, redis_host, redis_port, redis_sentinel_name, redis_password, healthcheck_path, ): self.queue_name = queue_name self.queue_time_name = queue_time_name self.max_runtime = max_runtime self.redis_host = redis_host self.redis_port = redis_port self.redis_sentinel_name = redis_sentinel_name self.redis_password = redis_password self.healthcheck_path = healthcheck_path self.worker_name = self._new_worker_name() self.version = self._get_version() self.client, self.sentinel_client = self._new_redis_client() self._redrive_key = f"work:{self.worker_name}" self._liveness_key = f"worker:{self.worker_name}" def _new_worker_name(self): random_id = "".join( random.choices(string.ascii_lowercase + string.digits, k=10) ) return f"{self.queue_name}:{random_id}" def _get_version(self): try: with open("version", "r") as f: return f.read().strip() except FileNotFoundError: return "unknown" def _new_redis_client(self): socket_kwargs = dict( socket_keepalive=True, socket_keepalive_options=({socket.TCP_NODELAY: 1}), socket_connect_timeout=5, socket_timeout=60, ) redis_kwargs = dict( socket_kwargs, db=0, retry=redis.retry.Retry(redis.backoff.ExponentialBackoff(), 5), retry_on_error=[ redis.exceptions.ConnectionError, ConnectionError, TimeoutError, ], password=self.redis_password, ) sentinel_kwargs = dict( socket_kwargs, password=self.redis_password, ) if self.redis_sentinel_name is not None: sentinel = redis.Sentinel( [(self.redis_host, self.redis_port)], sentinel_kwargs=sentinel_kwargs, **redis_kwargs, ) return sentinel.master_for(self.redis_sentinel_name), sentinel return ( redis.Redis( host=self.redis_host, port=self.redis_port, **redis_kwargs, ), None, ) @classmethod def from_env(cls, *args): return cls( *args, getenv("QUEUE_NAME"), getenv("QUEUE_TIME_NAME", None), getenv("MAX_RUNTIME", None, int_upon_1000), getenv("REDIS_HOST", "localhost"), getenv("REDIS_PORT", 6379, int), getenv("REDIS_SENTINEL_NAME", None), getenv("REDIS_PASSWORD", None), getenv("HEALTHCHECK_PATH", "/ready"), ) @contextlib.contextmanager def _short_switch_interval(self): swi = sys.getswitchinterval() sys.setswitchinterval(0.001) try: yield finally: sys.setswitchinterval(swi) @contextlib.contextmanager def _sigterm_as_sigint(self): def handler(sig, frame): raise KeyboardInterrupt() old_handler = signal.signal(signal.SIGTERM, handler) try: yield finally: signal.signal(signal.SIGTERM, old_handler) def warmup(self): pass def handle_request(self, request_data, recorder, notify_progress, times_out_at): raise NotImplementedError() def write_healthcheck(self): if self.healthcheck_path is not None: try: with open(self.healthcheck_path, "w") as f: f.write("ready") except (PermissionError, FileNotFoundError): logging.info(f"Failed to write {self.healthcheck_path}") def _liveness_loop(self, quit_evt, ready_evt): while not quit_evt.is_set(): try: self.client.set(self._liveness_key, "alive", ex=3) ready_evt.set() except Exception as exn: logging.exception(exn) time.sleep(1) @contextlib.contextmanager def _alive_to_redis(self): quit_evt = threading.Event() ready_evt = threading.Event() liveness_thread = threading.Thread( target=self._liveness_loop, args=(quit_evt, ready_evt) ) liveness_thread.start() ready_evt.wait() try: yield finally: quit_evt.set() liveness_thread.join() def run(self): with self._short_switch_interval(), self._sigterm_as_sigint(), self._alive_to_redis(): self.warmup() if self.queue_time_name is not None: self.client.delete(self.queue_time_name) self.write_healthcheck() while True: self._run_once() def _clear_redrive(self, conn): conn.delete(self._redrive_key) def _run_once(self): brpop_res = self.client.brpoplpush( self.queue_name, self._redrive_key, timeout=10 ) if brpop_res is None: return with SignalProtector(): time_dequeued = time.time() request = json.loads(brpop_res.decode("utf-8")) request_id = request.get("id", None) if request_id is None: logging.warning("Ignoring request without id") self._clear_redrive(self.client) return request_data = request.get("data", None) if request_data is None: logging.warning("Ignoring request without data") self._clear_redrive(self.client) return recorder = JSONPerfRecorder() response = { "version": self.version, } def notify_progress(progress): with self.client.pipeline() as pipe: pipe.rpush(request_id, json.dumps({"progress": progress})) pipe.expire(request_id, 60) pipe.execute() try: times_out_at = request.get("timesOutAt", None) if ( times_out_at is not None and self.max_runtime is not None and time_dequeued + self.max_runtime > times_out_at ): raise RuntimeError("Not enough time to process request") response.update( self.handle_request( request_data, recorder, notify_progress, times_out_at ) ) except Exception as exn: logging.exception(exn) response.update( error=str(exn), backtrace=traceback.format_exception(exn) ) response.update(perf=recorder.to_object()) with self.client.pipeline() as pipe: self._clear_redrive(pipe) pipe.lpush(request_id, json.dumps(response)) pipe.expire(request_id, 60) pipe.execute() time_submitted = request.get("timeSubmitted", None) if time_submitted is not None: if self.queue_time_name is not None: with self.client.pipeline() as pipe: pipe.lpush(self.queue_time_name, time_dequeued - time_submitted) pipe.ltrim(self.queue_time_name, 0, 100) pipe.expire(self.queue_time_name, 60) pipe.execute() put_queue_wait_metric(self.queue_name, time_dequeued - time_submitted)