import matplotlib.pyplot as plt
import numpy as np
import popsummary
import random

import matplotlib as mpl
mpl.rcParams['figure.figsize'] = (8,6)
mpl.rcParams['xtick.labelsize'] = 22
mpl.rcParams['ytick.labelsize'] = 22
mpl.rcParams['axes.grid'] = False
mpl.rcParams['grid.linestyle'] = ':'
mpl.rcParams['grid.color'] = 'grey'
mpl.rcParams['lines.linewidth'] = 2
mpl.rcParams['axes.labelsize'] = 24
mpl.rcParams['legend.handlelength'] = 2
mpl.rcParams['legend.fontsize'] = 22

from matplotlib import rc
rc('font', **{'family': 'serif', 'serif': ['Computer Modern']})
rc('text', usetex=True)

import seaborn as sns
sea = sns.color_palette("Set1")


# Load PI curves
pi_curveO4a = np.loadtxt("../../PI_curves/PI_O1_to_O4a_HLV.txt").T
pi_curveO5 = np.loadtxt("../../PI_curves/PI_Design_Apl_T2000012v2_Vpl_24884O5Tier1LowSens.txt").T


# Load NSBH and BNS rates
h5_file = '../Data/popsummary_BNS-NSBHrates.h5'
result = popsummary.popresult.PopulationResult(fname = h5_file)
bnsrate_samples = result.get_hyperparameter_samples(hyperparameters=['BNS'])
nsbh3msunrate_samples = result.get_hyperparameter_samples(hyperparameters=['NSBH3Msun'])
R0_BNS = np.percentile(bnsrate_samples, 50)
R0_NSBH = np.percentile(nsbh3msunrate_samples,50)
R0up_BNS = np.percentile(bnsrate_samples, 95)
R0up_NSBH = np.percentile(nsbh3msunrate_samples,95)
R0down_BNS = np.percentile(bnsrate_samples, 5)
R0down_NSBH = np.percentile(nsbh3msunrate_samples,5)

# Load BNS background spectrum
BNSdata = np.loadtxt('../Data/OmegaO4a_BNS.txt')
bns_f = BNSdata[:,0]
bns_med = BNSdata[:,1]
bns_up = bns_med / R0_BNS * R0up_BNS
bns_down = bns_med / R0_BNS * R0down_BNS

# Load NSBH background sprectrum
NSBHdata = np.loadtxt('../Data/OmegaO4a_NSBH.txt')
nsbh_f = NSBHdata[:,0]
nsbh_med = NSBHdata[:,1]
nsbh_up = nsbh_med / R0_NSBH * R0up_NSBH
nsbh_down = nsbh_med / R0_NSBH * R0down_NSBH

# Load BBH background spectra
# Load BBH spectra
omega_data = np.load('./../Data/omegas_o4a_default_spin_planck15lal_final.npz')
bbh_f = np.array(omega_data['freqs'])
omegadata_final = np.array(omega_data['omega_gw'])
bbh_up = np.percentile(omegadata_final, 95, axis=0)
bbh_down = np.percentile(omegadata_final, 5, axis=0)
bbh_med = np.percentile(omegadata_final, 50, axis=0)

# Downsample and scale BNS and NSBH spectra with rate R0 samples

freqs = np.logspace(np.log10(10), np.log10(2000), 150)

omegadata_bns = []
omegadata_nsbh = []
omegadata_bbh = []

nsamp1 = random.sample(range(len(bnsrate_samples)), len(omegadata_final))
nsamp2 = random.sample(range(len(nsbh3msunrate_samples)), len(omegadata_final))

for i in range(len(omegadata_final)):
    ri_bns = float(bnsrate_samples[nsamp1[i]])
    ri_nsbh = float(nsbh3msunrate_samples[nsamp2[i]])
    omegai_bns = bns_med / R0_BNS * ri_bns
    omegai_nsbh = nsbh_med / R0_NSBH * ri_nsbh
    omegai_bns = omegai_bns[:8193]
    omegai_nsbh = omegai_nsbh[:8193]
    omegai_bns = np.interp(freqs, bns_f[:8193], omegai_bns)
    omegai_nsbh = np.interp(freqs, nsbh_f[:8193], omegai_nsbh)
    omegai_bbh = np.interp(freqs, bbh_f, omegadata_final[i]) 
    omegadata_bns.append(omegai_bns)
    omegadata_nsbh.append(omegai_nsbh)
    omegadata_bbh.append(omegai_bbh)

# Get quantiles describing BBH+BNS+NSBH Omega(f)
n = len(omegadata_final)
combs = 5000000
net_05 = np.zeros(len(freqs))
net_50 = np.zeros(len(freqs))
net_95 = np.zeros(len(freqs))
for i in range(len(freqs)):
    # At each frequency, draw many combinations of OmegaBBH+OmegaBNS
    bns = np.array(omegadata_bns)[np.random.choice(range(0,n),size=combs,replace=True),i]
    bbh = np.array(omegadata_bbh)[np.random.choice(range(0,n),size=combs,replace=True),i]
    nsbh = np.array(omegadata_nsbh)[np.random.choice(range(0,n),size=combs,replace=True),i]

    # Compute an ensemble of possible totals
    total = bns+bbh+nsbh

    # Record quantiles
    net_05[i] = np.quantile(total,0.05)
    net_50[i] = np.quantile(total,0.50)
    net_95[i] = np.quantile(total,0.95)

# Generate macro files to include into the paper

