###########################################################################
# MLHandler.py
#
# Modified Version of lcm_image_test.py from Ben and Jon at CVision
#
# The modifications support:
#   (1) variable input image size
#   (2) Variable image display size
#   (3) Variable image format (mono or rgb)
#   (4) Real-time display of images using either qt signals or opencv
#
# Adapted:  Nov 2019      Paul Roberts
#
###########################################################################

import datetime
import cv2
import os
import numpy as np
import time
import logging
from collections import deque
from image import image_t
from mbari import ext_target_t, ext_target_wheader_t
from cvision import bounding_box_t
from pyqtgraph.Qt import QtCore
from stereo_calc import StereoCalc
from lcm import LCM
from annotation import Annotation

logging.basicConfig(
    format='%(asctime)s %(levelname)-8s %(message)s',
    level=logging.DEBUG,
    datefmt='%Y-%m-%d %H:%M:%S',
)
logger = logging.getLogger(__name__)


class Handler(QtCore.QThread):

    stereoFrame = QtCore.Signal(object)

    @classmethod
    def init_class(cls, args):
        fourcc = cv2.VideoWriter_fourcc('M', 'J', 'P', 'G')
        # cls.writer = cv2.VideoWriter(video_path, fourcc, 10, (1032, 1544))
        cls.display_width = args.display_width
        cls.display_height = args.display_height
        cls.box_width = 120
        cls.box_height = 120
        cls.imgs = [deque([], maxlen=10), deque([], maxlen=10)]

        cls.img_buffer = np.zeros((cls.display_height * 2, cls.display_width, 3), dtype='uint8')
        cls.start_time = datetime.datetime.now()

        # setup KCF tracker
        cls.kcf_trackers = []
        cls.kcf_bboxes = []

        # Setup SimpleBlobDetector parameters.
        cls.params = cv2.SimpleBlobDetector_Params()

        cls.bg_detector_slow = cv2.createBackgroundSubtractorKNN(500)
        cls.bg_detector_mid = cv2.createBackgroundSubtractorKNN(50)
        cls.bg_detector_fast = cv2.createBackgroundSubtractorKNN(5)

        # Change thresholds
        cls.params.minThreshold = 5
        cls.params.maxThreshold = 100

        # Filter by Area.
        cls.params.filterByArea = True
        cls.params.minArea = 50

        # Filter by Circularity
        cls.params.filterByCircularity = False
        cls.params.minCircularity = 0

        # Filter by Convexity
        cls.params.filterByConvexity = False
        cls.params.minConvexity = 0

        # Filter by Inertia
        cls.params.filterByInertia = False
        cls.params.minInertiaRatio = 0

        cls.last_save_utime = int(time.time())

        # Create a detector with the parameters
        cls.detector = cv2.SimpleBlobDetector_create(cls.params)


    @classmethod
    def set_lcm_dirs(cls, folder, file):
        cls.lcm_log_file = file
        cls.lcm_log_folder = folder

    @classmethod
    def set_data_dirs(cls, left_dir, right_dir):

        # Create dirs for storing output
        cls.left_data_dir = left_dir + str(int(time.time()))
        cls.right_data_dir = right_dir + str(int(time.time()))
        if not os.path.exists(Handler.left_data_dir):
            os.makedirs(os.path.join(Handler.left_data_dir, 'xml'))
            os.makedirs(os.path.join(Handler.left_data_dir, 'images'))
            os.makedirs(os.path.join(Handler.left_data_dir, 'boxed_images'))
            os.makedirs(os.path.join(Handler.left_data_dir, 'rois'))
        if not os.path.exists(Handler.right_data_dir):
            os.makedirs(os.path.join(Handler.right_data_dir, 'xml'))
            os.makedirs(os.path.join(Handler.right_data_dir, 'images'))
            os.makedirs(os.path.join(Handler.right_data_dir, 'boxed_images'))
            os.makedirs(os.path.join(Handler.right_data_dir, 'rois'))

    @classmethod
    def set_box_width(cls, width):
        cls.box_width = width

    @classmethod
    def set_box_height(cls, height):
        cls.box_height = height

    @classmethod
    def set_score_threshold(cls, score_threshold):
        logger.info(score_threshold)
        cls.score_threshold = float(score_threshold)/100

    @classmethod
    def set_match_classes(cls, should_match):
        logger.info(should_match)
        cls.match_classes = should_match

    @classmethod
    def set_show_bg_mask(cls, show_bg_mask):
        logger.info(show_bg_mask)
        cls.show_bg_mask = show_bg_mask

    @classmethod
    def set_show_ml(cls, show_ml):
        logger.info(show_ml)
        cls.show_ml = show_ml

    @classmethod
    def set_save_kcf(cls, save_kcf):
        logger.info(save_kcf)
        cls.save_kcf_rois = save_kcf

    @classmethod
    def set_save_ml(cls, save_ml):
        logger.info(save_ml)
        cls.save_ml_rois = save_ml

    @classmethod
    def set_match_method(cls, method):
        cls.match_method = method
        logger.info(method)

    @classmethod
    def set_target_class(cls, target_class):
        cls.target_class_name = target_class
        logger.info(cls.target_class_name)

    @classmethod
    def set_box_deviation(cls, dev):
        logger.info(dev)
        cls.box_match_threshold = dev

    def __init__(self, index):
        QtCore.QThread.__init__(self)
        self.index = index
        self.num_images = 0
        self.init_kcf = False


    def __call__(self, channel, data):
        Handler.imgs[self.index].append(image_t.decode(data))
        if (len(Handler.imgs[0]) > 0) and (len(Handler.imgs[1]) > 0):
            left = Handler.imgs[0].popleft()
            right = Handler.imgs[1].popleft()
            Handler.img_buffer[:self.display_height, :, :] = self.unpack_image(left)
            Handler.img_buffer[self.display_height:, :, :] = self.unpack_image(right)
            self.stereoFrame.emit(Handler.img_buffer)

