module qCentroids
using SparseArrays
using DataStructures
using DataFrames
using SpecialFunctions
using Peaks
using Statistics

function fastSplice!(X::Vector{Float64},idx::Vector{Int64})
    if idx[1] != 1
        prepend!(idx,1)
    end
    if idx[end] != length(X)
        append!(idx,length(X))
    end   
    append!(X,zeros((length(idx)-1)*2-1))
    X .= reduce(vcat,[vcat(X[u:v],0) for (u,v) in zip(idx[1:end-1],idx[2:end])])
    nothing
end

function peak_indices(ydiff)
    yDiff = zeros(length(ydiff)+1)
    yDiff1 = prepend!(map(u->sign(u)==1, ydiff),0)
    yDiff2 = prepend!(yDiff1[1:end-1],0)
    yDiff = map(u->u==1 , yDiff1 - yDiff2)

    yDiffNeg = map(u->sign(u)==-1, ydiff)
    yDiffNeg = prepend!(yDiffNeg,0)
    yDiffNeg = append!(diff(yDiffNeg),0)
    yDiffNeg = map(u->u==-1,yDiffNeg)
    yDiff -= yDiffNeg

    yDiff[map(u->u==1, yDiff)] = cumsum(yDiff[map(u->u==1, yDiff)])
    yDiff[map(u->u==-1, yDiff)] = cumsum(yDiff[map(u->u==-1, yDiff)])
    _filter = cumsum(yDiff)
    _filter = round.(Int, _filter)
    return(_filter)
end

function peak_indices2(y_raw::Vector{Float64})::Vector{Int64}
    y_signs = prepend!(diff(sign.(y_raw)),0)
    y_signs[y_signs .== 1] .= cumsum(y_signs[y_signs .== 1])
    y_signs[y_signs .== -1] .= cumsum(y_signs[y_signs .== -1])
    cumsum(y_signs) 
end

function split_peaks!(x_raw::Vector{Float64}, 
            y_raw::Vector{Float64}, 
            _filter::Vector{Int64},
            model::String)

    if model == "gaussian"        
        addPtsIdx = (y_raw .!= 0) .* (_filter .== 0)
        addPtsIdx = findall(addPtsIdx)
    end

    if model == "bigaussian"
        addPtsIdx, vals = findmaxima(y_raw)
        _,proms = peakproms!(addPtsIdx,y_raw)
        _f = (proms ./ vals) .< .5
        deleteat!(addPtsIdx,_f)
        deleteat!(proms,_f)
        _, widths, _, _ = peakwidths!(addPtsIdx, y_raw, proms)
        _f = widths .< 1
        deleteat!(addPtsIdx,_f)
        re = [FindPeakEdge(y_raw,u,1) for u in addPtsIdx]
        le = [FindPeakEdge(y_raw,u,-1) for u in addPtsIdx]
        _f = ((addPtsIdx .- le) .< 2) .| ((re .- addPtsIdx) .< 2)
        deleteat!(addPtsIdx,_f)
    end

    if !isempty(addPtsIdx)
        fastSplice!(y_raw,addPtsIdx)
        fastSplice!(x_raw,addPtsIdx)
        append!(_filter, zeros((length(addPtsIdx)-1)*2-1))
        _filter .= peak_indices2(y_raw)            
    end
    nothing
end

function peak_filter!(_filter::Vector{Int64},
            x_raw::Vector{Float64},
            y_raw::Vector{Float64},
            model::String)
    if model=="gaussian"
        minPts = 4
    elseif model=="bigaussian"
        minPts = 3
        idxZeros = findall(y_raw .== 0)
        binsizes = append!(diff(idxZeros),length(y_raw)-idxZeros[end])
        [u <= 6 ? y_raw[U-u+1:U] .= 0 : nothing for (u,U) in zip(binsizes,cumsum(binsizes))]    
    end 
    _filter .= peak_indices2(y_raw)
    split_peaks!(x_raw, y_raw, _filter,model)

    binSizes = SortedDict(counter(_filter))
    bins,binsizes = collect(keys(binSizes)), collect(values(binSizes))
    _index2zero = findall(in(bins[map(u->u<minPts,binsizes)]),_filter)
    y_raw[_index2zero] .*= 0
    _filter .= peak_indices2(y_raw)
    nothing
