"""Retarded-field surface-flux qualification, independent of the grid solver."""

from __future__ import annotations

import argparse
import json
import math

import retarded_oracle

STATUS = "retarded-boundary qualification only; not a canonical result or physical validation"
CONTROL_CELLS = {
    "coarse": 8, "medium": 10, "medium_half_dt": 10,
    "medium_quarter_dt": 10, "medium_wider": 12,
    "medium_outer_wide": 10, "fine": 12,
    "medium_late": 10, "medium_late_outer_wide": 10,
    "fine_late": 12,
    "medium_late_half_dt": 10, "fine_late_half_dt": 12,
    "frozen_half_late": 20, "frozen_half_late_outer_wide": 20,
    "frozen_half_late_half_dt": 20, "frozen_third_late": 30,
}


def surface_nodes(spacing: float, control_cells: int):
    """Yield unique cubic surface nodes with their outward face quadrature weights."""
    for ix in range(-control_cells, control_cells + 1):
        for iy in range(-control_cells, control_cells + 1):
            for iz in range(-control_cells, control_cells + 1):
                coordinates = (ix, iy, iz)
                faces = []
                for axis, cell in enumerate(coordinates):
                    if abs(cell) != control_cells:
                        continue
                    other = [coordinates[index] for index in range(3) if index != axis]
                    edge_weight = math.prod(0.5 if abs(value) == control_cells else 1.0 for value in other)
                    faces.append((axis, 1 if cell > 0 else -1, edge_weight * spacing**2))
                if faces:
                    yield (ix * spacing, iy * spacing, iz * spacing), faces


def kernels(spacing: float, half_cells: int, control_cells: int, width: float = retarded_oracle.WIDTH, source_spacing: float | None = None):
    source_spacing = spacing if source_spacing is None else source_spacing
    source_half_cells = round(half_cells * spacing / source_spacing)
    left = retarded_oracle.profile(source_spacing, source_half_cells, -retarded_oracle.SEPARATION / 2, width)
    right = retarded_oracle.profile(source_spacing, source_half_cells, retarded_oracle.SEPARATION / 2, width)
    for position, faces in surface_nodes(spacing, control_cells):
        paths = []
        for phase_index, source in enumerate((left, right)):
            for x, y, z, weight in source:
                delta = (position[0] - x, position[1] - y, position[2] - z)
                radius = math.sqrt(sum(value * value for value in delta))
                if radius <= 0:
                    raise ValueError("source intersects control surface")
                paths.append((phase_index, radius, weight / (4 * math.pi * radius),
                              tuple(value / radius for value in delta)))
        yield faces, paths


def rates(time: float, phase: float, prepared) -> tuple[float, float]:
    energy_flux = momentum_flux = 0.0
    for faces, paths in prepared:
        ut = 0.0
        gradient = [0.0, 0.0, 0.0]
        for phase_index, radius, amplitude, direction in paths:
            value, derivative = retarded_oracle.pulse_and_derivative(
                time - radius, phase if phase_index else 0.0
            )
            ut += amplitude * derivative
            radial_gradient = -amplitude * (derivative + value / radius)
            for axis in range(3):
                gradient[axis] += radial_gradient * direction[axis]
        lagrangian = 0.5 * (ut * ut - sum(value * value for value in gradient))
        for axis, normal, area in faces:
            energy_flux -= normal * ut * gradient[axis] * area
            momentum_flux += normal * (
                gradient[axis] * gradient[2] + (lagrangian if axis == 2 else 0.0)
            ) * area
    return energy_flux, momentum_flux


def evaluate(resolution: str, phase: float, width: float = retarded_oracle.WIDTH) -> dict:
    if resolution not in retarded_oracle.CONFIGS or phase not in (-math.pi / 2, 0.0, math.pi / 2) or width not in retarded_oracle.WIDTHS:
        raise ValueError("unregistered boundary-oracle configuration")
    # These inner-volume widths are fixed independently of the grid solver.
    control_cells = CONTROL_CELLS[resolution]
    spacing, half_cells, time_step, steps = retarded_oracle.CONFIGS[resolution]
    source_spacing, _ = retarded_oracle.source_grid(resolution)
    prepared = list(kernels(spacing, half_cells, control_cells, width, source_spacing))
    integrated = [0.0, 0.0]
    previous = None
    for step in range(steps + 1):
        current = rates(step * time_step, phase, prepared)
        if previous is not None:
            for axis in range(2):
                integrated[axis] += time_step * (previous[axis] + current[axis]) / 2
        previous = current
    return {
        "status": STATUS, "resolution": resolution, "phase": phase, "source_width": width,
        "source_spacing": source_spacing,
        "spacing": spacing, "time_step": time_step, "steps": steps,
        "control_half_width": control_cells * spacing,
        "surface_nodes": len(prepared),
        "integrated_energy_flux": integrated[0],
        "integrated_momentum_flux": integrated[1],
    }


def assess(run: dict) -> list[str]:
    required = {
        "status", "resolution", "phase", "source_width", "source_spacing", "spacing", "time_step", "steps",
        "control_half_width", "surface_nodes", "integrated_energy_flux",
        "integrated_momentum_flux",
    }
    if not isinstance(run, dict) or set(run) != required:
        return ["incomplete boundary-oracle record"]
    numeric = required - {"status", "resolution"}
    if any(type(run[key]) not in (int, float) or not math.isfinite(run[key]) for key in numeric):
        return ["non-finite boundary-oracle measurement"]
    if not isinstance(run["resolution"], str) or run["resolution"] not in retarded_oracle.CONFIGS or run["phase"] not in (-math.pi / 2, 0.0, math.pi / 2) or run["source_width"] not in retarded_oracle.WIDTHS:
        return ["unregistered boundary-oracle configuration"]
    spacing, _, time_step, steps = retarded_oracle.CONFIGS[run["resolution"]]
    cells = CONTROL_CELLS[run["resolution"]]
    expected = {
        "source_spacing": retarded_oracle.source_grid(run["resolution"])[0],
        "spacing": spacing, "time_step": time_step, "steps": steps,
        "control_half_width": cells * spacing,
    }
    if run["status"] != STATUS or any(not math.isclose(run[key], value, abs_tol=1e-12, rel_tol=0) for key, value in expected.items()):
        return ["boundary-oracle cohort or configuration drifted"]
    if run["surface_nodes"] != sum(1 for _ in surface_nodes(spacing, cells)):
        return ["boundary-oracle surface cohort drifted"]
    return []


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--resolution", choices=retarded_oracle.CONFIGS, default="medium")
    parser.add_argument("--phase", choices=("forward", "zero", "reverse"), default="forward")
    parser.add_argument("--width", type=float, choices=retarded_oracle.WIDTHS, default=retarded_oracle.WIDTH)
    args = parser.parse_args()
    phase = {"forward": -math.pi / 2, "zero": 0.0, "reverse": math.pi / 2}[args.phase]
    run = evaluate(args.resolution, phase, args.width)
    run["qualification_failures"] = assess(run)
    print(json.dumps(run, indent=2, allow_nan=False))
    if run["qualification_failures"]:
        raise SystemExit(1)


if __name__ == "__main__":
    main()
