"""Bounded NumPy realization of the existing second-order scalar field stencil."""

import math

import numpy as np

import time_domain
import radiation

CONFIGS = {
    "cell_coarse": (0.2, 52, 20, 352, 0.04),
    "cell_fine": (0.1, 104, 40, 352, 0.04),
    "cell_fine_half_dt": (0.1, 104, 40, 704, 0.02),
    "cell_finest": (0.05, 208, 80, 704, 0.02),
    "cell_fine_outer_wide": (0.1, 120, 40, 352, 0.04),
    "cell_finest_half_dt": (0.05, 208, 80, 1408, 0.01),
    "cell_finest_outer_wide": (0.05, 240, 80, 704, 0.02),
}
STATUS = "finite-cell field qualification only; not a canonical result"


def advance(previous, current, following, spacing, dt):
    inner = (slice(1, -1),) * 3
    out = following[inner]
    np.add(current[2:, 1:-1, 1:-1], current[:-2, 1:-1, 1:-1], out=out)
    out += current[1:-1, 2:, 1:-1]
    out += current[1:-1, :-2, 1:-1]
    out += current[1:-1, 1:-1, 2:]
    out += current[1:-1, 1:-1, :-2]
    out -= 6 * current[inner]
    out /= spacing**2
    out *= dt**2
    out += 2 * current[inner]
    out -= previous[inner]


def diagnostics(previous, current, following, indices, left, right, drives, h, dt, half, cells):
    lo, hi = half - cells, half + cells + 1
    center = (slice(lo, hi),) * 3
    velocity = (following[center] - previous[center]) / (2 * dt)
    gradients = []
    for axis in range(3):
        plus, minus = list(center), list(center)
        plus[axis], minus[axis] = slice(lo + 1, hi + 1), slice(lo - 1, hi - 1)
        gradients.append((current[tuple(plus)] - current[tuple(minus)]) / (2 * h))
    weights = np.ones(2 * cells + 1)
    weights[[0, -1]] = 0.5
    volume_weights = weights[:, None, None] * weights[None, :, None] * weights[None, None, :]
    squares = sum(g * g for g in gradients)
    lagrangian = 0.5 * (velocity**2 - squares)
    energy = float(np.sum(0.5 * (velocity**2 + squares) * volume_weights)) * h**3
    momentum = -float(np.sum(velocity * gradients[2] * volume_weights)) * h**3
    energy_flux = momentum_flux = 0.0
    face_weights = weights[:, None] * weights[None, :]
    for axis in range(3):
        for index, normal in ((0, -1), (-1, 1)):
            face = [slice(None)] * 3
            face[axis] = index
            face = tuple(face)
            energy_flux -= normal * float(np.sum(velocity[face] * gradients[axis][face] * face_weights)) * h**2
            stress = gradients[axis][face] * gradients[2][face]
            if axis == 2:
                stress = stress + lagrangian[face]
            momentum_flux += normal * float(np.sum(stress * face_weights)) * h**2
    drive = left * drives[0] + right * drives[1]
    source_velocity = (following[indices] - previous[indices]) / (2 * dt)
    plus = (indices[0], indices[1], indices[2] + 1)
    minus = (indices[0], indices[1], indices[2] - 1)
    source_gradient = (current[plus] - current[minus]) / (2 * h)
    work = float(np.sum(drive * source_velocity)) * h**3
    force = float(np.sum(drive * source_gradient)) * h**3
    return energy, momentum, energy_flux, momentum_flux, work, force


def sources(h, half, width):
    axes = radiation.axes(h, width)
    values = {}
    for slot, zz in enumerate(axes[2:]):
        for ix, _, wx in axes[0]:
            for iy, _, wy in axes[1]:
                for iz, _, wz in zz:
                    entry = values.setdefault((ix + half, iy + half, iz + half), [0.0, 0.0])
                    entry[slot] = wx * wy * wz / h**3
    keys = np.array(list(values), dtype=np.int64)
    strengths = np.array(list(values.values()), dtype=np.float64)
    return tuple(keys[:, axis] for axis in range(3)), strengths[:, 0], strengths[:, 1]


def evolve(config, source_data, phase):
    h, half, cells, steps, dt = config
    side = 2 * half + 1
    previous = np.zeros((side,) * 3, dtype=np.float64)
    current = np.zeros_like(previous)
    following = np.zeros_like(previous)
    indices, left, right = source_data
    integrated = np.zeros(4)
    rates = None
    first = None
    peak = 0.0
    for step in range(steps + 1):
        drives = time_domain.temporal_drive(step * dt, 0.0), time_domain.temporal_drive(step * dt, phase)
        advance(previous, current, following, h, dt)
        following[indices] += dt**2 * (left * drives[0] + right * drives[1])
        values = diagnostics(previous, current, following, indices, left, right, drives, h, dt, half, cells)
        if first is None:
            first = values
        if rates is not None:
            integrated += dt * (rates + np.array(values[2:])) / 2
        rates = np.array(values[2:])
        peak = max(peak, abs(values[1]))
        previous, current, following = current, following, previous
    energy_change = values[0] - first[0]
    momentum_change = values[1] - first[1]
    return {
        "source_nodes": len(left), "energy_change": energy_change, "momentum_change": momentum_change,
        "integrated_energy_flux": float(integrated[0]), "integrated_momentum_flux": float(integrated[1]),
        "integrated_source_work": float(integrated[2]), "integrated_source_force": float(integrated[3]),
        "energy_residual": float(energy_change + integrated[0] - integrated[2]),
        "momentum_residual": float(momentum_change + integrated[1] + integrated[3]),
        "peak_abs_stored_momentum": peak,
    }


def simulate(name, width=0.25, phase=-math.pi / 2):
    if name not in CONFIGS or width not in (0.2, 0.25) or phase not in (-math.pi / 2, 0.0, math.pi / 2):
        raise ValueError("unregistered finite-cell run")
    if np.__version__ != "2.5.3":
        raise ValueError("unqualified numerical backend version")
    h, half, cells, steps, dt = CONFIGS[name]
    support = 3 * width + h / 2
    boundary_return = 2 * half * h - cells * h - math.pi / 4 - support
    if steps * dt >= boundary_return or 3 * (dt / h)**2 >= 1:
        raise ValueError("unsafe finite-cell domain or time step")
    return {"status": STATUS, "backend": "numpy-2.5.3", "configuration": name, "width": width, "phase": phase,
            "source_rule": "clipped-cell-mass-at-node-v1", "spacing": h, "time_step": dt,
            "steps": steps, "control_half_width": cells * h, "outer_half_width": half * h,
            "earliest_boundary_return": boundary_return,
            **evolve(CONFIGS[name], sources(h, half, width), phase)}
