# -*- coding: utf-8 -*-
"""
Created on Mon Nov 18 12:28:08 2024

@author: chuffard
"""

# libraries
import seaborn as sns
import matplotlib.pyplot as pl
import math
import pandas as pd
from shapely.geometry import Polygon
import datetime
import seaborn as sns
import statistics
import matplotlib.pyplot as plt
import os
import re
from scipy.cluster.hierarchy import dendrogram, linkage
import plotly.figure_factory as ff
import pandas as pd
import numpy as np
import scipy.cluster.hierarchy
from scipy.cluster.hierarchy import ward, dendrogram, leaves_list
from scipy.spatial.distance import pdist
import numpy.ma as ma
from scipy import stats
from scipy.stats import zscore
from numpy.lib.stride_tricks import as_strided
import scipy.stats
from numpy.lib import pad
import tanglegram as tg
import warnings
from matplotlib.dates import DateFormatter
import matplotlib.dates as mdates

def norm_to_zero_one(df):
    return (df - df.min()) * 1.0 / (df.max() - df.min())

def day_of_year_to_month_day(df, day_of_year_column):
    """
    Converts a day-of-year column in a DataFrame to month and day columns.
    """

    df['date'] = pd.to_datetime(df[day_of_year_column], format='%j', errors='coerce')
    df['month'] = df['date'].dt.month
    df['day'] = df['date'].dt.day

    return df

def pltcolor(lst):
    cols=[]
    for l in lst:
        if l=='A':
            cols.append('red')
        elif l=='B':
            cols.append('blue')
        else:
            cols.append('green')
    return cols

SESatn = pd.read_excel("G:/Analyses/2024b MARS deployment/Station M seasonal.xlsx", sheet_name="AllSESSoFarMARS")
SESatn.dtypes
SESatn["date_time"] = pd.to_datetime(SESatn["datetime"])
SESatn['day_of_year'] = SESatn['date_time'].dt.dayofyear
SESatn = SESatn.sort_values(by=['date_time'])
SESatn['rm150'] = SESatn['atn'].rolling(150, min_periods=23).mean()
SESatn['slope'] = SESatn['rm150'].diff()
triggerSES = 0.011
SESatn["triggersloperm3MASS"] = np.where(SESatn['slope'] >= triggerSES, 'TriggerSlope', 'WaitSlope')
count = SESatn["triggersloperm3MASS"] .str.count('TriggerSlope').sum()

print(count) 

sns.scatterplot(x='date_time', y='atn', hue='triggersloperm3MASS', data=SESatn)


#Wintertrigger 0.007


#Extract volumns of interest

SESdata = SESatn[['atn', 'datetime','day_of_year', 'month', 'slope']]
#average and SE by day of year
# Calculate mean and standard error by group
Summary = SESdata.groupby('day_of_year')['atn'].agg(['mean', 'sem'])

# Rename columns for clarity
Summary.columns = ['ATNaverage', 'ATNstandard_error']

# Calculate mean and standard error by group
Summaryslope = SESdata.groupby('day_of_year')['slope'].agg(['mean', 'sem'])

# Rename columns for clarity
Summaryslope.columns = ['slopeaverage', 'slopestandard_error']


StationM = pd.read_excel("G:/Analyses/2024b MARS deployment/Station M seasonal.xlsx", sheet_name="StationM")
StationM.dtypes
StationM["date_time"] = pd.to_datetime(StationM["Collect_date"])
StationM['combined'] = StationM['PULSE'].astype(str) + StationM['Cup_number'].astype(str)
StationM['day_of_year'] = StationM['date_time'].dt.dayofyear
StationM.rename(columns={'GAPFILL600_POCflux_mg_C_m-2_d-1': 'POC'}, inplace=True)
StationM.rename(columns={'GAPFILL600_Massflux_mg_C_m-2_d-1': 'Mass'}, inplace=True)

StationMcup = StationM.groupby('combined').agg(['mean'])
StationMcup.columns = ['_'.join(col) for col in StationMcup.columns]
StationMcup = StationMcup.sort_values(by=['date_time_mean'])

StationMdata = StationMcup[['POC_mean', 'Mass_mean','day_of_year_mean', 'Month_mean', 'date_time_mean']]
StationMdata['rm3Mass'] = StationMdata['Mass_mean'].rolling(3, min_periods=3).mean()

