import numpy as np
import matplotlib.pyplot as plt
import fpfind.lib.parse_timestamps as parser
from S15lib.g2lib.g2lib import histogram
from fpfind import TSRES
from lmfit import Model,Parameters
import pickle
import matplotlib.ticker as ticker

def g2exp(x,peak,cen,wid,base):
    return base+peak*np.exp(-2*np.abs(x-cen)/wid)

spectrum=[]
spectrum.append(np.genfromtxt('Laser7mA_4nmFilter0.csv',delimiter=',',skip_header=1))
spectrum.append(np.genfromtxt('Laser7mA_ThorEDFA1000mA_noFilter0.csv',delimiter=',',skip_header=1))
spectrum.append(np.genfromtxt('Laser7mA_ThorEDFA1000mA_noFilter_10WEDFA1120mA_noFilter0.csv',delimiter=',',skip_header=1))
spectrum.append(np.genfromtxt('Laser7mA_ThorEDFA1000mA_4nmFilter_10WEDFA1120mA_noFilter0.csv',delimiter=',',skip_header=1))

"""
plt.plot(spectrum[1][:,0],spectrum[1][:,1]-np.mean(spectrum[1][:,1]))
plt.plot(spectrum[2][:,0],spectrum[2][:,1]-np.mean(spectrum[2][:,1]))
plt.plot(spectrum[3][:,0],spectrum[3][:,1]-np.mean(spectrum[3][:,1]))
plt.xlim(1520,1600)
plt.show()
"""

g2=[]
foldername=''
with open(foldername+'20260430_1033_IDQ10_laser7mA_4nmFilter_tacq600s_partial.pickle','rb') as f:
	g2.append(pickle.load(f))
with open(foldername+'20260430_1053_IDQ10_laser7mA_EDFA1000mA_tacq600s_partial.pickle','rb') as f:
	g2.append(pickle.load(f))
with open(foldername+'20260430_1110_IDQ10_laser7mA_EDFA1000mA_10WEDFA1522mA551mW_tacq600s_partial.pickle','rb') as f:
	g2.append(pickle.load(f))
with open(foldername+'20260430_1126_IDQ10_laser7mA_EDFA1000mA_4nmFilter_maxPol_10WEDFA1522mA551mW_tacq600s_partial.pickle','rb') as f:
	g2.append(pickle.load(f))
#print(g2)

g2_params=[]
for i in range(len(g2)):
    hx=g2[i][0]
    hy=g2[i][1]
    max_time=hx[np.argmax(hy)]/256
    acc=8900
    gmodel = Model(g2exp)
    params=Parameters()
    params.add('peak',value=0.3*acc,min=0,max=0.7*acc)
    params.add('cen',value=0.99*max_time,min=0.9*max_time,max=1.1*max_time)
    params.add('wid',value=4,min=0,max=100)
    params.add('base',value=acc,min=-0.9*acc,max=1.1*acc)

    result = gmodel.fit(hy, params, weights=1/np.sqrt(hy), x=hx/256)
    #print(result.fit_report())
    g2_params.append(result.summary()['params'])
    #print(g2_params)
    #print(result.summary()['params'][0][7])
    print(result.summary()['params'][0][1]/result.summary()['params'][3][1])
    print(result.summary()['params'][0][7]/result.summary()['params'][3][1])
    print(result.summary()['params'][2][1])
    print(result.summary()['params'][2][7])
    print('\n')

x_tick_size=14
y_tick_size=14
y_axis_label_size=14
x_axis_label_size=14

fig = plt.figure(figsize=(8.27,11.69))
gs = fig.add_gridspec(4, 2, left=0.11,hspace=0.05, wspace=0.4)
axs = gs.subplots(sharex='col')
#fig.suptitle('Sharing x per column, y per row')
fig.text(0.03, 0.5, 'Intensity (dBc)', horizontalalignment='center',verticalalignment='center',rotation='vertical',fontsize=y_axis_label_size)
fig.text(0.975, 0.5, 'Coincidences (kcounts)', horizontalalignment='center',verticalalignment='center',rotation='vertical',fontsize=y_axis_label_size)
fig.text(0.5, 0.5, r'g$^{(2)}(\tau)$', horizontalalignment='center',verticalalignment='center',rotation='vertical',fontsize=y_axis_label_size)
fig.text(0.15, 0.945, 'A', horizontalalignment='center',verticalalignment='center',fontsize=16)
fig.text(0.15, 0.945-0.225, 'B', horizontalalignment='center',verticalalignment='center',fontsize=16)
fig.text(0.15, 0.945-0.225*2, 'C', horizontalalignment='center',verticalalignment='center',fontsize=16)
fig.text(0.15, 0.945-0.225*3, 'D', horizontalalignment='center',verticalalignment='center',fontsize=16)

