import glob
import os
import threading
import random
from typing import List, Tuple, Callable, Optional


import numpy as np
import pandas as pd
import cv2
import pydicom
import torchvision
from pydicom.pixel_data_handlers.util import apply_voi_lut
from tqdm.notebook import tqdm

import torch
from torch.utils.data import Dataset, DataLoader

import matplotlib.pyplot as plt
from matplotlib import animation, rc
import seaborn as sns
rc('animation', html='jshtml')


def save_image_statistic(input_folder:str, output_file:str, image_size: int = 224, n_threads: int = 8):
    # get all files in folder
    paths = glob.glob(input_folder + r'/**/*.dcm', recursive=True)
    # compute chunk size
    
    # split paths into lists of equal size
    thread_paths = list(split(paths, n_threads))
    
    # create and run threads
    threads = []
    file_names = []
    for i in range(n_threads):
        file_names.append(output_file[:-4] + f'_temp_{i}.csv')
        t = threading.Thread(target=get_images_stats, args=[thread_paths[i], file_names[i], image_size], name="get_statistic")
        threads.append(t)
        t.start()
    
    # waiting for all threads to finish
    for t in threads:
        t.join()
    
    # read all files of threads
    df_list = []
    for file_name in file_names:
        df_list.append(pd.read_csv(file_name))
        os.remove(file_name)
    
    data = pd.concat(df_list)
    data.to_csv(output_file, index=False)
    return data

    
def split(a, n):
    k, m = divmod(len(a), n)
    return (a[i*k+min(i, m):(i+1)*k+min(i+1, m)] for i in range(n))


def get_images_stats(paths, file_name, image_size) -> pd.DataFrame:
    stats = []
    for path in tqdm(paths):
        image = load_dicom_image(path, image_size)
        split = path.split('/')
        stats.append({
            'path' : path,
            'patient' : split[-3],
            'mri_type' : split[-2],
            'min' : image.min(),
            'max' : image.max(),
            'mean' : image.mean(),
            'median' : np.median(image),
            'std' : image.std()
        })
    
    pd.DataFrame(stats).to_csv(file_name, index=False)
    
    
def load_dicom_image(path, img_size=224, voi_lut=True):   
    dicom = pydicom.read_file(path)
    if voi_lut:
        data = apply_voi_lut(dicom.pixel_array, dicom)
    else:
        data = dicom.pixel_array
        
    data = cv2.resize(data, (img_size, img_size))
    return data


def save_image_metadata(input_folder:str, output_file:str) -> pd.DataFrame :
    """
    recursivly scan images in input_folder, 
    get metadata
    write to .csv file and return pd.DataFrame
    """
    metadata = []
    for path in tqdm(glob.glob( input_folder + r'/**/*.dcm', recursive=True)):
        dicom = pydicom.read_file(path)
        dicom.decode()
        df = pd.Series(dicom.values())[:-1]  # exclude pixel data
        
        keys = np.append(df.apply(lambda x: x.name).values, ['path'])
        values = np.append(df.apply(lambda x: x.value).values, [path])
        metadata.append(dict(zip(keys, values)))
    df = pd.DataFrame(metadata)
    df = df.sort_values(by='path')
    df.to_csv(output_file, index=False)
    print(f'{output_file} save successful.')
    
    return df


def filter_visualization(train:pd.DataFrame, test:pd.DataFrame, 
                         filter_function:Callable[..., pd.DataFrame], 
                         to_vis:Callable[[pd.DataFrame], pd.Series], 
                         title:str='', **kwargs
                        ) -> Tuple[pd.DataFrame, pd.DataFrame]:
    
    fig, axs = plt.subplots(2, 2, figsize=(30,8))
    fig.suptitle(title, fontsize=16, y=1)
    
    axs[0,0].set_title(f'train before')
    to_vis(train).plot(ax=axs[0,0], kind='hist', bins=200)
    axs[0,1].set_title(f'test before')
    to_vis(test).plot(ax=axs[0,1], kind='hist', bins=200, color='orange')
    
    print('train - ', end='')
    f_train = filter_function(train, **kwargs)
    print('test  - ', end='')
    f_test = filter_function(test, **kwargs)
    
    axs[1,0].set_title(f'train after')
    to_vis(f_train).plot(ax=axs[1,0], kind='hist', bins=200)
    axs[1,1].set_title(f'test after')
    to_vis(f_test).plot(ax=axs[1,1], kind='hist', bins=200, color='orange')
    
    return f_train, f_test


def show_images(images:Optional[List[np.array]]=None, 
                paths:Optional[List[str]]=None, 
                dataloader:Optional[DataLoader]=None,
                images_to_show:Optional[int]=None,
                title:Optional[str]=None,
                dpi:int=30):
    """
    Function to show multiple images.
    :param images: if not None, show this images.
    :param paths: if not None, read images from paths and show them instead of the images argument
    :param dataloader: if not None, get batch of images and show them instead of the images argument
    :param images_to_show: number of first images to show. If None, show all
    :param title: if not None, add title
    
    """
    normalize = True
    
    if paths is not None:
        images = [torch.as_tensor(load_dicom_image(p)).unsqueeze(0) for p in paths]
            
    if dataloader is not None:
        batch = next(iter(dataloader))
        images = batch['X'].cpu()  # for single images
        b,c,h,w = batch['X'].shape  # for stacked images
        images = batch['X'].cpu().reshape((-1,1,h,w))  # colapse batch and channels in one dimension
#         images = images[np.random.randint(0, b*c - 1, images_to_show),:,:,:]  # get only N random images. N = images_to_show
        normalize = False
    
    if images_to_show is not None:
        images = images[:images_to_show]
        
    
    grid = torchvision.utils.make_grid(images, normalize=normalize)

    fig = plt.figure(dpi=dpi, figsize=(grid.size()[2] / dpi, grid.size()[1] / dpi))
    
    
    if title is not None:
#         fig.suptitle(title, y=0.95, fontsize=6)  # i can't draw it correct
        print(title)

    plt.axis('off')
    plt.imshow(np.transpose(grid, (1,2,0)))
    
    
def create_animation(ims):
    # https://www.kaggle.com/ihelon/brain-tumor-eda-with-animations-and-modeling
    fig = plt.figure(figsize=(6, 6))
    plt.axis('off')
    im = plt.imshow(ims[0], cmap="gray")

    def animate_func(i):
        im.set_array(ims[i])
        return [im]

    return animation.FuncAnimation(fig, animate_func, frames = len(ims), interval = 1000//24)
    
    
def remove_low_contrast_images(df:pd.DataFrame, threshold:float=0.02) -> pd.DataFrame:
    before = df.shape[0]
    df = df[(df['std'] / df['gr_range']) > threshold]
    after = df.shape[0]
    print(f'before removing: {before:>7}, after:{after:>7}, removed: {before-after:>7} / {(before-after)/before:>3.0%}')
    return df