end

function inverse2x2(beta)
    D0 = map(u->beta[u,u],(1:size(beta)[1]))
    D1 = map(u->beta[u+1,u],(1:size(beta)[1]-1))
    append!(D1, 0)
    ad_bc = D0[1:2:end] .* D0[2:2:end] .- D1[1:2:end].^2
    inverse = copy(beta)
    [inverse[u,u] = v for (u,v) in zip(1:2:length(D0),D0[2:2:end])]
    [inverse[u,u] = v for (u,v) in zip(2:2:length(D0),D0[1:2:end])]
    [inverse[u:u+1,u:u+1] ./= v for (u,v) in zip(1:2:length(D0),ad_bc)]
    [inverse[u+1,u] = -1* inverse[u+1,u] for u in 1:2:length(D0)]
    [inverse[u,u+1] = -1* inverse[u,u+1] for u in 1:2:length(D0)]
    return(inverse)
end

function inverse3x3(beta)
    D0 = map(u->beta[u,u],(1:size(beta)[1]))
    D1 = map(u->beta[u+1,u],(1:size(beta)[1]-1))
    D2 = map(u->beta[u+2,u],(1:size(beta)[1]-2))
    append!(D1, 0)
    append!(D2, [0;0])

    M00 = D0[2:3:end] .* D0[3:3:end] - D1[2:3:end].^2
    M10 = D1[1:3:end] .* D0[3:3:end] - D1[2:3:end] .* D2[1:3:end]
    M11 = D0[1:3:end] .* D0[3:3:end] - D2[1:3:end].^2
    M20 = D1[1:3:end] .* D1[2:3:end] - D0[2:3:end] .* D2[1:3:end]
    M21 = D0[1:3:end] .* D1[2:3:end] - D1[1:3:end] .* D2[1:3:end]
    M22 = D0[1:3:end] .* D0[2:3:end] - D1[1:3:end].^2
    M10 *= -1
    M21 *= -1
    DET = D0[1:3:end] .* M00 + D1[1:3:end] .* M10 + D2[1:3:end] .* M20

    i00 = M00./DET
    i10 = M10./DET
    i20 = M20./DET
    i21 = M21./DET
    i11 = M11./DET
    i22 = M22./DET

    nRowsInv = size(beta)[1]
    idxInv = round.(Int,cumsum(ones(nRowsInv)))

    M = i00
    append!(M, i11)
    append!(M, i22)
    append!(M, i10)
    append!(M, i10)
    append!(M, i20)
    append!(M, i20)
    append!(M, i21)
    append!(M, i21)

    IDX1 = idxInv[1:3:end]
    append!(IDX1, idxInv[2:3:end])
    append!(IDX1, idxInv[3:3:end])
    append!(IDX1, idxInv[2:3:end])
    append!(IDX1, idxInv[1:3:end])
    append!(IDX1, idxInv[3:3:end])
    append!(IDX1, idxInv[1:3:end])
    append!(IDX1, idxInv[3:3:end])
    append!(IDX1, idxInv[2:3:end])
    IDX2 = idxInv[1:3:end]
    append!(IDX2, idxInv[2:3:end])
    append!(IDX2, idxInv[3:3:end])
    append!(IDX2, idxInv[1:3:end])
    append!(IDX2, idxInv[2:3:end])
    append!(IDX2, idxInv[1:3:end])
    append!(IDX2, idxInv[3:3:end])
    append!(IDX2, idxInv[2:3:end])
    append!(IDX2, idxInv[3:3:end])
    inverse = sparse(IDX1,IDX2,M)
    return(inverse)
end

