import utils

import numpy as np
import scipy as sp
import hyperspy.api as hs
from numba import jit

from skimage import measure, draw, morphology, segmentation

from tqdm import tqdm

params = {
    "MAX_AREA_CLOSING" : 800,
    "MAX_AREA_OPENING" : 200,
    "SD_THRESHOLD" : 2.5,
    "MIN_AREA_BLACK_WALLS" : 400,
    "MIN_AREA_WHITE_WALLS" : 200,
    "HORISTONTAL_THRESH" : 30,
    "NEIGHBOURHOOD_RADIUS" : 20,
    "RESIDUAL_THRESHOLD" : 5,
    "MAX_TRIALS" : 500,
    "NUMB_INTERMEDIATE_CIRCLES" : 5,
    "MIN_DOMAIN_AREA" : 800,
    "MAX_DOMAIN_AREA" : 10000,
}

def find_domains_and_areas(signal, params):
    """Returns a list of arrays with domains and a list of the correseponding areas.
    
    The main functionality of the developed algorithm for this project, namely locating 
    the ferro magnetic domains and determining their areas, can be found in this function. 

    Parameters:

    signal (Signal2D) : Hyperspy signal with dimensions { | X, Y}, ie. not a stack but a single signal

    params (dict) : A dictionary with the parameters used for the processing. 
    
    Returns:

    domains (np.arr) : Array containing the individual domains found by this
        algorithm. The np.arrays have a constant integer value at the (row,col) positions corresponding
        to the domain, otherwise zero. The indexing of this list corresponds to the indexign of the other
        return, $area.

    areas (list) : A list np.floats with the values of the areas found by this algorithm. The indexing
        corresponds to the other return, $domains_segmented.


    
     """

    
    arr = utils.shift_to_zero(utils.percentiles(signal.data, 1, 99))

    # noise filtering
    arr = morphology.area_closing(arr, params["MAX_AREA_CLOSING"])
    arr = morphology.area_opening(arr, params["MAX_AREA_OPENING"])
    mu, sd = np.mean(arr), np.std(arr)*params["SD_THRESHOLD"]
    arr[(arr < mu+sd) & (arr > mu-sd)] = mu
    arr = arr.astype("int")
    print("Noise reduction completed..")

    # segment black objects (diverging domain walls)
    black = np.copy(arr)
    mu = np.median(black)
    black[black > mu] = mu
    black_binary = np.copy(black)
    black_binary[black_binary == mu] = 0
    black_binary[black_binary > 0] = 1
    black_label = measure.label(black_binary, connectivity=2)

    black_label_values = np.unique(black_label)
    black_segments = []
    for val in black_label_values:
        segment = np.zeros(black.shape)
        segment[black_label == val] = val
        black_segments.append(segment)

    # relabel black objects after filtering out smaller objects
    black_segments = [s for s in black_segments if utils.area(s) > params["MIN_AREA_BLACK_WALLS"]]

    black_binary = np.copy(np.sum(black_segments, axis=0))
    black_binary[black_binary > 0] = 1
    black_label = measure.label(black_binary, connectivity=2)

    black_label_values = np.unique(black_label)[1:]
    black_segments = []
    for val in black_label_values:
        segment = np.zeros(black.shape)
        segment[black_label == val] = val
        black_segments.append(segment)

    # segment white objects (converging domain walls)
    white = np.copy(arr)
    mu = np.median(white)
    white[white < mu] = mu
    white_binary = np.copy(white)
    white_binary[white_binary == mu] = 0
    white_binary[white_binary > 0] = 1
    white_label = measure.label(white_binary, connectivity=2)

    white_label_values = np.unique(white_label)
    white_segments = []
    for val in white_label_values:
        segment = np.zeros(white.shape)
        segment[white_label == val] = val
        white_segments.append(segment)

    # relabel white objects after filtering out smaller objects
    white_segments = [s for s in white_segments if utils.area(s) > params["MIN_AREA_WHITE_WALLS"]]

    white_binary = np.copy(np.sum(white_segments, axis=0))
    white_binary[white_binary > 0] = 1
    white_label = measure.label(white_binary, connectivity=2)

    white_label_values = np.unique(white_label)[1:]
    white_segments = []
    for val in white_label_values:
        segment = np.zeros(white.shape)
        segment[white_label == val] = val
        white_segments.append(segment)

    # extract points closest to and furthest away from an approx. ring center
    cm_black = sp.ndimage.center_of_mass(black_binary)
    cm_white = sp.ndimage.center_of_mass(white_binary)

    black_segments_points = []
    for segment in black_segments:
        closest_point, outermost_point = utils.find_min_and_max_distances(segment, cm_black)
        dc = utils.distance(closest_point, cm_black) 
        df = utils.distance(outermost_point, cm_black)
        if np.abs(dc - df) > params["HORISTONTAL_THRESH"]:
            # filter out horisontal objects
            black_segments_points.append((closest_point, outermost_point))

    white_segments_points = []
    for segment in white_segments:
        closest_point, outermost_point = utils.find_min_and_max_distances(segment, cm_white)
        dc = utils.distance(closest_point, cm_white) 
        df = utils.distance(outermost_point, cm_white)
        if np.abs(dc - df) > params["HORISTONTAL_THRESH"]:
            # filter out horisontal objects
            white_segments_points.append((closest_point, outermost_point))

    # patch up objects that are disconnected
    black_segments_points_merged = utils.connect_neighbours(black_segments_points, R=params["NEIGHBOURHOOD_RADIUS"])  
    white_segments_points_merged = utils.connect_neighbours(white_segments_points, R=params["NEIGHBOURHOOD_RADIUS"])

    # RANSAC fitting of inner and outer circles
    inner_circle_points, outer_circle_points = [], []
    for points in black_segments_points:
        inner_circle_points.append(points[0])
        outer_circle_points.append(points[1])

    min_samples_inner = len(inner_circle_points) - 1
    min_samples_outer = len(outer_circle_points) - 1

    inner_circle, inner_circle_params = utils.draw_circle_based_on_ransac_fit(arr_shape=black.shape, 
                                                         circle_points=inner_circle_points,
                                                         min_samples=min_samples_inner,
                                                         residual_threshold=params["RESIDUAL_THRESHOLD"],
                                                         max_trials=params["MAX_TRIALS"])
    outer_circle, outer_circle_params = utils.draw_circle_based_on_ransac_fit(arr_shape=black.shape, 
                                                         circle_points=outer_circle_points,
                                                         min_samples=min_samples_inner,
                                                         residual_threshold=params["NEIGHBOURHOOD_RADIUS"],
                                                         max_trials=params["MAX_TRIALS"])
    circles = outer_circle + inner_circle

    print("Peripheries determined..")

    # draw intermediate circles for determination of points on the walls
    intermediate_circles = []
    dR = outer_circle_params["R"] - inner_circle_params["R"]
    for n in range(params["NUMB_INTERMEDIATE_CIRCLES"]):

        R = inner_circle_params["R"] + dR//(params["NUMB_INTERMEDIATE_CIRCLES"] + 1)*(n+1)
        cx, cy = inner_circle_params["x"], inner_circle_params["y"]

        rr, cc = draw.circle_perimeter(cx, cy, R, shape=inner_circle.shape)
        circle = np.zeros(inner_circle.shape)
        circle[rr, cc] = 1
        intermediate_circles.append(circle)


    # determine points on the walls crossing the intermediate circles
    black_object_points = []
    black_domain_walls = []
    for object_ in black_segments:
        object_binary = np.copy(object_)
        object_binary[object_binary > 0] = 1
        points = []
        for circle in intermediate_circles:
            intersection_points = np.where(object_binary + circle == 2)
            if intersection_points[0].size != 0 and intersection_points[1].size != 0:
                x, y = np.mean(intersection_points, axis=1) ##
                points.append(np.array([int(x), int(y)]))
        black_object_points.append(points)

        # approximate domain walls
        if len(points) < 2:
            wall = np.zeros(black.shape)
        else:
            wall = utils.draw_line_between_points(points, black.shape)
        black_domain_walls.append(wall)

    # repeat for white walls
    white_object_points = []
    white_domain_walls = []
    for object_ in white_segments:
        object_binary = np.copy(object_)
        object_binary[object_binary > 0] = 1
        points = []
        for circle in intermediate_circles:
            intersection_points = np.where(object_binary + circle == 2)
            if intersection_points[0].size != 0 and intersection_points[1].size != 0:
                x, y = np.mean(intersection_points, axis=1) ##
                points.append(np.array([int(x), int(y)]))
        white_object_points.append(points)

        if len(points) < 2:
            wall = np.zeros(white.shape)
        else:
            wall = utils.draw_line_between_points(points, white.shape)
        white_domain_walls.append(wall)

    # fit the extended walls to the ring structure
    mask_inner = utils.create_circular_mask(inner_circle_params["x"], 
                                            inner_circle_params["y"],
                                            inner_circle_params["R"],
                                            arr_shape=inner_circle.shape)
    mask_outer = utils.create_circular_mask(outer_circle_params["x"], 
                                            outer_circle_params["y"],
                                            outer_circle_params["R"],
                                            arr_shape=outer_circle.shape)

    mask_ring = ~(~mask_inner ^ mask_outer)

    black_domain_walls_fitted = np.sum(black_domain_walls, axis=0)
    black_domain_walls_fitted[mask_ring == 0] = 0

    white_domain_walls_fitted = np.sum(white_domain_walls, axis=0)
    white_domain_walls_fitted [mask_ring == 0] = 0

    circles_overlayed_with_domains = black_domain_walls_fitted + white_domain_walls_fitted + circles
    print("Domains determined..")

    # segmentation of domains (domain areas) by Felzenszwalb
    segments_fz = segmentation.felzenszwalb(circles_overlayed_with_domains, 
                                            multichannel=False)
    # filter out segments that are not domains
    values_fz = np.unique(segments_fz)
    domains_segmented = []
    domains = np.zeros(arr.shape)
    for value in values_fz:
        segment = np.zeros(arr.shape)
        segment[segments_fz == value] = 1
        segment_area = utils.area(segment)
        if segment_area > params["MIN_DOMAIN_AREA"] and segment_area < params["MAX_DOMAIN_AREA"]:
            domains_segmented.append(segment)
            domains[segment == 1] = 1
    domains = measure.label(domains, connectivity=2)

    pixel_to_area_scaling_factor = signal.axes_manager["x"].scale*signal.axes_manager["y"].scale
    areas = [np.sum(domain)*pixel_to_area_scaling_factor for domain in domains_segmented] # in um**2
    print("Areas determined..\n")
    
    return domains, areas

