#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Some processing functions

"""

import os
import sys
import csv
import copy
import math
import numpy as np
import netCDF4 as nc
import scipy.stats as stats
import matplotlib.pyplot as plt
import shapefile
import statsmodels.api as sm
import datetime
from shapely.geometry import Point,Polygon,MultiPoint,LineString
from mpl_toolkits.basemap import Basemap
from matplotlib import cm
from matplotlib.colors import Normalize
from scipy.interpolate import interpn
from scipy.interpolate import griddata

### Some operators ###
lg_a = np.logical_and
lg_o = np.logical_or
######################



def prepare_for_2d_pdf(x,y,xbins,ybins):
    '''
    Source: stackoverflow.com/questions/31430390/return-the-value-of-a-2d-pdf-given-x-and-y-in-python
    Preparing data for 2D pdf count scatterplot
    Returns [x,y,counts]
    Counts is number of data points falling within each 2d bin
    Code for plotting the output:
        ax.contourf(x,y,counts,levels=mylevels,cmap=mycmap,extend='max',norm=LogNorm())
        or
        ax.contourf(x,y,counts,levels=mylevels,cmap=mycmap,extend='max')
    '''
    counts_transp,xxedges,yyedges = np.histogram2d(x,y,bins=[xbins,ybins])
    counts = counts_transp.T
    xxcents = xxedges[0:-1]+np.diff(xxedges)        
    yycents = yyedges[0:-1]+np.diff(yyedges)  
    xxpl,yypl = np.meshgrid(xxcents,yycents)    
    return(xxpl,yypl,counts)

def getlandoceanmask(grid_lat,grid_lon,fmask,caspian_is_ocean=True,is_binary=True,super_high_res=False):
    '''
    Computes land (1) and ocean (0) mask over grid
    Uses the high-resolution land-ocean mask from TerraClimate soil data
    Option to set Caspian sea as ocean (default) or not
    '''
    print('Retrieving Land/Ocean mask')
    # Original resolution is 0.04 degree: we can usually afford to sub-sample #
    if(super_high_res):
        subsmp = 1
    else:
        subsmp = 4
    # Load #
    ds = nc.Dataset(fmask)
    mklat0 = np.array(ds['lat'][::subsmp,::subsmp])
    mklon0 = np.array(ds['lon'][::subsmp,::subsmp])
    mask0  = np.array(ds['mask'][::subsmp,::subsmp])
    ds.close()
    ### Put Caspian sea as ocean ###
    if(caspian_is_ocean):
        lscasp =[(51.92,36.42),(49.26,37.44),(48.82,38.68),(49.57,40.33),(47.80,42.50),(46.78,44.53),(48.11,46.30),(51.12,47.02),
                 (53.12,46.88),(53.20,45.33),(50.99,44.27),(53.78,42.23),(54.93,41.03),(54.09,39.93),(54.09,36.65)]
        polycasp    = Polygon(lscasp)
        iscasp      = np.zeros_like(mklat0).astype(bool)
        inds_search = lg_a(lg_a(lg_a(mklon0>=np.min(lscasp),mklon0<=np.max(lscasp)),
                                mklat0>=np.min(lscasp)),mklat0<=np.max(lscasp))
        i_inds,j_inds = np.where(inds_search)
        for kk in range(len(i_inds)):
            if(polycasp.contains(Point(mklon0[i_inds[kk],j_inds[kk]],mklat0[i_inds[kk],j_inds[kk]]))):
                iscasp[i_inds[kk],j_inds[kk]] = True
        mask0[iscasp] = 0
    # Interpolate #
    if(is_binary):
        minterp = 'nearest'
    else:
        minterp = 'linear'
    mask = griddata((mklat0.flatten(),mklon0.flatten()),mask0.flatten(),(grid_lat,grid_lon),method=minterp)
    # Return #
    print('Land(1)/Ocean(0) mask retrieved')
    return(mask)


def dist_haversine(latpoint,lonpoint,lats,lons):
    '''
    Computes Haversine formula for distance between a point and an array of points
    see en.wikipedia.org/wiki/Haversine_formula
    '''
    e_rad = 6371008.8 #en.wikipedia.org/wiki/Earth_radius
    latpoint = latpoint*2*np.pi/360
    lonpoint = lonpoint*2*np.pi/360
    lats     = lats*2*np.pi/360
    lons     = lons*2*np.pi/360
    dists = 2*e_rad*np.arcsin(np.sqrt(
        (np.sin((lats-latpoint)/2))**2+
        (np.sin((lons-lonpoint)/2))**2*
            (1 - (np.sin((lats-latpoint)/2))**2 - (np.sin((lats+latpoint)/2))**2)
        ))
    return(dists)

def scaling_minmax01(indata):
    '''
    Normalizes data in [0,1] range via (x-min)/(max-min)
    '''
    inmin = np.min(indata)
    inmax = np.max(indata)
    if(inmax>inmin):
        outdata = (indata-inmin)/(inmax-inmin)
    else:
        print('scaling_minmax01 received constant data: returning nans')
        outdata = np.nan*np.ones_like(indata)
    return(outdata,inmin,inmax)


def linear_trend_test(timearray,vals,one_sided=False):
    '''
    Performs test for presence of linear trend
    '''
    
    ntot     = len(vals)
    xrg      = np.column_stack((np.ones(ntot),timearray))
    yrg      = np.reshape(vals,(ntot,1))
    rgmod    = sm.OLS(yrg,xrg)
    rgmodres = rgmod.fit()
    
    trendval = rgmodres.params[1] #trend value
    t_stat   = rgmodres.tvalues[1] #t-statistic of the linear coefficient
    pval     = rgmodres.pvalues[1] #p-value of the t-statistic
    if(one_sided):
        pval = pval/2
    return(trendval,pval,t_stat)


            
def lat_lon_cell_area(lats,lons):
    """
    Calculate the area of a cell, in meters^2, on a lat/lon grid.
    This applies the following equation from Santini et al. 2010.
    S = (λ_2 - λ_1)(sinφ_2 - sinφ_1)R^2
    S = surface area of cell on sphere
    λ_1, λ_2, = bands of longitude in radians
    φ_1, φ_2 = bands of latitude in radians
    R = radius of the sphere
    """
    AVG_EARTH_RADIUS_METERS = 6371008.8 #en.wikipedia.org/wiki/Earth_radius#Mean_radius
    
    south,north = lats[:]
    west,east   = lons[:]

    west  = math.radians(west)
    east  = math.radians(east)
    south = math.radians(south)
    north = math.radians(north)
    
    area = (east-west)*(math.sin(north)-math.sin(south))*(AVG_EARTH_RADIUS_METERS**2) 
    return(area)


def averaging_daily_to_monthly(dayarray,vals,axis_time=None,disp=True):
    '''
    Averages values provided at daily time steps to monthly time steps
    '''
    
    ndim   = len(np.shape(vals))
    msteps = np.arange(1/24,23/24+1e-5,1/12)
    m00    = np.argmin(abs(msteps+np.floor(dayarray[0])-dayarray[0])) #month index of first time step
    m11    = np.argmin(abs(msteps+np.floor(dayarray[-1])-dayarray[-1])) #month index of last time step
    tm_out  = np.arange(np.floor(dayarray[0])+msteps[m00],np.floor(dayarray[-1])+msteps[m11]+1e-10,1/12) #monthly time
    nmonths = len(tm_out)
    
    # Move time axis at index 0 #
    if(ndim>=2):
        if(axis_time is None):
            print('Error: provide axis_time in averaging_daily_to_monthly for multi-dimensional arrays')
            sys.exit()
        mv_vals = np.moveaxis(vals,axis_time,0)
    else:
        mv_vals = np.copy(vals)
    
    if(ndim==1):
        avg_vals = np.nan*np.ones(nmonths)
    elif(ndim==2):
        avg_vals = np.nan*np.ones((nmonths,np.shape(mv_vals)[1]))
    elif(ndim==3):
        avg_vals = np.nan*np.ones((nmonths,np.shape(mv_vals)[1],np.shape(mv_vals)[2]))
    else:
        print('averaging_daily_to_monthly not implemented for dimension>2')
        sys.exit()
        
    # Monthly averaging #
    for ii,iitm in enumerate(tm_out):
        inds = abs(dayarray-iitm)<=1/24
        if(np.any(inds)):
            avg_vals[ii,] = np.mean(mv_vals[inds,],axis=0)
        else:
            if(disp):
                print('month without data in averaging_daily_to_monthly: returning nan')
    
    # Re-move to original time axis #
    if(ndim>=2):
        avg_vals = np.moveaxis(avg_vals,0,axis_time)
    
    return(tm_out,avg_vals)

def corrcoef_3d_vectorized(vals0,vals1,axis,mask_in=None,get_2tailed_pval=True):
    '''
    Computes correlation coefficient between two 3d arrays along given axis
    mask_in is a 2-dimensional array to specify for which indices correlation should be computed
    Indices out of mask_in are assigned -9999 filling value 
    '''
    vcop0 = np.copy(vals0)
    vcop1 = np.copy(vals1)
    vcop0[:,mask_in==False] = -9999
    vcop1[:,mask_in==False] = -9999
    numer = np.nanmean(((vcop0-np.nanmean(vcop0,axis=axis))*(vcop1-np.nanmean(vcop1,axis=axis))),axis=axis)
    denom = 1e-10+np.nanstd(vcop0,axis=0)*np.nanstd(vcop1,axis=0)
    corr_out = numer/denom
    corr_out[mask_in==False] = -9999
    if(get_2tailed_pval==False):
        return(corr_out)
    elif(get_2tailed_pval):
        nsamples   = len(vals0)
        temp       = np.maximum(np.minimum(corr_out,0.99999),-0.99999)
        tstat      = corr_out*np.sqrt(nsamples-2)/np.sqrt(1-temp**2) #test statistic
        vcdf_abst  = stats.t.cdf(abs(tstat),df=nsamples-2) #cdf value of test statistic
        twotailedp = 2*(1-vcdf_abst) #two-tailed t-test
        return(corr_out,twotailedp)

def smooth_movingaverage_time(timer,vals,window,axis_time=None):
    
    ntot   = len(timer) #number of time steps
    ndim   = len(np.shape(vals))    
    # Move time axis at index 0 #
    if(ndim>=2):
        if(axis_time is None):
            print('Error: provide axis_time in smooth_movingaverage_time for multi-dimensional arrays')
            sys.exit()
        mv_vals = np.moveaxis(vals,axis_time,0)
    else:
        mv_vals = np.copy(vals)
    
    # Smoothing #
    outvals = np.zeros_like(mv_vals)
    for tt in range(ntot):
        inds = abs(timer-timer[tt])<=(window/2)
        outvals[tt] = np.nanmean(mv_vals[inds,],axis=0)
    
    # Re-move to original time axis #
    if(ndim>=2):
        outvals = np.moveaxis(outvals,0,axis_time)
    
    return(outvals)


def smooth_movingaverage_in2d(vals,windowy,windowx,indy,indx,grid2dy,grid2dx,connminmax_y=None,connminmax_x=None):
    '''
    Performs moving average in 2 dimensions
    Along y- and x-directions
    indy: index of the vals array along y-direction
    indx: index of the vals array along y-direction
    coordinates grids passed (grid2dy,grid2dx) must be 2-dimensional
    connminmax_ can be used to allow connection between lower- and upper-end:
        connminmax_ = [lower_edge,upper_edge]
    '''
    ndim = len(np.shape(vals))  
    if(indy!=0):
        mv_vals = np.swapaxes(vals,indy,0)
    if(indx!=1 and ((indy==1 and indx==0)==False)):
        mv_vals = np.swapaxes(mv_vals,indx,1)
    if(ndim==2):
        ntot = 1
    elif(ndim==3):
        ntot = np.shape(mv_vals)[2]
    else:
        print('Error: smooth_movingaverage_in2d is not implemented for more than 3 dimensional arrays')
        return(np.isnan)
    
    # Compute spatial moving average #
    nny = np.shape(mv_vals)[0]
    nnx = np.shape(mv_vals)[1]
    if(np.shape(grid2dy)[0]!=nny and np.shape(grid2dy)[0]==nnx):
        grid2dy = np.moveaxis(grid2dy,1,0)
        grid2dx = np.moveaxis(grid2dx,1,0)
    outvals = np.nan*np.ones_like(mv_vals)
    for ii in range(nny):
        #print(f'smoothing y-band: {ii+1}/{nny}')
        for jj in range(nnx):
            # Distances in y-direction #
            distsyy = np.sqrt((grid2dy[ii,jj]-grid2dy)**2)
            if(connminmax_y is not None):
                # Close to lower edge and connection to upper edge #
                if(grid2dy[ii,jj]<connminmax_y[0]+windowy/2):
                    conninds = grid2dy>connminmax_y[1]+(grid2dy[ii,jj]-connminmax_y[0]-windowy/2)
                    distsyy[conninds] = np.sqrt((connminmax_y[1]-grid2dy[conninds]+grid2dy[ii,jj]-connminmax_y[0])**2)
                # Close to upper edge and connection to lower edge #
                if(grid2dy[ii,jj]>connminmax_y[1]-windowy/2):
                    conninds = grid2dy<connminmax_y[0]+(connminmax_y[1]-grid2dy[ii,jj]+windowy/2)
                    distsyy[conninds] = np.sqrt((grid2dy[conninds]-connminmax_y[0]+connminmax_y[1]-grid2dy[ii,jj])**2)
            # Distances in x-direction #
            distsxx = np.sqrt((grid2dx[ii,jj]-grid2dx)**2)
            if(connminmax_x is not None):
                # Close to lower edge and connection to upper edge #
                if(grid2dx[ii,jj]<connminmax_x[0]+windowx/2):
                    conninds = grid2dx>connminmax_x[1]+(grid2dx[ii,jj]-connminmax_x[0]-windowx/2)
                    distsxx[conninds] = np.sqrt((connminmax_x[1]-grid2dx[conninds]+grid2dx[ii,jj]-connminmax_x[0])**2)
                # Close to upper edge and connection to lower edge #
                if(grid2dx[ii,jj]>connminmax_x[1]-windowx/2):
                    conninds = grid2dx<connminmax_x[0]+(connminmax_x[1]-grid2dx[ii,jj]+windowx/2)
                    distsxx[conninds] = np.sqrt((grid2dx[conninds]-connminmax_x[0]+connminmax_x[1]-grid2dx[ii,jj])**2)
            # Indices within window #
            inds = lg_a(distsyy<=windowy/2,distsxx<=windowx/2)
            # Average #
            if(ntot==1):
                outvals[ii,jj] = np.nanmean(mv_vals[inds])
            elif(ntot>1):
                outvals[ii,jj,:] = np.nanmean(mv_vals[inds,:],axis=0)

    # Re-move to original axes #
    outvals = np.swapaxes(outvals,0,indy)
    if(indy!=1):
        outvals = np.swapaxes(outvals,1,indx)
    else:
        outvals = np.swapaxes(outvals,0,indx)
    return(outvals)



