import os
os.environ['OPENBLAS_NUM_THREADS'] = '1'

import pandas as pd
import numpy as np

from RR import RRAlgotrithm

STAT_TIMESTEP = 1.0

def determine_zone(bpm, zones):
    if len(zones) == 3:
        if (bpm > 0) & (bpm < zones[0]):
            return "Zone 2 (Easy)"
        elif (bpm >= zones[0]) & (bpm < zones[1]):
            return "Zone 3 (Aerobic)"
        elif (bpm >= zones[1]) & (bpm < zones[2]):
            return "Zone 4 (Threshold)"
        elif (bpm >= zones[2]):
            return "Zone 5 (Maximum)"
            
    elif len(zones) == 4:
        if (bpm > 0) & (bpm < zones[0]):
            return "Zone 1 (Warm Up)"
        elif (bpm >= zones[0]) & (bpm < zones[1]):
            return "Zone 2 (Easy)"
        elif (bpm >= zones[1]) & (bpm < zones[2]):
            return "Zone 3 (Aerobic)"
        elif (bpm >= zones[2]) & (bpm < zones[3]):
            return "Zone 4 (Threshold)"
        elif (bpm >= zones[3]):
            return "Zone 5 (Maximum)"

        
def get_onratio(filename):
    df = pd.read_csv(filename, header=0).iloc[:-100]
    
    # define variable names
    colTime, colTempOral, colTempNasal = 'timeSeconds', 'tempOral', 'tempNasal'
    signalTime, signalNasal, signalOral = df[colTime].to_numpy(), df[colTempNasal].to_numpy(), df[colTempOral].to_numpy()
    
    OralRatio = estimateONratioCH02r1(signalTime, signalOral, signalNasal)
    return OralRatio
    
def get_onratio_ch02r0(filename):
    df = pd.read_csv(filename, header=0).iloc[:-100]
    
    # define variable names
    colTime, colTempOral, colTempNasal = 'timeSeconds', 'tempOral', 'tempNasal'
    signalTime, signalNasal, signalOral = df[colTime].to_numpy(), df[colTempNasal].to_numpy(), df[colTempOral].to_numpy()
    
    OralRatio = estimateONratioCH02r0(signalTime, signalOral, signalNasal)
    return OralRatio

def get_onratio_ch02r1(filename):
    df = pd.read_csv(filename, header=0).iloc[:-100]
    
    # define variable names
    colTime, colTempOral, colTempNasal = 'timeSeconds', 'tempOral', 'tempNasal'
    signalTime, signalNasal, signalOral = df[colTime].to_numpy(), df[colTempNasal].to_numpy(), df[colTempOral].to_numpy()
    
    OralRatio = estimateONratioCH02r1(signalTime, signalOral, signalNasal)
    return OralRatio
    
