import os
import glob
import numpy as np
import matplotlib as mpl
mpl.use('TkAgg')
import matplotlib.pyplot as plt
from matplotlib.widgets import Slider, Button, RadioButtons

import scipy.constants as scc
from darm_lib import goodTicks, ward_model

mpl.rcParams['axes.titlesize'] = 12
mpl.rcParams['axes.labelsize'] = 12
mpl.rcParams['xtick.labelsize'] = 12
mpl.rcParams['ytick.labelsize'] = 12
mpl.rcParams['legend.fontsize'] = 12

''' Read in DARM Data '''
script_dir = os.path.dirname(os.path.realpath(__file__))
data_dir = script_dir.replace('code', 'data')

DARM_plant_txts = glob.glob(data_dir + '/*DARM_plant_mA_per_pm*.txt')

DARMdict = {}
for dpt in DARM_plant_txts:
    temp_data = np.loadtxt(dpt, skiprows=1)
    ff = temp_data[:,0]
    TF = temp_data[:,1] * np.exp(1j * temp_data[:,2])
    unc = temp_data[:,3]

    temp = dpt.split('/')[-1]
    date = temp[0:10].replace('_', ' ')
    description = temp[32:-4].replace('_', ' ')

    DARMdict[dpt] = {}
    DARMdict[dpt]['ff'] = ff
    DARMdict[dpt]['TF'] = TF
    DARMdict[dpt]['unc'] = unc
    DARMdict[dpt]['date'] = date
    DARMdict[dpt]['description'] = description


''' IFO Parameters '''
lamb = 1064e-9
w0 = 2 * np.pi * scc.c/lamb
L = 3994.5 # meters
l_SRC0 = 56.01 # meters
FSR = scc.c/(2*L)
M = 40 # kg

requested_input_power = 36.5 # W
input_power_on_PRM = 32.0 # W
PRG = 44.0
Pbs0 = input_power_on_PRM * PRG # Input Power watts * PRG = Power on beamsplitter in watts


arm_loss0 = 38e-6
SRC_loss0 = 0.01

T_ITM0 = 0.014 # amplitude transmission of ITMs
T_ETM0 = 4e-6
T_SRM0 = 0.3234

t_ITM0 = np.sqrt(T_ITM0) # amplitude transmission of ITMs
t_ETM0 = np.sqrt(T_ETM0)
t_SRM0 = np.sqrt(T_SRM0)
r_ITM0 = np.sqrt(1 - t_ITM0**2)
r_ETM0 = np.sqrt(1 - t_ETM0**2 - arm_loss0)
r_SRM0 = np.sqrt(1 - t_SRM0**2 - SRC_loss0)

wa0 = -FSR * np.log(r_ITM0 * r_ETM0) # Hz, arm pole

phi0 = 89.5 * np.pi/180 # rads
zeta0 = 90.0 * np.pi/180 # rads

responsivity = scc.e * lamb/(scc.h * scc.c) # A/W
DCPD_QE = 0.98 # W/W
DCPD_power = 20 * 1e-3 / responsivity / DCPD_QE # W = mA * A/mA * W/A
E_LO0 = np.sqrt(DCPD_power) # sqrt(W), always 20 mA on DCPDs

# numbers from https://git.ligo.org/haocun.yu/lho_squeezing/wikis/squeezer-budget
T_OFI = 0.965
R_OMs = 0.98
T_OMC = 0.957
ModeMatchOMC = 0.95

post_SRM_loss0 = 1 - T_OFI * R_OMs * T_OMC * ModeMatchOMC # losses from back of SRM to OMC DCPDs, 1 - 85% ~ 15%

ff = np.logspace(0, np.log10(5000), 1000)
ward_params = [phi0, zeta0, E_LO0, L, l_SRC0, T_SRM0, T_ITM0, T_ETM0, Pbs0, arm_loss0, SRC_loss0, post_SRM_loss0]
origData = ward_model(ff, *ward_params)

''' Get Measurement Data '''
fit_dpt = data_dir + '/2019_08_19_DARM_plant_mA_per_pm_Aug_Spots__No_SRCL_Offset.txt'
fitDict = DARMdict[fit_dpt]

