"""Shared Planck-2018 pilot vector and massive-neutrino background closure.

The values here define a pre-data-contact pilot reference, not an Apeiron fit.
Both Apeiron AP1-M1 and the LambdaCDM null model must use the same object.
"""
from __future__ import annotations

from dataclasses import asdict, dataclass
from functools import lru_cache
import json
import math
from pathlib import Path

import numpy as np
from scipy.integrate import quad


MPL_REDUCED_EV = 2.435e27
HBAR_EV_S = 6.582119569e-16
KB_EV_K = 8.617333262e-5
MPC_M = 3.0856775814913673e22
FD_MASSLESS_INTEGRAL = 7.0 * math.pi**4 / 120.0


@dataclass(frozen=True)
class Planck2018Pilot:
    H0_km_s_Mpc: float = 67.4
    omega_b: float = 0.0224
    omega_c: float = 0.120
    T_gamma0_K: float = 2.7255
    N_eff: float = 3.046
    neutrino_scheme: str = "two_massless_plus_one_massive"
    massive_neutrino_eV: float = 0.06
    chemical_potential: float = 0.0

    def validate(self) -> None:
        numeric = [
            self.H0_km_s_Mpc, self.omega_b, self.omega_c, self.T_gamma0_K,
            self.N_eff, self.massive_neutrino_eV, self.chemical_potential,
        ]
        if not np.all(np.isfinite(numeric)):
            raise ValueError("pilot vector must be finite")
        if min(self.H0_km_s_Mpc, self.omega_b, self.omega_c, self.T_gamma0_K, self.N_eff) <= 0.0:
            raise ValueError("positive Planck pilot parameters required")
        if self.massive_neutrino_eV < 0.0 or self.chemical_potential != 0.0:
            raise ValueError("only the frozen zero-chemical-potential baseline is supported")
        if self.neutrino_scheme != "two_massless_plus_one_massive":
            raise ValueError("unsupported neutrino scheme")


def h0_to_mpl(H0_km_s_Mpc: float) -> float:
    return (H0_km_s_Mpc * 1000.0 / MPC_M) * HBAR_EV_S / MPL_REDUCED_EV


def derived_reference(p: Planck2018Pilot) -> dict[str, float]:
    p.validate()
    h = p.H0_km_s_Mpc / 100.0
    H0 = h0_to_mpl(p.H0_km_s_Mpc)
    rho_crit0 = 3.0 * H0 * H0
    T_gamma0_Mpl = KB_EV_K * p.T_gamma0_K / MPL_REDUCED_EV
    rho_gamma0 = math.pi**2 * T_gamma0_Mpl**4 / 15.0
    Omega_gamma0 = rho_gamma0 / rho_crit0
    return {
        "h": h,
        "H0_Mpl": H0,
        "rho_crit0_Mpl4": rho_crit0,
        "Omega_b0": p.omega_b / h**2,
        "Omega_c0": p.omega_c / h**2,
        "Omega_gamma0": Omega_gamma0,
        "Omega_nu_massless_limit0": p.N_eff * (7.0 / 8.0) * (4.0 / 11.0) ** (4.0 / 3.0) * Omega_gamma0,
        "T_gamma0_Mpl": T_gamma0_Mpl,
        "T_nu0_eV": (4.0 / 11.0) ** (1.0 / 3.0) * KB_EV_K * p.T_gamma0_K,
    }


@lru_cache(maxsize=4096)
def _fd_ratios(y_key: float) -> tuple[float, float]:
    y = float(y_key)
    if y < 0.0 or not math.isfinite(y):
        raise ValueError("finite non-negative m a/T_nu required")
    if y == 0.0:
        return 1.0, 1.0

    def distribution(q: float) -> float:
        if q > 50.0:
            return math.exp(-q)
        return 1.0 / (math.exp(q) + 1.0)

    energy = quad(
        lambda q: q * q * math.sqrt(q * q + y * y) * distribution(q),
        0.0, 60.0, epsabs=1e-11, epsrel=2e-11, limit=200,
    )[0]
    pressure = quad(
        lambda q: q**4 / math.sqrt(q * q + y * y) * distribution(q),
        0.0, 60.0, epsabs=1e-11, epsrel=2e-11, limit=200,
    )[0]
    return energy / FD_MASSLESS_INTEGRAL, pressure / FD_MASSLESS_INTEGRAL