def estimateONratioCH02r0(signalTime, signalOral, signalNasal, N_window =30, alpha=0.75, flagPlot=False ):
    # Function for estimating the Oral/Nasal (O/N) breathing ratio
    # Assumes signal data has a sampling rate of 10 Hz
    #
    # Some definitions
    #   signalTime: np.array with time signal data in seconds
    #   signalOral: np.array with oral temperature signal in degrees Celcius
    #   signalNasal: np.array with nasal temperature signal in degrees Celcius
    #   N_window: size of the moving window used for DC and RMS calculations (default = 30, which implies a window of 3 seconds)
    #   alpha: reduction factor for Oral AC RMS signal to decide between prue oral and oral exp/nasal insp cases
    
    
    # compute AC / DC temperature signals
    N = len(signalTime)

    # Option 3 for DC: use moving average for DC signal
    signalNasalDC, signalOralDC = np.zeros(N), np.zeros(N)
    for i in range(N_window,N):
        signalOralDC[i] = np.mean(signalOral[i-N_window:i])
        signalNasalDC[i] = np.mean(signalNasal[i-N_window:i])
    
    # compute AC and plot AC & DC signals
    signalNasalAC = signalNasal - signalNasalDC
    signalOralAC = signalOral - signalOralDC
        
    # compute moving RMS continuous time for AC signal
    signalOralACrms = np.zeros(N)
    signalNasalACrms = np.zeros(N)
    
    if (N_window > 2*N): print('Error: Acquisition time is too short. Try longer activity times')
    
    for i in range(2*N_window,N):
        window_OralAC = signalOralAC[i-N_window:i]
        window_NasalAC = signalNasalAC[i-N_window:i]
        
        signalOralACrms[i] = np.sqrt(1/N_window*np.sum(window_OralAC**2))
        signalNasalACrms[i] =  np.sqrt(1/N_window*np.sum(window_NasalAC**2))    
    
    i_start = 2*N_window # needed to avoid spurious signal in the beginning
    
    # Classification algorithm
    
    signalOralSwitch,signalNasalSwitch = np.zeros(N), np.zeros(N)
    
    alpha = 0.850  # AC signal threshold in %
    
    for i in range(i_start,N):
        if (signalNasalDC[i] > signalOralDC[i]):    #Pure nasal breathing
            signalOralSwitch[i] = 0.0
            signalNasalSwitch[i] = 1.0
        else:
     
            if( signalNasalACrms[i] < alpha*signalOralACrms[i] ) : # Pure oral breathing 
                signalOralSwitch[i] = 1.0
                signalNasalSwitch[i] = 0.0
            else:   # Nasal Insp / Oral exp
                signalOralSwitch[i] = 0.0
                signalNasalSwitch[i] = 1.0           
            
    # pad switch signals with initial value
    signalOralSwitch[0:i_start] = signalOralSwitch[i_start]
    signalNasalSwitch[0:i_start] = signalNasalSwitch[i_start]
    
    # Compute Oral ratio
    OralRatio = np.mean(signalOralSwitch)
  
    return OralRatio
    
def estimateONratioCH02r1(signalTime, signalOral, signalNasal, N_window =30, alpha=5.0, flagPlot=False):
    # Function for estimating the Oral/Nasal (O/N) breathing ratio
    # Assumes signal data has a sampling rate of 10 Hz
    #
    # Some definitions
    #   signalTime: np.array with time signal data in seconds
    #   signalOral: np.array with oral temperature signal in degrees Celcius
    #   signalNasal: np.array with nasal temperature signal in degrees Celcius
    #   N_window: size of the moving window used for DC and RMS calculations (default = 30, which implies a window of 3 seconds)
    #   alpha: reduction factor for Oral AC RMS signal to decide between pure oral and oral exp/nasal insp cases
    
    
    # compute AC / DC temperature signals
    N = len(signalTime)
    
    if flagPlot:
        plt.plot(signalTime,signalOral)
        plt.plot(signalTime,signalNasal)
        plt.legend(["Oral","Nasal"])
        plt.xlabel("Time [s]")
        plt.ylabel("Temp [C]")
        plt.show()
        

    # Option 3 for DC: use moving average for DC signal 
    signalNasalDC, signalOralDC = np.zeros(N), np.zeros(N)
    for i in range(N_window,N):
        signalOralDC[i] = np.mean(signalOral[i-N_window:i])
        signalNasalDC[i] = np.mean(signalNasal[i-N_window:i])
    
    # compute AC and plot AC & DC signals
    
    signalNasalAC = signalNasal - signalNasalDC
    signalOralAC = signalOral - signalOralDC
        
    # compute moving RMS continuous time for AC signal
    signalOralACrms = np.zeros(N)
    signalNasalACrms = np.zeros(N)
    
    if (N_window > 2*N): print('Error: Acquisition time is too short. Try longer activity times')
    
    for i in range(2*N_window,N):
        window_OralAC = signalOralAC[i-N_window:i]
        window_NasalAC = signalNasalAC[i-N_window:i]
        
        signalOralACrms[i] = np.sqrt(1/N_window*np.sum(window_OralAC**2))
        signalNasalACrms[i] =  np.sqrt(1/N_window*np.sum(window_NasalAC**2))    
        
    i_start = 2*N_window # needed to avoid spurious signal in the beginning
    
    # Classification algorithm
    
    signalOralSwitch,signalNasalSwitch = np.zeros(N), np.zeros(N)
    
    for i in range(i_start,N):
        
        if( signalNasalACrms[i] < alpha*signalOralACrms[i] ) : # Pure oral breathing 
                    signalOralSwitch[i] = 1.0
                    signalNasalSwitch[i] = 0.0
        else:   # Nasal Insp / Oral exp
            signalOralSwitch[i] = 0.0
            signalNasalSwitch[i] = 1.0   

             
            
    # pad switch signals with initial value
    signalOralSwitch[0:i_start] = signalOralSwitch[i_start]
    signalNasalSwitch[0:i_start] = signalNasalSwitch[i_start]
    
    # Compute Oral ratio
    OralRatio = np.mean(signalOralSwitch)
    
    return OralRatio