for i in range(4):
    axs[i,0].plot(spectrum[i][:,0],spectrum[i][:,1]-np.max(spectrum[i][:,1]))
    #axs[i,0].set_ylim(-5,40)
    #axs[i,0].get_yaxis().set_visible(False)
    axs[i,0].set_yticks([-40,-30,-20,-10,0], [-40,-30,-20,-10,0])
    axs[i,0].tick_params(axis='y', labelsize=y_tick_size)

    """
    if i==0:
        axs[i,0].plot(spectrum[i][:,0],spectrum[i][:,1])#-np.mean(spectrum[i][:,1]))
    if i!=0:
        axs[i,0].plot(spectrum[i][:,0],spectrum[i][:,1]-np.mean(spectrum[i][:,1])+np.mean(spectrum[0][:,1]))
    axs[i,0].set_ylim(-70,-25)
    axs[i,0].set_yticklabels([])
    axs[i,0].set_yscale('log')
    axs[i,0].yaxis.set_minor_locator(ticker.LogLocator(base=10.0,subs=np.arange(0.1,1,0.1)))
"""
axs[0,0].set_ylim(-45,5)
axs[1,0].set_ylim(-45,5)
axs[2,0].set_ylim(-45,5)
axs[3,0].set_ylim(-45,5)
axs[3,0].set_xlim(1525,1575)
axs[3,0].set_xticks([1530,1550,1570])
axs[3,0].tick_params(axis='x', labelsize=x_tick_size)
axs[3,0].tick_params(axis='y', labelsize=y_tick_size)
axs[3,0].set_xlabel('Wavelength (nm)',fontsize=y_axis_label_size)

for i in range(4):
    axs[i,1].set_yticks(np.arange(1,1.6,0.1))
    axs[i,1].set_ylim(0.95,1.55)
    axs[i,1].tick_params(axis='y', labelsize=y_tick_size)

    def coin_g2(x):
        return x/base
    def g2_coin(x):
        return x*base
    secax=axs[i,1].secondary_yaxis('right',functions=(g2_coin,coin_g2))
    secax.set_yticks(np.arange(9000,14000,1000), ['9','10','11','12','13'])
    secax.tick_params(axis='y', labelsize=y_tick_size)

    hx=g2[i][0]
    hy=g2[i][1]
    peak=g2_params[i][0][1]
    cen=g2_params[i][1][1]
    wid=g2_params[i][2][1]
    base=g2_params[i][3][1]
    fit_time=np.linspace(-15,15,int(1e5))
    axs[i,1].errorbar(hx/256-cen,hy/base,yerr=np.sqrt(base)/base,fmt='o',capsize=3,fillstyle='none',markersize=3)
    axs[i,1].plot(fit_time,g2exp(fit_time,peak,0,wid,base)/base,c='k')


    """
    hx=g2[i][0]
    hy=g2[i][1]
    peak=g2_params[i][0][1]
    cen=g2_params[i][1][1]
    wid=g2_params[i][2][1]
    base=g2_params[i][3][1]
    fit_time=np.linspace(-15,15,int(1e5))
    axs[i,1].errorbar(hx/256-cen,hy,yerr=np.sqrt(base),fmt='o',capsize=3,fillstyle='none',markersize=3)
    axs[i,1].plot(fit_time,g2exp(fit_time,peak,0,wid,base),c='k')
    axs[i,1].set_yticks(np.arange(9000,14000,1000), ['9','10','11','12','13'])
    axs[i,1].set_ylim(8500,13500)

    def coin_g2(x):
        return x/base
    def g2_coin(x):
        return x*base
    secax = axs[i,1].secondary_yaxis('right', functions=(coin_g2,g2_coin))
    secax.set_yticks(np.arange(1,1.6,0.1))
    axs[i,1].tick_params(axis='y', labelsize=y_tick_size)
    secax.tick_params(axis='y', labelsize=y_tick_size)
    """

axs[3,1].set_xlim(-10,10)
axs[3,1].tick_params(axis='x', labelsize=x_tick_size)
axs[3,1].set_xlabel(r'Time delay $\tau$ (ns)',fontsize=y_axis_label_size)
#axs[0,1].set_xticks([1530,1540,1550,1560,1570])

fig.subplots_adjust(left=0.06, bottom=0.075, right=0.9, top=0.97)
plt.show()