import cv2
import imageio
import numpy as np
import matplotlib.pyplot as plt
from scipy.spatial import KDTree
from skimage.color import label2rgb
from skimage.measure import label, regionprops

class DepthImage:

    def __init__(self, file_prefix, types="XYZA", debug=False):
        self.depth_images = []
        self.depth_mask = None
        self.types = types
        self.z_index = self.types.find("Z")
        self.file_prefix = file_prefix
        self.debug = debug
        self.load()

    def load(self):
        for t in self.types:
            file_name = self.file_prefix + '-DepthMap-' + t + '.tiff'
            self.depth_images.append(np.flipud(imageio.imread(file_name)))

    def mask(self, image_mask):

        # rescale mask as needed
        if not image_mask.shape == self.depth_images[0].shape:
            image_mask = cv2.resize(image_mask, self.depth_images[0].shape)
        # apply mask to each channel
        for i in range(len(self.types)):
            self.depth_images[i] = self.depth_images[i] * image_mask
        # apply mask to depth mask
        self.depth_mask = self.depth_mask * image_mask

    def get_depth_mask(self, depth_range=[-500, 500]):
        if self.depth_mask is None:
            tmp_img = self.depth_images[self.z_index]
            tmp_img[np.isnan(tmp_img)] = depth_range[0] - 100
            self.depth_mask = np.logical_and(tmp_img > depth_range[0], tmp_img <= depth_range[1]).astype('uint8')
        if self.debug:
            cv2.imshow('depth_mask', cv2.resize(self.depth_mask, (1024, 1024))*255)
            cv2.waitKey()
        return self.depth_mask

    def generate_point_cloud(self, depth_range=[-500, 500]):
        if self.depth_mask is None:
            self.get_depth_mask(depth_range)
        # label connected components in the depth image
        label_image = label(self.depth_mask)
        if self.debug:
            cv2.imshow('label_image', label2rgb(label_image, image=self.depth_images[2]))
            cv2.waitKey()
        self.point_cloud = []
        for region in regionprops(label_image):
            coords = region['coords']

            p = []
            p.append(np.mean(self.depth_images[0][coords[:, 0], coords[:, 1]]))
            p.append(np.mean(self.depth_images[1][coords[:, 0], coords[:, 1]]))
            p.append(np.mean(self.depth_images[2][coords[:, 0], coords[:, 1]]))
            p.append(len(coords[:, 0]))
            self.point_cloud.append(p)

        self.point_cloud = np.array(self.point_cloud)
        self.num_particles = self.point_cloud.shape[0]

        if self.point_cloud.shape[0] < 10:
            self.nnstat = 0
            self.nn_points = 0
        else:
            self.kdtree = KDTree(self.point_cloud)
            self.nndist, self.nn_points = self.kdtree.query(self.point_cloud, 2)
            self.nnstat = 2 * np.pi * self.num_particles * np.sum(self.nndist[:, 1] ** 2)

        if self.debug:
            print(self.nndist[:, 1])
            plt.hist(self.nndist[:, 1], bins=500)
            plt.show()
        return self.point_cloud