import numpy as np
from scipy.sparse import csr_matrix
from scipy import special
import pandas as pd

def findCentroids(x,y,t=1,k=1):
    x,y,_filter = peak_filter(x,y)
    # Create Sparse Matrix
    _filterData = _filter > 0
    columnIndex = _filter[_filterData]
    columnIndex -= 1
    rowIndex = np.arange(len(columnIndex))

    X = x[_filterData]
    Y = y[_filterData]

    x_sparse = csr_matrix((X, (rowIndex,columnIndex)))
    y_sparse = csr_matrix((Y, (rowIndex,columnIndex)))
    ones_sparse = csr_matrix((X**0, (rowIndex,columnIndex)))

    x_sparse_centroids = x_sparse.T.dot(y_sparse)
    x_sparse_centroids.data = x_sparse_centroids.data / ones_sparse.T.dot(y_sparse).data
    numberOfnnz = ones_sparse.T.dot(ones_sparse).data.astype(int)
    x_centroids = x_sparse_centroids.data
    x_centroids_repeated = []
    weight_sparse_means_repeated = []
    weight_sparse = csr_matrix((Y**2,(rowIndex,columnIndex)))
    weight_sparse_means = weight_sparse.T.dot(weight_sparse).data / ones_sparse.T.dot(ones_sparse).data
    for u in range(len(numberOfnnz)):
        x_centroids_repeated = np.append(x_centroids_repeated, x_centroids[u] * np.ones(numberOfnnz[u]))
        weight_sparse_means_repeated = np.append(weight_sparse_means_repeated, weight_sparse_means[u] * np.ones(numberOfnnz[u]))
    x = X-x_centroids_repeated
    weight_sparse = csr_matrix((Y**2 / weight_sparse_means_repeated,(rowIndex,rowIndex)))

    # Linear Regression
    # Creating Vandermonde Matrix
    xV = np.concatenate((np.ones(len(x)), x, x**2))
    rI = np.concatenate((rowIndex, rowIndex, rowIndex))
    cI = np.concatenate((columnIndex*3, columnIndex*3+1, columnIndex*3+2)).astype(int)
    xV = csr_matrix((xV,(rI,cI)))
    # Calculate Coefficients
    beta = xV.T.dot(weight_sparse).dot(xV)
    inverse = inverse3x3(beta)
    ylog = np.log(Y)
    beta = inverse.dot(xV.transpose()).dot(weight_sparse).dot(ylog)
    # Transform Linear Coefficients into Gaussian Coefficients
    x0 = beta[1::3] / (-2) / beta[2::3] + x_centroids
    width = abs(1 / 2 / beta[2::3])**.5
    height = np.exp(beta[0::3] - beta[1::3]**2 / 4 / beta[2::3])
    area = height * width * (2*np.pi)**.5
    # Error Propagation
    yfit = xV*beta
    residues = yfit - ylog
    residues = csr_matrix((residues,(rowIndex, columnIndex)))
    RSS = residues.T.dot(weight_sparse).dot(residues)
    df  = numberOfnnz-np.ones(len(numberOfnnz))*3
    MSE = RSS.data / df
    varC = np.repeat(MSE,3)*inverse.diagonal()
    dx0 = (abs(varC[1::3] / 4 / beta[2::3]**2) + (beta[1::3] / 2 / beta[2::3]**2)**2 * varC[2::3])**.5
    k1 = np.exp(2*beta[0::3] - beta[1::3]**2 / 2 / beta[2::3])
    k2 = (2*np.pi)**.5 * np.exp(beta[0::3]-beta[1::3]**2 / 4 / beta[2::3])
    k3 = (1 / 2 / abs((beta[2::3])))**.5
    d0 = abs((np.pi * k1 / beta[2::3] * varC[0::3]))
    d1 = abs((np.pi * beta[1::3]**2 * k1 / 4 / beta[2::3]**3)) * varC[1::3]
    d2 = abs((k2 / 4 / beta[2::3]**2 / k3 - k2 * beta[1::3]**2 * k3 / 4 / beta[2::3]**2))**2 * varC[2::3]
    darea = (d0 + d1 + d2)**.5
    # DQS
    dqsArea = special.erfc(darea / area)

    ### DELETE PEAKS WITH Beta_2 > 0
    dqsArea[beta[2::3]>=0] *= 0
    _filter_delete = dqsArea>0
    x0, height, width, area, dqsArea = x0[_filter_delete], height[_filter_delete],  width[_filter_delete], area[_filter_delete], dqsArea[_filter_delete]

    ### EXPORT DataFrame
    df = pd.DataFrame()
    df["Centroids"] = x0
    df["Peak width [sigma]"] = width
    df["Peak height"] = height
    df["Peak area"] = area
    df["DQS"] = dqsArea
    df["RT [s]"] = t * np.ones(len(x0))
    df["scans"] = k * np.ones(len(x0))
    return(df)