def fd_ratios(y: float) -> tuple[float, float]:
    # Rounding is far below the numerical quadrature error relevant to AP1-R2a.
    return _fd_ratios(round(float(y), 12))


def neutrino_background(N: np.ndarray, p: Planck2018Pilot) -> dict[str, np.ndarray]:
    """Return total neutrino rho and p in reduced-Planck density units."""

    d = derived_reference(p)
    N = np.asarray(N, dtype=float)
    if not np.all(np.isfinite(N)):
        raise ValueError("finite N grid required")
    a = np.exp(N)
    one_massless0 = d["rho_crit0_Mpl4"] * d["Omega_nu_massless_limit0"] / 3.0
    massless_two = 2.0 * one_massless0 * a**-4
    y = a * p.massive_neutrino_eV / d["T_nu0_eV"]
    ratios = np.array([fd_ratios(float(value)) for value in y.ravel()]).reshape(y.shape + (2,))
    massive_rho = one_massless0 * a**-4 * ratios[..., 0]
    massive_p = one_massless0 * a**-4 * ratios[..., 1] / 3.0
    return {
        "rho": massless_two + massive_rho,
        "pressure": massless_two / 3.0 + massive_p,
        "rho_massless": massless_two,
        "rho_massive": massive_rho,
        "pressure_massive": massive_p,
        "y_massive": y,
    }


def shared_standard_background(N: np.ndarray, p: Planck2018Pilot) -> dict[str, np.ndarray]:
    d = derived_reference(p)
    N = np.asarray(N, dtype=float)
    rho_c0 = d["rho_crit0_Mpl4"]
    nu = neutrino_background(N, p)
    return {
        "rho_b": rho_c0 * d["Omega_b0"] * np.exp(-3.0 * N),
        "rho_c": rho_c0 * d["Omega_c0"] * np.exp(-3.0 * N),
        "rho_gamma": rho_c0 * d["Omega_gamma0"] * np.exp(-4.0 * N),
        "rho_nu": nu["rho"],
        "p_nu": nu["pressure"],
    }


def build_reference_report() -> dict:
    p = Planck2018Pilot()
    d = derived_reference(p)
    nu0 = neutrino_background(np.array([0.0]), p)
    Omega_nu0 = float(nu0["rho"][0] / d["rho_crit0_Mpl4"])
    return {
        "schema": "apeiron-ap1-planck2018-neutrino-closure-v1.0",
        "status": "REFERENCE_AND_NEUTRINO_BACKGROUND_CLOSURE_PASS",
        "reference": asdict(p),
        "derived": {**d, "Omega_nu0": Omega_nu0,
                    "Omega_m0_including_massive_neutrino": d["Omega_b0"] + d["Omega_c0"] + Omega_nu0},
        "provenance": {
            "Planck_2018": "https://arxiv.org/abs/1807.06209",
            "NIST_hbar": "https://physics.nist.gov/cgi-bin/cuu/Value?hbarev=",
            "NIST_kB": "https://physics.nist.gov/cgi-bin/cuu/Value?kev=",
        },
        "role": "shared pre-data-contact pilot vector for Apeiron and LambdaCDM; not an Apeiron fit",
    }


def main() -> None:
    import argparse

    parser = argparse.ArgumentParser()
    parser.add_argument("--output", type=Path)
    args = parser.parse_args()
    text = json.dumps(build_reference_report(), indent=2) + "\n"
    if args.output:
        args.output.write_text(text, encoding="utf-8")
    else:
        print(text, end="")


if __name__ == "__main__":
    main()
