# -*- coding: utf-8 -*-
"""
Created on Fri Apr 21 09:51:59 2017

@author: kaandorp
"""
from datetime import*
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import dateutil
from dateutil import *
from dateutil.rrule import rrule, DAILY
from dateutil.rrule import rrule, MONTHLY
from dateutil.relativedelta import relativedelta
import matplotlib
import xarray
import math 
import seaborn as sns

sns.set_style('whitegrid')
font = {'family' : 'normal',
        'weight' : 'normal',
        'size'   : 20}

matplotlib.rc('font', **font)
new_style = {'grid': False}
matplotlib.rc('axes', **new_style)
set_fontsize = 20

#read particle data:
df = pd.read_csv("Particles_TravelTimes_SpringendalseBeek.txt")
df.date = pd.to_datetime(df.date)
df.releasedate = pd.to_datetime(df.releasedate)
df = df[df.date != '2018-01-01']

#read agricultural input:
agriin_dummy = pd.read_csv('input_curve_Nitrate_Chloride.txt', delimiter='\t') 
df['releaseyear'] = df['releasedate'].dt.year 

#add NO3 input to df:
df['N_agri'] = np.nan
df['N_natural'] = np.nan
for i in agriin_dummy['Year'].values:
    df['N_agri'][df['releaseyear'] == i] = agriin_dummy['NO3[mg/l]'][agriin_dummy['Year'] == i].values

#add landuse:
df['lgnpre_reverse'] = (df['lgnpre'] -1)*-1  #here nature = 1 and agri = 0
df['lgn1998_reverse'] = (df['lgn1998'] -1)*-1

df['N_agri'][df.releaseyear < 1950] = 0  #mg/L 
df['N_natural'] = 5   #mg/L 

df['NO3_in'] = df['N_natural'] * df['lgnpre_reverse'] + df['N_agri'] * df['lgnpre']
df['NO3_in'][df.releaseyear >= 1998] = df['N_natural'] * df['lgn1998_reverse'] + df['N_agri'] * df['lgn1998']

df['NO3_in_weighted'] = df['NO3_in'] * df.RCH

#add Chloride:
df['Cl_natural'] =15     #mg/L 

df['Cl_agri'] = np.nan
for i in agriin_dummy['Year'].values:
    df['Cl_agri'][df['releaseyear'] == i] = agriin_dummy['Cl[mg/l]'][agriin_dummy['Year'] == i].values
df['Cl_agri'][df.releaseyear < 1950] = 0#18

df['Cl_in'] = df['Cl_natural'] * df['lgnpre_reverse'] + df['Cl_agri'] * df['lgnpre']
df['Cl_in'][df.releaseyear >= 1998] = df['Cl_natural'] * df['lgn1998_reverse'] + df['Cl_agri'] * df['lgn1998']

df['Cl_in_weighted'] = df['Cl_in'] * df.RCH

#add tritium
tritium = pd.read_csv('input_curve_Tritium.txt', delimiter='\t')
tritium.columns = ['releasedate', '3H']
tritium['releasedate'] = pd.to_datetime(tritium.releasedate)
df['3H'] = np.nan

for i in tritium['releasedate'].values:
    df['3H'][df['releasedate'] == i] = tritium['3H'][tritium['releasedate'] == i].values

#Tritium decay based on particle TT:
df['3Hdecayed'] = df['3H'] * math.e**(-0.056262*(df.time/365.25))
df['3Hdecay_weighted'] = df['3Hdecayed'] * df.RCH

agriin_dummy['releasedate'] = pd.to_datetime(agriin_dummy.Year*10000+1*100+1,format='%Y%m%d') #convert years to 1st day of year for plotting

#read measurements:
springendalsebeek = pd.read_csv(r"P:\archivedprojects\1209285-mars-phd\Chemie_inputreeksen\coupling_chemieTTD_backup_WCF141\dataspringendalsebeek_2018.txt", delimiter='\t', index_col='date', parse_dates = ['date'], dayfirst=True)
springendalsebeek.NO3[springendalsebeek.NO3 == 0] = np.nan  #remove mistakes where no3 = 0
springendalsebeek= springendalsebeek.convert_objects(convert_numeric=True)

#set indices:
agriin_dummy.index = agriin_dummy.releasedate
tritium.index = tritium['releasedate']
springendalsebeek = springendalsebeek.sort_index()

#Calculate WQ and plot:
fig, axes = plt.subplots(nrows=3, ncols=2, sharex=True, sharey=False, figsize=(14,12))
axes[0,1].set_title('Total catchment')
axes[0,0].set_title('Upstream catchment')

(df.groupby('date').sum()['Cl_in_weighted']/df.groupby('date').sum().RCH).plot(ax=axes[1,1],c='k',ls=':', lw=1)
ax1 = agriin_dummy['Cl[mg/l]'].plot(ax=axes[1,1], c='b', ls=':', secondary_y=True)
axes[1,1].scatter(springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].index,springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].Cl, label='Measurements downstream', c='k', s=10)
springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].Cl.rolling('1095D').mean().plot(ax=axes[1,1],c='r',ls='--', lw=1)
axes[1,1].set_ylim(0,50)
axes[1,1].grid()
axes[1,1].yaxis.set_ticks(np.arange(0, 50, 10))
axes[1,1].text("2015-1-01",43, 'd')
ax1.set_ylabel('Input Cl [mg/L]')