fit_ff = fitDict['ff']
fit_TF = fitDict['TF'] * 1e-3 * 1e12 / responsivity / DCPD_QE * -1 # W/m = mA/pm * A/mA * pm/m * W/A * W/W, unknown sign flip
fit_unc = fitDict['unc']
fit_date = fitDict['date']
fit_description = fitDict['description']

DARMPlantff = fit_ff
DARMPlantTF = fit_TF


''' Slider Plot '''
fig, ax = plt.subplots(2,1, sharex=True, figsize=(12,9))
s1 = ax[0]
s2 = ax[1]

# Plot measurement
s1.loglog(DARMPlantff, np.abs(DARMPlantTF), marker='.', label='March 30, 2019 H1 DARM Plant Measurement')
s2.semilogx(DARMPlantff, 180/np.pi*np.angle(DARMPlantTF), marker='.', label='March 30, 2019 H1 DARM Plant Measurement')

l3, = s1.loglog(ff, np.ones_like(ff), lw=3, ls='--', label='Saved Trace')
l4, = s2.semilogx(ff, np.zeros_like(ff), lw=3, ls='--', label='Saved Trace')

l1, = s1.loglog(ff, np.abs(origData), label='Ward DARM Model')
l2, = s2.semilogx(ff, 180/np.pi*np.angle(origData), label='Ward DARM Model')

s1.set_title('Ward DARM Optical Plant with Parameter Sliders  ' +
    r'$E_{LO}\, \frac{d E_{\zeta}}{d L_{DARM}} = \sqrt{\frac{2 P_{bs} \omega_0^2}{L^2 (\omega_{arm}^2 + \omega^2)}} \frac{t_s e^{i \beta} \left( (1-r_s e^{2 i \beta}) \cos{\phi} \cos{\zeta}  - (1+r_s e^{2 i \beta}) \sin{\phi} \sin{\zeta} \right)}{1 + r_s^2 e^{4 i \beta} - 2 r_s e^{2 i \beta} \left( \cos{2 \phi} + \frac{\kappa}{2} \sin{2 \phi} \right) }$',
    y=1.02)
s1.set_ylabel('Magnitude [W/m]')
s1.set_xlim([ff[0], ff[-1]])
s1.set_ylim([1e9, 3e11])
s1.grid()
s1.grid(which='minor', ls='--', alpha=0.3)
s1.legend(loc='upper right')

s2.set_xlabel('Frequency [Hz]')
s2.set_ylabel('Phase [degs]')
s2.set_xlim([ff[0], ff[-1]])
s2.set_ylim([-180, 180])
s2.set_yticks(-45*np.arange(-4,5))
s2.grid()
s2.grid(which='minor', ls='--', alpha=0.3)

plt.tight_layout()
plt.subplots_adjust(bottom=0.32)

axcolor = 'lightgoldenrodyellow'
axphi             = plt.axes([0.25, 0.22, 0.65, 0.01], facecolor=axcolor)
axzeta            = plt.axes([0.25, 0.20, 0.65, 0.01], facecolor=axcolor)
axE_LO            = plt.axes([0.25, 0.18, 0.65, 0.01], facecolor=axcolor)
axTS              = plt.axes([0.25, 0.16, 0.65, 0.01], facecolor=axcolor)
axTI              = plt.axes([0.25, 0.14, 0.65, 0.01], facecolor=axcolor)
axTE              = plt.axes([0.25, 0.12, 0.65, 0.01], facecolor=axcolor)
axPbs             = plt.axes([0.25, 0.10, 0.65, 0.01], facecolor=axcolor)
axarm_loss        = plt.axes([0.25, 0.08, 0.65, 0.01], facecolor=axcolor)
axSRC_loss        = plt.axes([0.25, 0.06, 0.65, 0.01], facecolor=axcolor)
axpost_SRM_loss   = plt.axes([0.25, 0.04, 0.65, 0.01], facecolor=axcolor)