def process_whole_stack(hspy_filename, store_as_files=True):
    """Processing the whole dataset and storing the results as np.arrays in .txt files.
    
    The different steps of the alogrithms have not been optimized for performance yet, so processing
    the whole dataset takes about 60 mins. 

    Parameters:

    hspy_filename (str) : Name of the .hspy file containing the data set

    store_as_files (bool) : Wheter or not to store the results as .py files for lates use, default is True.
    
    Returns:

    areas (np.array) : Multidimensional array containing the areas found for the whole dataset. Shape
        (25, N_i), where N_i is the number of domains found in signal i.

    domains (np.array) : Multidimensiol array containing the domains found for the whole dataset, shape 
        (25, N_i, X, X), where N_i is the number of domains in signal i."""
    # load file
    s = hs.load(hspy_filename)

    # align ring structures in stack
    print("\nAligning..")
    shifts = s.estimate_shift2D()
    s.align2D(shifts=shifts)
    print("Aligningment finished..")

    # crop to isolate ring structure
    N, M = s.data[0].shape
    cx, cy = N//2, M//2
    R = 0.40*N

    xmin, xmax = int(cx - R), int(cx + R)
    ymin, ymax = int(cy - R), int(cy + R)
    s = s.isig[xmin:xmax, ymin:ymax]

    # find domains and areas for the whole stack
    domains_list = []
    areas_list = []
    idx = 1
    for signal in tqdm(s):
        print(f"\nImage {idx} of {25}")
        d, a = find_domains_and_areas(signal, params)
        domains_list.append(d)
        areas_list.append(a)
        idx += 1

    # store resutls as numpy files
    areas = np.asarray(areas_list)
    domains = np.asarray(domains_list)
    if store_as_files:
        np.save('areas.npy', areas)
        np.save('domains.npy', domains)

    return domains, areas
        

######################
LOAD_RESULTS = False
######################

def main():
    """Run program"""

    hspy_filename = "2021_03_26_FA721_A6.5_in_situ_stack.hspy"

    if LOAD_RESULTS:
        domains, areas = np.load("areas.npy"), np.load("domains.npy")
    else:
        domains, areas = process_whole_stack(hspy_filename)

    s_fz = hs.signals.Signal2D(domains, stack=True)
    s_fz.plot(navigator="slider", scalebar=False)

if __name__ == "__main__":
    main()