import matplotlib.pyplot as plt
import numpy as np
import json

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")


# Define Madau Dickinson model
def madauDickinsonModel(z, alpha, beta, zp):

    """
    Madau-Dickinson-like SFR model for the CBC volumetric merger rate

    Parameters
    z : float or array
        Location at which to evaluate model
    alpha : float
        Power-law index governing low-redshift behavior, which grows as `(1+z)**alpha`
    beta : float
        Power-law index governing high-redshift behavior, which falls as `(1+z)**(-beta)`
    zp : float
        Transition point marking approximate location of "cosmic noon"

    Returns
    -------
    f_z : float or array
        Unnormalized MD model evaluated across `z`
    """

    return (1.+z)**alpha/(1.+((1.+z)/(1.+zp))**(alpha+beta))

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

# Load samples from O4a
samples = np.load('./../Data/default_O4A_population_MD_rate_stochfinal.npy')[()]


# Load rate samples for O3
O3_data = np.load('./../Data/processed_emcee_samples_noStochastic_r00r01.npy')
R0_O3 = O3_data[:, 2]
alpha_z_O3 = O3_data[:, 10]
beta_z_O3 = O3_data[:, 11]
zPeak_O3 = O3_data[:, 12]

z_grid = np.linspace(0, 5, 1000)

R_zs = samples['R0'][:, np.newaxis] \
            * madauDickinsonModel(
                  z_grid[np.newaxis, :],
                  samples['alpha_z'][:, np.newaxis],
                  samples['beta_z'][:, np.newaxis],
                  samples['zPeak'][:, np.newaxis]) \
            / madauDickinsonModel(
                  0.,
                  samples['alpha_z'][:, np.newaxis],
                  samples['beta_z'][:, np.newaxis],
                  samples['zPeak'][:, np.newaxis]) 

R_zs_O3 = R0_O3[:, np.newaxis] \
            * madauDickinsonModel(
                  z_grid[np.newaxis, :],
                  alpha_z_O3[:, np.newaxis],
                  beta_z_O3[:, np.newaxis],
                  zPeak_O3[:, np.newaxis]) \
            / madauDickinsonModel(
                  0.,
                  alpha_z_O3[:, np.newaxis],
                  beta_z_O3[:, np.newaxis],
                  zPeak_O3[:, np.newaxis]) 

# Load BBH spectra
omega_data = np.load('./../Data/omegas_o4a_default_spin_planck15lal_final.npz')
    
freqs = np.array(omega_data['freqs'])
omega_gw = np.array(omega_data['omega_gw'])



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

ylim_up = 1e-6
ylim_down = 1e-11
ax = fig.add_subplot(121)
ax.set_ylim(ylim_down,ylim_up)
ax.set_xlim(10,1800)
ax.yaxis.grid()
ax.xaxis.grid()
ax.set_xlabel(r'$f$ [Hz]')
ax.set_ylabel(r'$\Omega_\mathrm{BBH} (f)$')
ax.set_xscale('log')
ax.set_yscale('log')
# Plot individual contributions
ax.plot(freqs, np.array(omega_data['omega_gw']).T, lw=0.12, alpha=0.12, color=sea[0])
ax.plot(freqs, np.quantile(omega_data['omega_gw'], 0.05, axis=0), color='black')
ax.plot(freqs, np.quantile(omega_data['omega_gw'], 0.95, axis=0), color='black')
ax.plot(pi_curveO4a[0],2*pi_curveO4a[1], color=sea[1], label='O1-O4a ($2\sigma$)')
ax.legend(loc='upper left')

ax = fig.add_subplot(122)
ax.set_yscale('log')
ax.yaxis.grid()
ax.xaxis.grid()
ax.set_xlim(0,5)
ax.set_ylim(1e-1,1e4)
ax.set_xlabel('Redshift')
ax.set_ylabel('$R_\mathrm{BBH}(z)$ $[\mathrm{Gpc}^{-3}\,\mathrm{yr}^{-1}]$')

ax.plot(z_grid, R_zs.T, lw=0.12, alpha=0.12, color=sea[0])
ax.plot(z_grid, np.quantile(R_zs_O3, 0.05, axis=0), color='black', ls=':', label='O3')
ax.plot(z_grid, np.quantile(R_zs_O3, 0.95, axis=0), color='black', ls=':')
ax.plot(z_grid, np.quantile(R_zs, 0.05, axis=0), color='black', label='O4a')
ax.plot(z_grid, np.quantile(R_zs, 0.95, axis=0), color='black')
ax.legend(loc='lower left')

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

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

exit()