from __future__ import annotations

import unittest

import numpy as np

from planck2018_neutrino_closure import (
    Planck2018Pilot,
    derived_reference,
    fd_ratios,
    neutrino_background,
    shared_standard_background,
)


class Planck2018NeutrinoClosureTests(unittest.TestCase):
    def setUp(self):
        self.p = Planck2018Pilot()

    def test_physical_densities_reproduce_reference_omegas(self):
        d = derived_reference(self.p)
        self.assertAlmostEqual(d["Omega_b0"], 0.04931, delta=2e-5)
        self.assertAlmostEqual(d["Omega_c0"], 0.26416, delta=2e-5)
        self.assertAlmostEqual(d["Omega_gamma0"], 5.445e-5, delta=2e-8)

    def test_fd_massless_limit(self):
        energy, pressure = fd_ratios(0.0)
        self.assertEqual(energy, 1.0)
        self.assertEqual(pressure, 1.0)

    def test_relativistic_equation_of_state(self):
        n = neutrino_background(np.array([-20.0]), self.p)
        self.assertAlmostEqual(float(n["pressure"][0] / n["rho"][0]), 1.0 / 3.0, delta=2e-7)

    def test_massive_species_is_nonrelativistic_today(self):
        n = neutrino_background(np.array([0.0]), self.p)
        w_massive = float(n["pressure_massive"][0] / n["rho_massive"][0])
        self.assertLess(w_massive, 1.0e-4)

    def test_neutrino_continuity(self):
        h = 2.0e-4
        n = neutrino_background(np.array([-h, 0.0, h]), self.p)
        derivative = (n["rho"][2] - n["rho"][0]) / (2.0 * h)
        residual = derivative + 3.0 * (n["rho"][1] + n["pressure"][1])
        scale = max(abs(derivative), abs(3.0 * (n["rho"][1] + n["pressure"][1])))
        self.assertLess(abs(residual) / scale, 2.0e-7)

    def test_total_matter_matches_planck_reference(self):
        d = derived_reference(self.p)
        n = neutrino_background(np.array([0.0]), self.p)
        omega_nu = float(n["rho"][0] / d["rho_crit0_Mpl4"])
        total = d["Omega_b0"] + d["Omega_c0"] + omega_nu
        self.assertAlmostEqual(total, 0.315, delta=0.0015)

    def test_shared_background_positive_and_finite(self):
        f = shared_standard_background(np.array([-12.0, -5.0, 0.0]), self.p)
        for value in f.values():
            self.assertTrue(np.all(np.isfinite(value)))
        for key in ("rho_b", "rho_c", "rho_gamma", "rho_nu"):
            self.assertTrue(np.all(f[key] > 0.0))


if __name__ == "__main__":
    unittest.main()