(df[df.subcatchment==1].groupby('date').sum()['Cl_in_weighted']/df[df.subcatchment==1].groupby('date').sum().RCH).plot(ax=axes[1,0],c='k',ls=':', lw=1)
axes[1,0].scatter(springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].index,springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].Cl, label='Measurements upstream', c='k', s=10)
ax1=agriin_dummy['Cl[mg/l]'].plot(ax=axes[1,0], c='b', ls=':', secondary_y=True)
springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].Cl.rolling('1095D').mean().plot(ax=axes[1,0],c='r',ls='--', lw=1)
axes[1,0].grid()
axes[1,0].set_ylim(0,50)
axes[1,0].set_ylabel('')
axes[1,1].yaxis.set_ticklabels([])
axes[1,0].yaxis.set_ticks(np.arange(0, 50, 10))
axes[1,0].text("2015-1-01",43, 'c')
axes[1,0].set_ylabel('Cl [mg/L]')
ax1.yaxis.set_ticklabels([])

(df.groupby('date').sum()['NO3_in_weighted']/df.groupby('date').sum().RCH).plot(ax=axes[2,1],c='k',ls=':', lw=1)
axes[2,1].scatter(springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].index,springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].NO3, label='Measurements downstream', c='k', s=10)
ax1=agriin_dummy['NO3[mg/l]'].plot(ax=axes[2,1], c='b', ls=':', secondary_y=True)
springendalsebeek[(springendalsebeek.Meetpunt == 'Downstream')].NO3.rolling('1095D').mean().plot(ax=axes[2,1],c='r',ls='--', lw=1)
axes[2,1].set_ylim(0,80)
axes[2,1].grid()
axes[2,1].set_xlabel('')
axes[2,1].yaxis.set_ticks(np.arange(0, 80, 20))
axes[2,1].text("2015-1-01",70, 'f')
axes[2,1].tick_params(axis='x', rotation=90)
ax1.set_ylabel('Input NO$_3$ [mg/L]')

(df[df.subcatchment==1].groupby('date').sum()['NO3_in_weighted']/df[df.subcatchment==1].groupby('date').sum().RCH).plot(ax=axes[2,0],c='k',ls=':', lw=1, label='Initial model run')
axes[2,0].scatter(springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].index,springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].NO3, label='Measurements', c='k', s=10)
springendalsebeek[(springendalsebeek.Meetpunt == 'Upstream')].NO3.rolling('1095D').mean().plot(ax=axes[2,0],c='r',ls='--', lw=1, label='Measurements - 3 year mean')
ax2 = agriin_dummy['NO3[mg/l]'].plot(ax=axes[2,0], c='b', ls=':', secondary_y=True, label='Input curve agriculture (right axis)')
axes[2,0].grid()
axes[2,0].set_xlabel('')
axes[2,1].yaxis.set_ticklabels([])
axes[2,0].set_ylim(0,80)
axes[2,0].set_ylabel('')
axes[2,0].yaxis.set_ticks(np.arange(0, 80, 20))
axes[2,0].text("2015-1-01",70, 'e')
axes[2,0].tick_params(axis='x', rotation=90)
axes[2,0].set_ylabel('NO$_3$ [mg/L]')
ax2.yaxis.set_ticklabels([])

(df.groupby('date').sum()['3Hdecay_weighted']/df.groupby('date').sum().RCH).plot(ax=axes[0,1],c='k',ls=':', lw=1)
tritium['3H'].plot(ax=axes[0,1], c='b', ls=':', secondary_y=False)
axes[0,1].set_xlim("1969-02-01","2018-1-01")
axes[0,1].set_ylim(0,300)
axes[0,1].grid()
axes[0,1].set_xlabel('')
axes[0,1].yaxis.set_ticks(np.arange(0, 300, 50))
axes[0,1].text("2015-1-01",260, 'b')

(df[df.subcatchment==1].groupby('date').sum()['3Hdecay_weighted']/df[df.subcatchment==1].groupby('date').sum().RCH).plot(ax=axes[0,0],c='k',ls=':', lw=1, label='Initial run')
axes[0,0].scatter(springendalsebeek.index,springendalsebeek['3H'], label='Measurements', c='k', s=10)
tritium['3H'].plot(ax=axes[0,0], c='b', ls=':', secondary_y=False)
axes[0,0].set_xlim("1969-2-01","2018-1-01")
axes[0,0].grid()
axes[0,0].set_xlabel('')
axes[0,1].yaxis.set_ticklabels([])
axes[0,0].set_ylim(0,300)
axes[0,0].set_ylabel('')
axes[0,0].yaxis.set_ticks(np.arange(0, 300, 50))
axes[0,0].set_ylabel('3H [TU]')
axes[0,0].text("2015-1-01",260, 'a')

plt.subplots_adjust(wspace=0.02, hspace=0.1, bottom=.1)

handles, labels = axes[2,0].get_legend_handles_labels()
handles2, labels2 = ax2.get_legend_handles_labels()
handles = handles + handles2
labels = labels + labels2
lgd = axes[2,1].legend(handles, labels, ncol=2, loc='lower center', bbox_to_anchor=(-0.05,-0.7))