function createVMatrix(
            rowIndex::Vector{Int64},
            columnIndex::Vector{Int64},
            X::Vector{Float64},
            Y::Vector{Float64},
            model::String,
        )
    x_sparse = sparse(rowIndex,columnIndex,X)
    y_sparse = sparse(rowIndex,columnIndex,Y)
    ones_sparse = sparse(rowIndex,columnIndex, X.^0)

    if model == "gaussian"
        x_sparse_centroids = (x_sparse' * y_sparse) / (ones_sparse' * y_sparse)
        x_centroids = nonzeros(x_sparse_centroids)
        wf = 2
        weight_sparse = sparse(rowIndex,columnIndex, Y.^wf)
    elseif model == "bigaussian"
        max_IDX = argmax(y_sparse,dims=1)
        x_centroids = [x_sparse[u] for u in max_IDX]
        wf = 4
        weight_sparse = sparse(rowIndex,columnIndex, Y.^wf)
    end
    numberOfnnz = round.(Int, nonzeros(ones_sparse' * ones_sparse))

    x_centroids_repeated = []
    weight_sparse_means_repeated = []
    weight_sparse_means = nonzeros(weight_sparse'*weight_sparse / (ones_sparse' * ones_sparse))
    for u in (1:length(numberOfnnz))
        append!(x_centroids_repeated,x_centroids[u] * ones(numberOfnnz[u]))
        append!(weight_sparse_means_repeated,weight_sparse_means[u] * ones(numberOfnnz[u]))
    end
    x = X-x_centroids_repeated
    weight_sparse = sparse(rowIndex,rowIndex, Y.^wf ./ weight_sparse_means_repeated)
    # # Linear Regression
    # # Creating Vandermonde Matrix
    if model == "gaussian"
        xV = [ones(length(x)); x; x.^2]
        rI = [rowIndex; rowIndex; rowIndex]
        cI = [round.(Int,columnIndex*3-ones(length(x))*2); round.(Int,columnIndex*3-ones(length(x))); columnIndex*3]
    elseif model == "bigaussian"
        xV = [ones(length(x)); x.^2]
        rI = [rowIndex; rowIndex]
        cI = [round.(Int,columnIndex*2-ones(length(x))); columnIndex*2]
        x_centroids = x_centroids'
    end
    xV = sparse(rI,cI,xV)
    
    return(xV,weight_sparse,x_centroids,numberOfnnz)
end

function findCentroids(x::Vector{Float64},
            y::Vector{Float64},
            t=1. ::Float64, 
            k=1. ::Float64, 
            model="gaussian" ::String)
    _filter = zeros(Int64,length(x))
    peak_filter!(_filter,x,y,model)
    # # Create Sparse Matrix
    _filterData = _filter .> 0
    columnIndex = _filter[_filterData]
    rowIndex = collect(1:length(columnIndex))

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

    # # # Linear Regression
    # # # Creating Vandermonde Matrix
    xV,weight_sparse,x_centroids,numberOfnnz = createVMatrix(rowIndex,columnIndex,X,Y,model)
    # # Calculate Coefficients
    beta = xV' * weight_sparse * xV
    if model=="gaussian"
        inverse = inverse3x3(beta)
        ylog = map(u->log(u), Y)
    elseif model=="bigaussian"
        inverse = inverse2x2(beta)
        ylog = log.(Y)
    end
    beta = inverse * xV' * weight_sparse * ylog
    # # Transform Linear Coefficients into Gaussian Coefficients
    if model == "gaussian"
        b0 = beta[1:3:end]
        b1 = beta[2:3:end]
        b2 = beta[3:3:end]
        num_param = 3
        x0 = b1 / -2 ./ b2 .+ x_centroids
        width = (map(u->abs(u), (1 / 2 ./ b2)).^.5)
        height = map(u->exp(u), b0 - b1.^2 / 4 ./ b2)
        area = height .* width * (2pi).^.5
    elseif model == "bigaussian"
        b0 = beta[1:2:end]
        b2 = beta[2:2:end]
        num_param = 2
        x0 = x_centroids
        height = exp.(b0)
        width = sqrt.(abs.(.- 1 ./ 2 ./ b2))
        area = height .* width * (2pi).^.5
    end
    
    # # Error Propagation
    yfit = xV*beta
    residues = yfit - ylog
    residues = sparse(rowIndex, columnIndex, residues)
    RSS = residues' * weight_sparse * residues
    df  = numberOfnnz-ones(length(numberOfnnz))*num_param
    MSE = nonzeros(RSS) ./ df
    varC = map(u->inverse[u,u],(1:size(beta)[1])) .* reshape(repeat(MSE,1,num_param)',length(numberOfnnz)*num_param,1)
    if model == "gaussian"
        vC0 = varC[1:3:end]
        vC1 = varC[2:3:end]
        vC2 = varC[3:3:end]
        k1 = map(u->exp(u), 2*b0 - b1.^2 / 2 ./ b2)
        k2 = (2pi).^.5 * map(u->exp(u), b0-b1.^2 / 4 ./b2)
        k3 = ((1 / 2 ./ map(u->abs(u), (b2))).^.5)
        d0 = map(u->abs(u), (pi * k1 ./ b2 .* vC0))
        d1 = map(u->abs(u), (pi * b1.^2 .* k1 / 4 ./ b2.^3)) .* vC1
        d2 = map(u->abs(u), (k2 / 4 ./ b2.^2 ./ k3 - k2 .* b1.^2 .* k3 / 4 ./ b2.^2)).^2 .* vC2
        darea = (d0+ d1+ d2).^.5
    elseif model == "bigaussian"
        vC0 = varC[1:2:end]
        vC2 = varC[2:2:end]
        # formula: dArea = sqrt( |dArea/dheight|^2 * |dheight|^2 + |dArea/dwidth|^2 * |dwidth|^2 )
        dA_dh = width .* sqrt(2pi)
        dA_dw = height .* sqrt(2pi)
        dh = exp.(b0) .* sqrt.(vC0)
        dw = 1 ./ 2 ./ sqrt(2) ./ sqrt.(abs.(.- 1 ./ b2)) ./ b2.^2 .* sqrt.(vC2)
        darea = sqrt.(dA_dh.^2 .* dh.^2 .+ dA_dw.^2 .*dw.^2)
    end
    
    if model == "gaussian"
        # # DQS
        dqsArea = map(u->erfc(u), darea ./ area)
        # ### DELETE PEAKS WITH Beta_2 > 0
        dqsArea[map(u->u>=0,b2)] .*= 0
        _filter_delete = map(u->u>0, dqsArea)
        x0, height, width, area, dqsArea = x0[_filter_delete], height[_filter_delete],  width[_filter_delete], area[_filter_delete], dqsArea[_filter_delete]
        asymmetry = ones(length(x0))
    elseif model == "bigaussian"
        binSizes = SortedDict(counter(x0))
        bins,binsizes = collect(keys(binSizes)), collect(values(binSizes))
        _filter_delete = reduce(vcat, [u == 2 ? trues(2) : falses(1) for u in binsizes])
        x0, height, width, area, darea = x0[_filter_delete], height[_filter_delete],  width[_filter_delete], area[_filter_delete], darea[_filter_delete]

        asymmetry = [width[x0 .== u][1] / width[x0 .== u][2] for u in unique(x0)]
        width = [mean(width[x0 .== u]) for u in unique(x0)]
        height = [mean(height[x0 .== u]) for u in unique(x0)]
        area = [sum(area[x0 .== u]) for u in unique(x0)]
        darea = [sum(darea[x0 .== u]) for u in unique(x0)]
        dqsArea = map(u->erfc(u), darea ./ area)
        x0 = [mean(x0[x0 .== u]) for u in unique(x0)]

        _filter_delete = map(u->u>0, dqsArea)
        x0, height, width, area, dqsArea, asymmetry = x0[_filter_delete], height[_filter_delete],  width[_filter_delete], area[_filter_delete], dqsArea[_filter_delete], asymmetry[_filter_delete]
    end

    # ### EXPORT DataFrame
    df = DataFrame()
    df[!, "Centroids"] = x0
    df[!, "Peak width [sigma]"] = width
    df[!, "Peak asymmetry"] = asymmetry
    df[!, "Peak height"] = height
    df[!, "Peak area"] = area
    df[!, "DQS"] = dqsArea
    df[!, "RT [s]"] = t * ones(length(x0))
    df[!, "Scans"] = k * ones(length(x0))
    return(df)
end

function FindPeakEdge(x::Vector{Float64},x0_pos::Int64,dir::Int64)::Int64
    u = 0
    while true
        if x[x0_pos + u] < (x[x0_pos]/10)
            return(x0_pos + u)
        else
            u += dir
            if ((x0_pos+u)>length(x)) | ((x0_pos+u)<1)
                return(x0_pos + u-dir)
            end
        end  
    end
end

end # module