# Convert 'Timestamp' to datetime
StationMdata['Timestamp'] = pd.to_datetime(StationMdata['date_time_mean'])

# Calculate time difference
StationMdata['TimeDiff'] = StationMdata['Timestamp'].diff()
StationMdata['slope'] = StationMdata['rm3Mass'].diff()
filtered_StationMdata = StationMdata[StationMdata['TimeDiff'] < '11 days']
filtered_StationMdata['percentslope'] = filtered_StationMdata['slope']/filtered_StationMdata['rm3Mass']*100
filtered_StationMdata['day_of_year'] = filtered_StationMdata['day_of_year_mean'].apply(math.floor)

plt.scatter(filtered_StationMdata['Timestamp'], filtered_StationMdata['rm3Mass'], c=filtered_StationMdata['slope'], cmap='RdBu')

trigger = 70
filtered_StationMdata["triggersloperm3MASS"] = np.where(filtered_StationMdata['slope'] >= trigger, 'TriggerSlope', 'WaitSlope')
count = filtered_StationMdata["triggersloperm3MASS"].str.count('TriggerSlope').sum()

print(count) #14

WinterSpring_M = filtered_StationMdata[filtered_StationMdata['day_of_year'] < 90]

wintertriggertrigger = 70*.6
WinterSpring_M["triggersloperm3MASS"] = np.where(WinterSpring_M['slope'] >= wintertriggertrigger, 'TriggerSlope', 'WaitSlope')
count = WinterSpring_M["triggersloperm3MASS"].str.count('TriggerSlope').sum()

print(count) 
sns.scatterplot(x='date_time_mean', y='POC_mean', hue='triggersloperm3MASS', data=WinterSpring_M)




#average and SE by day of year
# Calculate mean and standard error by group
SummaryM = filtered_StationMdata.groupby('day_of_year')['POC_mean'].agg(['mean', 'sem'])

# Rename columns for clarity
SummaryM.columns = ['POCaverage', 'POCstandard_error']

#average and SE by day of year
# Calculate mean and standard error by group
SummaryslopeMass = filtered_StationMdata.groupby('day_of_year')['slope'].agg(['mean', 'sem'])

# Rename columns for clarity
SummaryslopeMass.columns = ['slopeaverage', 'slopestandard_error']

#average and SE by day of year
# Calculate mean and standard error by group
SummaryMass = filtered_StationMdata.groupby('day_of_year')['rm3Mass'].agg(['mean', 'sem'])

# Rename columns for clarity
SummaryMass.columns = ['rm3Massaverage', 'rm3Massstandard_error']

MassAtnPOC_concat = SummaryM.merge(SummaryMass, on='day_of_year', how='outer').merge(SummaryslopeMass, on='day_of_year', how='outer')

MassAtnPOC_concat.reset_index(drop=False, inplace=True)
MassAtnPOC_concat.set_index('day_of_year', inplace=True)

Norm = MassAtnPOC_concat.apply(norm_to_zero_one)

Norm = Norm.reset_index()

# Convert day of year to month and day
Norm = day_of_year_to_month_day(Norm, 'day_of_year')





#plot
# Plot the first series
#plt.errorbar(Norm['day_of_year'], Norm['POCaverage'], yerr=Norm['POCstandard_error'], label='POC', alpha = 0.1, color = "yellow")
#plt.plot(Norm['date'], Norm['POCaverage'], alpha = 0.5, color = "gray")

# Plot the second series
#plt.errorbar(Norm['day_of_year'], Norm['Massaverage'], yerr=Norm['Massstandard_error'], label='Mass', alpha = 0.1, color = "lightblue")
plt.plot(Norm['date'], Norm['rm3Massaverage'], alpha = 0.5,color = "brown")


# Plot the second series
#plt.errorbar(Norm['day_of_year'], Norm['ATNaverage'], yerr=Norm['ATNstandard_error'], label='ATN', alpha = .1, color = "gray")
plt.plot(Norm['date'], Norm['slopeaverage'], alpha = 0.5,color = "blue")


# Add labels and legend
plt.gca().xaxis.set_major_formatter(mdates.DateFormatter('%m/%d'))

plt.gcf().autofmt_xdate()
plt.legend()

# Show the plot
plt.show()


#plot

sns.scatterplot(x='date_time_mean', y='POC_mean', hue='triggersloperm3MASS', data=filtered_StationMdata)
