#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Tue Jun 15 17:09:57 2021

@author: Fiona Fix
this script contains functions used for the analysis in our paper
"""
from netCDF4 import Dataset, num2date
import numpy as np
from eofs.standard import Eof
import matplotlib.pyplot as plt
import plotting as plot
#--------------------------------------------------------
#%%     
def read_ncfile(infile, var_name):
    f = Dataset(infile, 'r')
    data = f.variables[var_name][:]
    return data

def get_pcs(data, npcs, pcscaling):
    solver = Eof(data)
    pcs    = solver.pcs(npcs=npcs, pcscaling = pcscaling) 
    return pcs
def get_eofs_and_explvar(data, npcs, eofscaling, Covariance):
    """
    gets Empirical Orthogobal Fuctions and Explaines Variances
    data (3D array)     : input data, eg SSTs
    npcs (integer)      : number of EOFs to be caculated
    eofscaling(integer) : set scaling option for eofs, not relevant if Covariance is True
    Covariance(boolean) : decide wheter you want eofs scaled or expressed as covariance
    """
    solver = Eof(data)
    if Covariance == True:
        eofs = solver.eofsAsCovariance(neofs=npcs, pcscaling =1)
    else: 
        eofs = solver.eofs(eofscaling=eofscaling, neofs=npcs)
        
    explVariances = solver.varianceFraction(neigs = npcs)
    return eofs, explVariances

def get_timeseries_occurrences(data, window_size, 
                               threshold, lengthcrit_events):
    """
    gets the timeseries of occurrences of El Nino/La NIna events
    data (1D array):       timeseries of eg SST annomalies
    window_size (int):     length of the mocing window
    lengthcrit_events(int): how many consecutive months needed for events
                           to be counted
    
    returns:               two 2D arrays: 
                           timeseries of occurrences of
                           nino events and conditions
                           nina events and conditions
    """
    list_ninos = []
    list_ninas = []
    for j in range((window_size//2),data.shape[0]-(window_size//2),1):
        window_data = data[j-(window_size//2):j+(window_size//2)]
        window_ninacond = count_condition(window_data, -threshold,  "smaller")
        window_ninaev   = count_events(window_data, -threshold, "smaller", lengthcrit_events)
        window_ninocond = count_condition(window_data,  threshold,  "greater")
        window_ninoev   = count_events(window_data,  threshold, "greater", lengthcrit_events)
        
        if lengthcrit_events==1:
            list_ninos.append(window_ninocond)
            list_ninas.append(window_ninacond)
        else:
            list_ninos.append(window_ninoev)
            list_ninas.append(window_ninaev)
    return list_ninos, list_ninas
        
def count_events(timeseries, threshold, rule, lengthcrit):
    """
    counts how often condition is met long enough for an event
    timeseries (array):     timeseries of eg SST annomalies
    threshold (float):      eg 0.5 for nino
    rule (string):          "greater" or "smaller"
    lengthcrit (int):       length of consecutive conditions 
                            to be counted an event
    returns:                number of events that meet condition 
                            for lengthcrit consecutive months
    """
    if rule=='greater':
        idx = np.where(timeseries>threshold)[0]
            
    elif rule=='smaller':
        idx = np.where(timeseries<threshold)[0]
    else:
        print('please state rule: "greater" or "smaller"')
    
    #number = len(idx)
    grouped     = np.split(idx, np.where(np.diff(idx) != 1)[0]+1)
    count_event = 0
    for g in grouped: 
        if len(g) >= lengthcrit:
            count_event = count_event+1
    return count_event

def count_condition(timeseries, threshold, rule):
    """
    counts how often condition is met 
    timeseries (array):     timeseries of eg SST annomalies
    threshold (float):      eg 0.5 for nino
    rule (string):          "greater" or "smaller"
    
    returns:                number of  months that meet condition 
    """   
    if rule=='greater':
        idx = np.where(timeseries>threshold)[0]
            
    elif rule=='smaller':
        idx = np.where(timeseries<threshold)[0]
    else:
        print('please state rule: "greater" or "smaller"')
    count_condition = len(idx)
    return count_condition

def calc_gradients(ENSO_counts):
    """
    calculates gradients of time series of occurences
    ENSO_counts (1Darray)   : time series
    
    returns: gradient and regression line
    """
    x        = np.arange(ENSO_counts.shape[0])
    coeffs   = np.ma.polyfit(x,ENSO_counts,1)
    return coeffs[0], np.poly1d(coeffs)(x)



def calc_stdsmeans(gradients, factor_changes):
    std = {key :[] for key in
            ['elnino', 'lanina', 'ninolike', 'ninalike'] }
    mean = {key :[] for key in
                    ['elnino', 'lanina', 'ninolike', 'ninalike'] }
    for k in ['elnino', 'lanina', 'ninolike', 'ninalike']:
        changes = np.array(gradients[k])*factor_changes
        std[k].append(np.std(changes))
        mean[k].append(np.mean(changes))
    return std, mean


#%%

def autocorr(x,lags):
    # lagged autocorrelation
    corr=[np.corrcoef(x[:],x[:])[0][1] if l==0 
                       else np.corrcoef(x[l:],x[:-l])[0][1]for l in lags]
    return np.asarray(corr)
#%%
def moving_average(x, w):
    return np.convolve(x, np.ones(w), 'valid') / w