def peak_filter(x_raw, y_raw):
    ydiff = np.diff(y_raw)
    _filter = peak_indices(ydiff)
    x_raw, y_raw, _filter = split_peaks(x_raw, y_raw, _filter)
    binsizes = np.bincount(_filter)
    bins = np.unique(_filter)
    _bins2zero = bins[binsizes <4]
    if len(_bins2zero) > 0:
        for _bin2zero in _bins2zero:
            _filter[_filter==_bin2zero] = 0
    ydiff = np.diff(_filter)
    _filter = peak_indices(ydiff)
    return(x_raw, y_raw, _filter)

def peak_indices(ydiff):
    yDiff = np.zeros(len(ydiff)+1)
    yDiff1 = np.append(0, np.sign(ydiff)==1)
    yDiff2 = np.append(0, yDiff1[0:-1])
    yDiff = np.float64((yDiff1 - yDiff2) == 1)

    yDiffNeg = np.append(0, np.sign(ydiff)==-1)
    yDiffNeg = np.append(np.diff(yDiffNeg),0)
    yDiffNeg = yDiffNeg == -1
    yDiff -= yDiffNeg

    yDiff[yDiff == 1] = np.cumsum(yDiff[yDiff == 1])
    yDiff[yDiff == -1] = np.cumsum(yDiff[yDiff == -1])
    _filter = np.cumsum(yDiff).astype("int")
    return _filter

def split_peaks(x_raw, y_raw, _filter):
    addPtsIdx = (y_raw != 0) * (_filter == 0)
    addPtsIdx = np.where(addPtsIdx)[0]
    if len(addPtsIdx) > 0:
        rawValuetoAdd_y = y_raw[addPtsIdx]
        rawValuetoAdd_x = x_raw[addPtsIdx]
        y_raw = np.insert(y_raw,np.append(addPtsIdx,addPtsIdx),np.append(rawValuetoAdd_y,rawValuetoAdd_y*0))
        x_raw = np.insert(x_raw,np.append(addPtsIdx,addPtsIdx),np.append(rawValuetoAdd_x,rawValuetoAdd_x))
        _filter = np.insert(_filter,np.append(addPtsIdx,addPtsIdx+1),np.append(_filter[addPtsIdx-1],_filter[addPtsIdx+1]))
    return(x_raw, y_raw, _filter)

def inverse3x3(beta):
    D0 = beta.diagonal()
    D1 = beta.diagonal(1)
    D2 = beta.diagonal(2)
    D1 = np.append(D1,0)
    D2 = np.append(D2,[0,0])
    
    M00 = D0[1::3] * D0[2::3] - D1[1::3]**2
    M10 = D1[0::3] * D0[2::3] - D1[1::3] * D2[0::3]
    M11 = D0[0::3] * D0[2::3] - D2[0::3]**2
    M20 = D1[0::3] * D1[1::3] - D0[1::3] * D2[0::3]
    M21 = D0[0::3] * D1[1::3] - D1[0::3] * D2[0::3]
    M22 = D0[0::3] * D0[1::3] - D1[0::3]**2
    M10 *= -1
    M21 *= -1
    DET = D0[0::3] * M00 + D1[0::3] * M10 + D2[0::3] * M20
    i00 = M00/DET
    i10 = M10/DET
    i20 = M20/DET
    i21 = M21/DET
    i11 = M11/DET
    i22 = M22/DET
    nRowsInv = np.shape(beta)[0]
    idxInv = np.arange(nRowsInv,dtype=int)
    M = i00
    M = np.append(M, i11)
    M = np.append(M, i22)
    M = np.append(M, i10)
    M = np.append(M, i10)
    M = np.append(M, i20)
    M = np.append(M, i20)
    M = np.append(M, i21)
    M = np.append(M, i21)
    IDX1 = idxInv[0::3]
    IDX1 = np.append(IDX1, idxInv[1::3])
    IDX1 = np.append(IDX1, idxInv[2::3])
    IDX1 = np.append(IDX1, idxInv[1::3])
    IDX1 = np.append(IDX1, idxInv[0::3])
    IDX1 = np.append(IDX1, idxInv[2::3])
    IDX1 = np.append(IDX1, idxInv[0::3])
    IDX1 = np.append(IDX1, idxInv[2::3])
    IDX1 = np.append(IDX1, idxInv[1::3])
    IDX2 = idxInv[0::3]
    IDX2 = np.append(IDX2, idxInv[1::3])
    IDX2 = np.append(IDX2, idxInv[2::3])
    IDX2 = np.append(IDX2, idxInv[0::3])
    IDX2 = np.append(IDX2, idxInv[1::3])
    IDX2 = np.append(IDX2, idxInv[0::3])
    IDX2 = np.append(IDX2, idxInv[2::3])
    IDX2 = np.append(IDX2, idxInv[1::3])
    IDX2 = np.append(IDX2, idxInv[2::3])
    inverse = csr_matrix((M,(IDX1,IDX2)))
    return inverse