import numpy as np
from scipy.signal.windows import exponential
from scipy.ndimage.filters import convolve1d
from statsmodels.stats.multitest import multipletests
import scipy.stats
import scipy.optimize
from numba import njit


@njit
def makePSTH_numba(spikes, startTimes, windowDur, binSize=0.001):
    bins = np.arange(0,windowDur+binSize,binSize)
    counts = np.zeros(bins.size-1)
    for i,start in enumerate(startTimes):
        startInd = np.searchsorted(spikes, start)
        endInd = np.searchsorted(spikes, start+windowDur)
        counts = counts + np.histogram(spikes[startInd:endInd]-start, bins)[0]
    
    counts = counts/len(startTimes)
    return counts/binSize, bins


def exponential_convolve(response_vector, tau=1, symmetrical=False):
    
    center = 0 if not symmetrical else None
    exp_filter = exponential(10*tau, center=center, tau=tau, sym=symmetrical)
    exp_filter = exp_filter/exp_filter.sum()
    filtered = convolve1d(response_vector, exp_filter[::-1])
    
    return filtered


def calcHitRate(hits, misses, adjusted=False):
    n = hits + misses
    if n == 0:
        return np.nan
    hitRate = hits / n
    if adjusted:
        if hitRate == 0:
            hitRate = 0.5 / n
        elif hitRate == 1:
            hitRate = 1 - 0.5 / n
    return hitRate


def calcDprime(hits, misses, falseAlarms, correctRejects):
    hitRate = calcHitRate(hits, misses, adjusted=True)
    falseAlarmRate = calcHitRate(falseAlarms, correctRejects, adjusted=True)
    z = [scipy.stats.norm.ppf(r) for r in (hitRate, falseAlarmRate)]
    return z[0] - z[1]


def multiple_comparisons(pvalues):
    if isinstance(pvalues, dict):
        old_values = list(pvalues.values())
    else:
        old_values = pvalues
    
    reject, corrected_pvalues, _, _ = multipletests(old_values, alpha=0.05, method='fdr_bh')

    if isinstance(pvalues, dict):
        return {list(pvalues.keys())[i]: corrected_pvalues[i] for i in range(len(pvalues))}

    return reject, corrected_pvalues


def comparison_matrix(*values, test_func=scipy.stats.wilcoxon):

    pvalue_matrix = np.full((len(values), len(values)), np.nan)
    for ind1, vals1 in enumerate(values):
        for ind2, vals2 in enumerate(values):
            if ind1==ind2:
                continue
            
            p = test_func(vals1, vals2, nan_policy='omit')
            pvalue_matrix[ind1, ind2] = p.pvalue
    
    diag_mask = np.ones(pvalue_matrix.shape, dtype=bool)
    diag_mask[np.diag_indices_from(diag_mask)] = False
    corrected_pvals = multipletests(pvalue_matrix[np.where(diag_mask)], method='fdr_bh')

    pvalue_matrix[np.where(diag_mask)] = corrected_pvals[1]
    sig_matrix = pvalue_matrix<0.05


    return pvalue_matrix, sig_matrix


def fitCurve(func,x,y,initGuess=None, bounds=None):
    if bounds is None:
        fit = scipy.optimize.curve_fit(func,x,y,p0=initGuess, maxfev=100000)[0]
    else:
        fit = scipy.optimize.curve_fit(func,x,y,p0=initGuess,bounds=bounds,maxfev=100000)[0]
    return fit


def calc_gompertz(x, a, b, c, d):
    return d + (a-d)*np.exp(-np.exp(-b*(x-c)))


def invert_gompertz(y, xs, a, b, c, d):

    vals = np.array([calc_gompertz(x, a, b, c, d) for x in xs])
    yind = np.where(vals<=y)[0]
    if len(yind)==0:
        return np.nan, np.nan
    else:
        return xs[yind[-1]], vals[yind[-1]]


def find_midpoint_raw(x, y):
    y = np.array(y)
    maxd = max(y)
    mind = min(y)
    midpoint = (maxd-mind)/2 + mind

    after_ind = np.where(y>midpoint)[0][0]
    before_ind = np.where(y[:after_ind]<midpoint)[0][-1]

    xbefore = x[before_ind]
    xafter = x[after_ind]

    ybefore = y[before_ind]
    yafter  = y[after_ind]

    slope = (yafter-ybefore)/(xafter-xbefore)

    rise_to_midpoint = midpoint-ybefore
    run_to_midpoint = rise_to_midpoint/slope

    return xbefore+run_to_midpoint, midpoint
