"""Complete-pulse radiation integrals for bounded separable source clouds.

Angular integration of the far-field pulse correlation gives the energy
kernel. This does not use the finite-difference field or its surface gradients.
"""

import math
from functools import lru_cache

import retarded_oracle
import source_sampling


def overlap_derivative(radius, first, second):
    """Derivative of the symmetrized temporal correlation of two pulse drives."""
    if radius >= 4:
        return 0.0
    b = math.pi / 2
    length = 4 - radius
    f = lambda k: 2 * math.sin(k * length / 2) / k
    c = lambda k: math.cos(k * length / 2)
    k_prime = (
        -0.5 * b * math.sin(b * radius) * f(2)
        - (1 + 0.5 * math.cos(b * radius)) * c(2)
        - 0.5 * b * math.sin(b * radius / 2) * (f(2 + b) + f(2 - b))
        - math.cos(b * radius / 2) * (c(2 + b) + c(2 - b))
        - 0.25 * (c(2 + 2 * b) + c(2 - 2 * b))
    ) / 4
    envelope, derivative = source_sampling.pulse_overlap(radius)
    return 0.5 * (math.cos(first - second) * (
        math.cos(radius) * derivative - math.sin(radius) * envelope
    ) + math.cos(4 + first + second) * k_prime)


def coincident_energy(first, second, intervals=2048):
    # The angular correlation limit is integral(f'_a f'_b)/(4 pi), not
    # evaluation of a singular instantaneous point-source potential.
    if intervals not in (1024, 2048, 4096):
        raise ValueError("unregistered temporal quadrature")
    step = 4 / intervals
    return math.fsum(
        (1 if i in (0, intervals) else 4 if i % 2 else 2)
        * retarded_oracle.pulse_and_derivative(i * step, first)[1]
        * retarded_oracle.pulse_and_derivative(i * step, second)[1]
        for i in range(intervals + 1)
    ) * step / (12 * math.pi)


def integrated(axes, phase):
    # Cache only deterministic integrals of immutable axes in this process.
    # A fresh CLI still computes every geometry; no saved result is loaded.
    values = _integrated(tuple(tuple(axis) for axis in axes), phase)
    return dict(zip(("energy_flux", "momentum_flux", "source_reaction"), values))


@lru_cache(maxsize=32)
def _integrated(axes, phase):
    """Energy and source reaction; axes are normalized x, y, left-z, right-z."""
    if phase not in (-math.pi / 2, 0.0, math.pi / 2):
        raise ValueError("unregistered phase")
    x, y, left, right = axes
    dx = source_sampling.differences(x, x)
    dy = source_sampling.differences(y, y)
    energy_terms = []
    for z1, z2, p1, p2, multiplicity in (
        (left, left, 0.0, 0.0, 1), (right, right, phase, phase, 1),
        (left, right, 0.0, phase, 2),
    ):
        dz = source_sampling.differences(z1, z2)
        limit = coincident_energy(p1, p2)
        def terms():
            for xx, wx in dx:
                for yy, wy in dy:
                    for zz, wz in dz:
                        radius = math.sqrt(xx * xx + yy * yy + zz * zz)
                        value = limit if radius == 0 else -overlap_derivative(radius, p1, p2) / (4 * math.pi * radius)
                        yield wx * wy * wz * value
        energy_terms.append(multiplicity * math.fsum(terms()))
    dz = source_sampling.differences(left, right)
    reaction = math.fsum(
        wx * wy * wz * (0.0 if xx == yy == zz == 0 else source_sampling.pair_impulse(xx, yy, zz, phase))
        for xx, wx in dx for yy, wy in dy for zz, wz in dz
    )
    return math.fsum(energy_terms), -reaction, reaction


def cell_axis(spacing, width, center):
    """Exact Gaussian mass of each clipped cell, assigned to its grid node."""
    if spacing not in (0.2, 0.1, 0.05) or width not in (0.2, 0.25) or center not in (0.0, -math.pi / 4, math.pi / 4):
        raise ValueError("unregistered finite-cell source")
    lower, upper = center - 3 * width, center + 3 * width
    normalizer = 2 * math.erf(3 / math.sqrt(2))
    nodes = []
    for i in range(math.floor(lower / spacing - 0.5), math.ceil(upper / spacing + 0.5) + 1):
        lo = max(lower, (i - 0.5) * spacing)
        hi = min(upper, (i + 0.5) * spacing)
        if hi <= lo:
            continue
        weight = (math.erf((hi - center) / (math.sqrt(2) * width))
                  - math.erf((lo - center) / (math.sqrt(2) * width))) / normalizer
        nodes.append((i, i * spacing, weight))
    return nodes


def axes(spacing, width):
    xy = cell_axis(spacing, width, 0.0)
    return xy, xy, cell_axis(spacing, width, -math.pi / 4), cell_axis(spacing, width, math.pi / 4)