sphi             = Slider(axphi,            r'Detuning $\phi$ [degs]',      85,  95, valinit=180/np.pi*phi0)
szeta            = Slider(axzeta,           r'Quadrature $\zeta$ [degs]', -180, 180, valinit=180/np.pi*zeta0)
sE_LO            = Slider(axE_LO,           r'DCPD (LO) Power $|E_\mathrm{LO}|^2$ [mW]', 10, 40, valinit=(E_LO0)**2*1e3)
saxTS            = Slider(axTS,             r'SRM Trans $t_s^2$ [%]', 0, 100, valinit=100*T_SRM0)
saxTI            = Slider(axTI,             r'ITM Trans $t_i^2$ [%]', 0, 10, valinit=100*T_ITM0)
saxTE            = Slider(axTE,             r'ETM Trans $t_e^2$ [ppm]', 0, 500, valinit=1e6*T_ETM0)
saxPbs           = Slider(axPbs,            r'Power on BS $P_{bs}$ [W]', 0, 3000, valinit=Pbs0)
saxarm_loss      = Slider(axarm_loss,       r'Arm Loss $\mathrm{Loss}_{arm}$ [ppm]', 0, 1000, valinit=arm_loss0*1e6)
saxSRC_loss      = Slider(axSRC_loss,       r'SRC Loss $\mathrm{Loss}_{SRC}$ [%]', 0, 100, valinit=SRC_loss0*1e2)
saxpost_SRM_loss = Slider(axpost_SRM_loss,  r'Post SRM Loss $\mathrm{Loss}_{OMC}$ [%]', 0, 100, valinit=post_SRM_loss0*1e2)

def update(val):
    phi = sphi.val * np.pi/180
    zeta = szeta.val * np.pi/180
    E_LO = np.sqrt(sE_LO.val * 1e-3)
    TS = saxTS.val * 1e-2
    TI = saxTI.val * 1e-2
    TE = saxTE.val * 1e-6
    Pbs = saxPbs.val
    arm_loss = saxarm_loss.val * 1e-6
    SRC_loss = saxSRC_loss.val * 1e-2
    post_SRM_loss = saxpost_SRM_loss.val * 1e-2

    new_params = [phi, zeta, E_LO, L, l_SRC0, TS, TI, TE, Pbs, arm_loss, SRC_loss, post_SRM_loss]

    new_data = ward_model(ff, *new_params)

    l1.set_ydata( np.abs(new_data) )
    l2.set_ydata( 180/np.pi * np.angle(new_data) )
    fig.canvas.draw_idle()

sphi.on_changed(update)
szeta.on_changed(update)
sE_LO.on_changed(update)
saxTS.on_changed(update)
saxTI.on_changed(update)
saxTE.on_changed(update)
saxPbs.on_changed(update)
saxarm_loss.on_changed(update)
saxSRC_loss.on_changed(update)
saxpost_SRM_loss.on_changed(update)


resetax = plt.axes([0.8, 0.00, 0.1, 0.02])
button_reset = Button(resetax, 'Reset', color=axcolor, hovercolor='0.975')
saveax = plt.axes([0.6, 0.00, 0.1, 0.02])
button_save = Button(saveax, 'Save Trace', color=axcolor, hovercolor='0.975')

def reset(event):
    sphi.reset()
    szeta.reset()
    sE_LO.reset()
    saxTS.reset()
    saxTI.reset()
    saxTE.reset()
    saxPbs.reset()
    saxarm_loss.reset()
    saxSRC_loss.reset()
    saxpost_SRM_loss.reset()
button_reset.on_clicked(reset)

def save(event):
    cur_data1 = l1.get_ydata()
    cur_data2 = l2.get_ydata()
    l3.set_ydata( np.abs(cur_data1) )
    l4.set_ydata( cur_data2 )
    fig.canvas.draw_idle()
button_save.on_clicked(save)

# rax = plt.axes([0.025, 0.5, 0.15, 0.15], facecolor=axcolor)
# radio = RadioButtons(rax, ('red', 'blue', 'green'), active=0)
#
#
# def colorfunc(label):
#     l.set_color(label)
#     fig.canvas.draw_idle()
# radio.on_clicked(colorfunc)

plt.show()