omega_bns_refs = np.array([np.interp(25, bns_f, [bns_med,bns_up,bns_down][i]) for i in range(3)])
omega_nsbh_refs = np.array([np.interp(25, nsbh_f, [nsbh_med,nsbh_up,nsbh_down][i]) for i in range(3)])
omega_cbc_refs = np.array([np.interp(25, freqs, [net_50,net_95,net_05][i]) for i in range(3)])

base_power = -11
with open("../Macros/bns_macros.txt", "w") as file:
    file.write("\\newcommand{{\\bnsMedian}}{{\ensuremath {0:.1f}}}\n".format(\
                omega_bns_refs[0]/10**base_power))
    file.write("\\newcommand{{\\bnsUpperError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_bns_refs[1] - omega_bns_refs[0])/10**base_power))
    file.write("\\newcommand{{\\bnsLowerError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_bns_refs[0] - omega_bns_refs[2])/10**base_power))
    file.write("\\newcommand{{\\bnsPower}}{{\ensuremath {0}}}\n".format(base_power))

    file.write("\\newcommand{{\\bnsrateMedian}}{{\ensuremath {0:.0f}}}\n".format(R0_BNS))
    file.write("\\newcommand{{\\bnsrateUpperError}}{{\ensuremath {0:.0f}}}\n".format(\
                R0up_BNS - R0_BNS))
    file.write("\\newcommand{{\\bnsrateLowerError}}{{\ensuremath {0:.0f}}}\n".format(\
                R0_BNS - R0down_BNS))

base_power = -11
with open("../Macros/nsbh_macros.txt", "w") as file:
    file.write("\\newcommand{{\\nsbhMedian}}{{\ensuremath {0:.1f}}}\n".format(\
                omega_nsbh_refs[0]/10**base_power))
    file.write("\\newcommand{{\\nsbhUpperError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_nsbh_refs[1] - omega_nsbh_refs[0])/10**base_power))
    file.write("\\newcommand{{\\nsbhLowerError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_nsbh_refs[0] - omega_nsbh_refs[2])/10**base_power))
    file.write("\\newcommand{{\\nsbhPower}}{{\ensuremath {0}}}\n".format(base_power))

    file.write("\\newcommand{{\\nsbhrateMedian}}{{\ensuremath {0:.0f}}}\n".format(R0_NSBH))
    file.write("\\newcommand{{\\nsbhrateUpperError}}{{\ensuremath {0:.0f}}}\n".format(\
                R0up_NSBH - R0_NSBH))
    file.write("\\newcommand{{\\nsbhrateLowerError}}{{\ensuremath {0:.0f}}}\n".format(\
                R0_NSBH - R0down_NSBH))
               

base_power = -9
with open("../Macros/cbc_macros.txt", "w") as file:
    file.write("\\newcommand{{\\cbcMedian}}{{\ensuremath {0:.1f}}}\n".format(\
                omega_cbc_refs[0]/10**base_power))
    file.write("\\newcommand{{\\cbcUpperError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_cbc_refs[1] - omega_cbc_refs[0])/10**base_power))
    file.write("\\newcommand{{\\cbcLowerError}}{{\ensuremath {0:.1f}}}\n".format(\
                (omega_cbc_refs[0] - omega_cbc_refs[2])/10**base_power))
    file.write("\\newcommand{{\\cbcPower}}{{\ensuremath {0}}}".format(base_power))

fig = plt.figure(figsize=(16,6))

ylim_up = 4e-8
ylim_down = 4e-12
ax = fig.add_subplot(121)
ax.set_ylim(ylim_down,ylim_up)
ax.set_xlim(10,1500)
ax.yaxis.grid()
ax.xaxis.grid()
ax.set_xlabel(r'$f$ [Hz]')
ax.set_ylabel(r'$\Omega_\mathrm{GW} (f)$')
ax.set_xscale('log')
ax.set_yscale('log')
# Plot individual contributions
ax.fill_between(bbh_f, bbh_up, bbh_down, facecolor=sea[0], alpha=0.7, lw=0.9, zorder=1, edgecolor='black', label='BBH')
ax.fill_between(bns_f, bns_up, bns_down, facecolor=sea[1], alpha=0.7, lw=0.9,zorder=2, edgecolor='black', label='BNS')
ax.fill_between(nsbh_f, nsbh_up, nsbh_down, facecolor=sea[2], alpha=0.7, lw=0.9,zorder=3, edgecolor='black', label='NSBH')
ax.legend(loc='upper left', ncol=2)

ax = fig.add_subplot(122)
ax.set_yscale('log')
ax.set_xscale('log')
ax.yaxis.grid()
ax.xaxis.grid()
ax.set_ylim(ylim_down,ylim_up)
ax.set_xlim(10,400)
ax.set_xlabel(r'$f$ [Hz]')
ax.set_ylabel(r'$\Omega_\mathrm{GW} (f)$')

ax.fill_between(freqs, net_05, net_95, facecolor=sea[3], alpha=0.7, lw=0.9,zorder=1, edgecolor='black', label='Total GWB')
ax.loglog(freqs, net_50, color=sea[3], label='Median GWB')
ax.plot(pi_curveO4a[0],2*pi_curveO4a[1], 'k-', label='O1-O4a ($2\sigma$)')
ax.plot(pi_curveO5[0],2*pi_curveO5[1], 'k', linestyle='--', lw=1.1, label='O5 target ($2\sigma$)')


ax.legend(loc=4)

fig.tight_layout()
fig.subplots_adjust(wspace=0.3)

plt.savefig("../../paper_figures/cbc_forecast.pdf", dpi = 500, bbox_inches = "tight")
plt.savefig("../../paper_figures/cbc_forecast.png", dpi = 500, bbox_inches = "tight")

exit()