"""Qualify optimized arithmetic and analytic kernels against separate paths."""

import math
import unittest
from unittest.mock import patch

import numpy as np

import boundary_oracle
import retarded_oracle
import source_sampling
import time_domain
import field
import radiation


class BackendTests(unittest.TestCase):
    def test_stencil_matches_scalar_implementation(self):
        side = 9
        current = np.arange(side**3, dtype=float).reshape((side,) * 3)
        current = np.sin(current / 11)
        previous = current * 0.9
        following = np.zeros_like(current)
        interior = [i * side**2 + j * side + k for i in range(1, side - 1)
                    for j in range(1, side - 1) for k in range(1, side - 1)]
        expected = time_domain._advance(previous.ravel().tolist(), current.ravel().tolist(), interior, side, 0.2, 0.04)
        field.advance(previous, current, following, 0.2, 0.04)
        np.testing.assert_allclose(following.ravel(), expected, atol=1e-14, rtol=0)

    def test_full_legacy_run_matches_scalar_path(self):
        h, half, cells, steps, ratio = time_domain.CONFIGS["medium"]
        side = 2 * half + 1
        points = time_domain.configured_sources("medium")
        coordinates = np.array([(index // side**2, index // side % side, index % side) for index, _, _ in points])
        sources = (tuple(coordinates[:, i] for i in range(3)),
                   np.array([left for _, left, _ in points]), np.array([right for _, _, right in points]))
        expected = time_domain.simulate("medium", -math.pi / 2)
        measured = field.evolve((h, half, cells, steps, ratio * h), sources, -math.pi / 2)
        for key, value in measured.items():
            self.assertAlmostEqual(value, expected[key], delta=1e-11, msg=key)

    def test_cell_mass_matches_independent_numerical_integration(self):
        for h in (0.2, 0.1, 0.05):
            for width in (0.2, 0.25):
                for center in (0.0, -math.pi / 4, math.pi / 4):
                    nodes = radiation.cell_axis(h, width, center)
                    self.assertAlmostEqual(sum(weight for _, _, weight in nodes), 1.0, delta=1e-14)
                    for i, _, weight in nodes:
                        lo, hi = max(center - 3 * width, (i - 0.5) * h), min(center + 3 * width, (i + 0.5) * h)
                        step = (hi - lo) / 128
                        integral = math.fsum((1 if k in (0, 128) else 4 if k % 2 else 2)
                                            * math.exp(-0.5 * ((lo + k * step - center) / width)**2)
                                            for k in range(129)) * step / 3
                        normalized = integral / (math.sqrt(2 * math.pi) * width * math.erf(3 / math.sqrt(2)))
                        self.assertAlmostEqual(weight, normalized, delta=1e-10)
                indices, left, right = field.sources(h, round(10.4 / h), width)
                self.assertEqual(len(set(zip(*indices))), len(left))
                self.assertTrue(np.all(left >= 0) and np.all(right >= 0))
                self.assertAlmostEqual(float(np.sum(left)) * h**3, 1.0, delta=1e-14)
                self.assertAlmostEqual(float(np.sum(right)) * h**3, 1.0, delta=1e-14)
                actual = dict(zip(zip(*indices), zip(left * h**3, right * h**3)))
                axes = radiation.axes(h, width)
                half = round(10.4 / h)
                for slot, zaxis in enumerate(axes[2:]):
                    for ix, _, wx in axes[0]:
                        for iy, _, wy in axes[1]:
                            for iz, _, wz in zaxis:
                                self.assertAlmostEqual(actual[ix + half, iy + half, iz + half][slot], wx * wy * wz, delta=1e-16)

    def test_energy_kernel_matches_direct_temporal_integration(self):
        for radius in (0.05, 0.4, 1.6, 3.8):
            for p, q in ((0.0, 0.0), (0.0, -math.pi / 2), (-math.pi / 2, -math.pi / 2),
                         (0.0, math.pi / 2), (math.pi / 2, math.pi / 2)):
                n = 2048
                dt = (4 - radius) / n
                def term(i):
                    t = radius + i * dt
                    a = retarded_oracle.pulse_and_derivative(t, p)[0]
                    b = retarded_oracle.pulse_and_derivative(t, q)[0]
                    da = retarded_oracle.pulse_and_derivative(t - radius, p)[1]
                    db = retarded_oracle.pulse_and_derivative(t - radius, q)[1]
                    return (1 if i in (0, n) else 4 if i % 2 else 2) * (a * db + b * da) / 2
                expected = math.fsum(term(i) for i in range(n + 1)) * dt / (12 * math.pi * radius)
                actual = -radiation.overlap_derivative(radius, p, q) / (4 * math.pi * radius)
                self.assertAlmostEqual(actual, expected, delta=1e-10)
                self.assertAlmostEqual(radiation.coincident_energy(p, q, 1024), radiation.coincident_energy(p, q, 4096), delta=1e-11)

    def test_coincident_complete_pulse_limits(self):
        for phase in (-math.pi / 2, 0.0, math.pi / 2):
            limit = radiation.coincident_energy(0.0, phase)
            radius = 1e-4
            self.assertAlmostEqual(-radiation.overlap_derivative(radius, 0.0, phase) / (4 * math.pi * radius), limit, delta=1e-8)
            values = [abs(source_sampling.pair_impulse(0, 0, -r, phase)) for r in (0.1, 0.01, 0.001)]
            self.assertLessEqual(values[2], values[1])
            self.assertLessEqual(values[1], values[0])

    def test_radiation_matches_existing_independent_surface_integral(self):
        xy = source_sampling.axis(0.25, "lattice", 0.4)
        axes = (xy, xy, source_sampling.axis(0.25, "lattice", 0.4, -math.pi / 4),
                source_sampling.axis(0.25, "lattice", 0.4, math.pi / 4))
        exact = radiation.integrated(axes, -math.pi / 2)
        surface = boundary_oracle.evaluate("medium_late", -math.pi / 2)
        self.assertLess(abs(exact["energy_flux"] - surface["integrated_energy_flux"]), 0.0003)
        self.assertLess(abs(exact["momentum_flux"] - surface["integrated_momentum_flux"]), 0.0002)

    def test_backend_and_input_bounds_refuse(self):
        with self.assertRaises(ValueError):
            radiation.cell_axis(0.025, 0.25, 0.0)
        with self.assertRaises(ValueError):
            field.simulate("unbounded")
        with patch.object(np, "__version__", "unsupported"):
            with self.assertRaises(ValueError):
                field.simulate("cell_coarse")

    def test_registered_finest_controls_preserve_geometry_and_time_window(self):
        with patch.object(field, "evolve", return_value={}) as evolve:
            for width in (0.25, 0.2):
                for name in ("cell_finest", "cell_finest_half_dt", "cell_finest_outer_wide"):
                    record = field.simulate(name, width)
                    self.assertEqual(record["spacing"], 0.05)
                    self.assertEqual(record["control_half_width"], 4.0)
                    self.assertAlmostEqual(record["time_step"] * record["steps"], 14.08)
                    self.assertGreater(record["earliest_boundary_return"], 14.08)
            self.assertEqual(evolve.call_count, 6)


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