def generate_report(filename, zones):
    if not filename:
        return None
    else:
        df = RRAlgotrithm(filename)
        data_report = generate_stats(df, zones)
        data_report["histogram"] = generate_histogram(df)
        return data_report


def generate_stats(resampled_df, zones):
    resampled_df['timeSeconds'] = resampled_df['timeSeconds'].astype('float')
    resampled_df['signalFrequencyBpm'] = resampled_df['signalFrequencyBpmPython'].astype('float')
    
    data_stats = dict()
    data_stats["min"] = resampled_df["signalFrequencyBpmPython"].min()
    data_stats["max"] = resampled_df["signalFrequencyBpmPython"].max()
    data_stats["avg"] = resampled_df["signalFrequencyBpmPython"].mean()
    
    resampled_df["bpmZone"] = resampled_df["signalFrequencyBpmPython"].apply(lambda x: determine_zone(x, zones))
    zones_dict = resampled_df["bpmZone"].value_counts(normalize=True).to_dict()
    
    if len(zones) == 4:
        five_zones_dict = {
            "Zone 1 (Warm Up)": 0,
            "Zone 2 (Easy)": 0,
            "Zone 3 (Aerobic)": 0,
            "Zone 4 (Threshold)": 0,
            "Zone 5 (Maximum)": 0
        }
        for key in zones_dict.keys():
            five_zones_dict[key] = zones_dict[key]
        data_stats["zones"] = list(five_zones_dict.values())
        
    elif len(zones) == 3:
        four_zones_dict = {
            "Zone 2 (Easy)": 0,
            "Zone 3 (Aerobic)": 0,
            "Zone 4 (Threshold)": 0,
            "Zone 5 (Maximum)": 0
        }
        for key in zones_dict.keys():
            four_zones_dict[key] = zones_dict[key]
        data_stats["zones"] = list(four_zones_dict.values())
        
    return data_stats

def generate_histogram(df):
    resampled_df = generate_resampled_data(df, 1, 100)
    resampled_array = resampled_df.values.tolist()
    return resampled_array


def generate_resampled_data(df, resamp_secs=1, data_points=None):

    if data_points:
        full_time = df['timeSeconds'][-1:].values[0]
        inc = full_time / data_points
        if inc > resamp_secs:
            resamp_secs = round(inc)
    

    # Convert timeSeconds to time
    df["time"] = pd.to_timedelta(df["timeSeconds"], unit="s")

    # Drop unnecessary columns
    df.drop(["timeSeconds", "tempOral", "tempNasal", "signalFrequencyBpmPhone"], axis=1, inplace=True)

    # Set time as index
    df.set_index("time", inplace=True)
    resample_bin = '{}S'.format(resamp_secs)
    resampled_df = df["signalFrequencyBpmPython"].resample(resample_bin).mean().reset_index()
    resampled_df["time"] = resampled_df["time"].dt.total_seconds()
    return resampled_df
 