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)
