"""Bounded 3D scalar-wave control-volume qualification for PF-001.

This owner-maintained finite-difference solver checks transient energy and
axial momentum accounting for two finite prescribed sources. It is not a
primary experiment result or evidence of a physical propulsion mechanism.
"""

from __future__ import annotations

import argparse
import json
import math

CONFIGS = {
    "coarse": (0.5, 14, 8, 38, 0.4),
    "medium": (0.4, 18, 10, 47, 0.4),
    "medium_half_dt": (0.4, 18, 10, 94, 0.2),
    "medium_quarter_dt": (0.4, 18, 10, 188, 0.1),
    "medium_wider": (0.4, 18, 12, 47, 0.4),
    "medium_outer_wide": (0.4, 21, 10, 47, 0.4),
    "fine": (1 / 3, 21, 12, 56, 0.4),
    "medium_late": (0.4, 26, 10, 88, 0.4),
    "medium_late_half_dt": (0.4, 26, 10, 176, 0.2),
    "medium_late_outer_wide": (0.4, 30, 10, 88, 0.4),
    "fine_late": (1 / 3, 32, 12, 105, 0.4),
    "fine_late_half_dt": (1 / 3, 32, 12, 210, 0.2),
    "frozen_half_late": (0.2, 52, 20, 176, 0.4),
    "frozen_half_late_outer_wide": (0.2, 60, 20, 176, 0.4),
    "frozen_half_late_half_dt": (0.2, 52, 20, 352, 0.2),
    "frozen_third_late": (0.4 / 3, 78, 30, 352, 0.3),
}
SOURCE_FACTORS = {
    "frozen_half_late": 2, "frozen_half_late_outer_wide": 2,
    "frozen_half_late_half_dt": 2, "frozen_third_late": 3,
}
WIDTH = 0.25
WIDTHS = (0.2, 0.25)
SEPARATION = math.pi / 2
PULSE_DURATION = 4.0
STATUS = "time-domain qualification only; not a canonical result or physical validation"


def source_points(spacing: float, half_cells: int, width: float = WIDTH) -> list[tuple[int, float, float]]:
    if width not in WIDTHS:
        raise ValueError("unregistered transient source width")
    side = 2 * half_cells + 1
    plane = side * side
    raw = []
    totals = [0.0, 0.0]
    for ix in range(side):
        x = (ix - half_cells) * spacing
        if abs(x) > 3 * width:
            continue
        for iy in range(side):
            y = (iy - half_cells) * spacing
            if abs(y) > 3 * width:
                continue
            for iz in range(side):
                z = (iz - half_cells) * spacing
                weights = []
                for center in (-SEPARATION / 2, SEPARATION / 2):
                    displacement = z - center
                    if abs(displacement) > 3 * width:
                        weights.append(0.0)
                    else:
                        weights.append(math.exp(-(x * x + y * y + displacement * displacement) / (2 * width**2)))
                if any(weights):
                    raw.append((ix * plane + iy * side + iz, weights[0], weights[1]))
                    totals[0] += weights[0] * spacing**3
                    totals[1] += weights[1] * spacing**3
    if min(totals) <= 0:
        raise ValueError("source profile vanished on the grid")
    return [(index, left / totals[0], right / totals[1]) for index, left, right in raw]


def _remap_sources(points, base_side: int, factor: int):
    """Embed a nodal cloud on a nested grid without changing nodal strengths."""
    side = (base_side - 1) * factor + 1
    remapped = []
    for index, left, right in points:
        ix, remainder = divmod(index, base_side**2)
        iy, iz = divmod(remainder, base_side)
        target = ix * factor * side**2 + iy * factor * side + iz * factor
        remapped.append((target, left * factor**3, right * factor**3))
    return remapped


def configured_sources(resolution: str, width: float = WIDTH):
    spacing, half_cells, _, _, _ = CONFIGS[resolution]
    factor = SOURCE_FACTORS.get(resolution, 1)
    if half_cells % factor:
        raise ValueError("source grid does not nest in field grid")
    points = source_points(spacing * factor, half_cells // factor, width)
    return points if factor == 1 else _remap_sources(points, 2 * (half_cells // factor) + 1, factor)


def temporal_drive(time: float, phase: float) -> float:
    if not 0 <= time <= PULSE_DURATION:
        return 0.0
    return math.sin(math.pi * time / PULSE_DURATION) ** 2 * math.cos(time + phase)


def diagnostics(
    previous: list[float],
    current: list[float],
    following: list[float],
    sources: dict[int, tuple[float, float]],
    drive_left: float,
    drive_right: float,
    spacing: float,
    time_step: float,
    half_cells: int,
    control_cells: int,
) -> tuple[float, float, float, float, float, float]:
    side = 2 * half_cells + 1
    plane = side * side
    lower = half_cells - control_cells
    upper = half_cells + control_cells
    energy = momentum = energy_flux = momentum_flux = source_work = source_force = 0.0
    volume = spacing**3
    area = spacing**2
    for ix in range(lower, upper + 1):
        wx = 0.5 if ix in (lower, upper) else 1.0
        for iy in range(lower, upper + 1):
            wy = 0.5 if iy in (lower, upper) else 1.0
            for iz in range(lower, upper + 1):
                wz = 0.5 if iz in (lower, upper) else 1.0
                index = ix * plane + iy * side + iz
                ut = (following[index] - previous[index]) / (2 * time_step)
                ux = (current[index + plane] - current[index - plane]) / (2 * spacing)
                uy = (current[index + side] - current[index - side]) / (2 * spacing)
                uz = (current[index + 1] - current[index - 1]) / (2 * spacing)
                lagrangian = 0.5 * (ut * ut - ux * ux - uy * uy - uz * uz)
                energy += 0.5 * (ut * ut + ux * ux + uy * uy + uz * uz) * wx * wy * wz * volume
                momentum -= ut * uz * wx * wy * wz * volume
                source_left, source_right = sources.get(index, (0.0, 0.0))
                source = source_left * drive_left + source_right * drive_right
                source_work += source * ut * volume
                source_force += source * uz * volume
                if ix in (lower, upper):
                    normal = -1 if ix == lower else 1
                    energy_flux -= normal * ut * ux * wy * wz * area
                    momentum_flux += normal * ux * uz * wy * wz * area
                if iy in (lower, upper):
                    normal = -1 if iy == lower else 1
                    energy_flux -= normal * ut * uy * wx * wz * area
                    momentum_flux += normal * uy * uz * wx * wz * area
                if iz in (lower, upper):
                    normal = -1 if iz == lower else 1
                    energy_flux -= normal * ut * uz * wx * wy * area
                    momentum_flux += normal * (uz * uz + lagrangian) * wx * wy * area
    return energy, momentum, energy_flux, momentum_flux, source_work, source_force


def _advance(previous, current, interior, side, spacing, time_step):
    """Shared second-order propagation step; callers own bounded configurations."""
    plane = side * side
    following = [0.0] * len(current)
    for index in interior:
        laplacian = (
            current[index + plane] + current[index - plane]
            + current[index + side] + current[index - side]
            + current[index + 1] + current[index - 1]
            - 6 * current[index]
        ) / spacing**2
        following[index] = 2 * current[index] - previous[index] + time_step**2 * laplacian
    return following


def simulate(resolution: str, phase: float, width: float = WIDTH) -> dict:
    if resolution not in CONFIGS:
        raise ValueError("unknown bounded resolution")
    if phase not in (-math.pi / 2, 0.0, math.pi / 2):
        raise ValueError("phase must be a registered control")
    if width not in WIDTHS:
        raise ValueError("unregistered transient source width")
    spacing, half_cells, control_cells, steps, time_ratio = CONFIGS[resolution]
    time_step = time_ratio * spacing
    outer_half_width = half_cells * spacing
    control_half_width = control_cells * spacing
    earliest_boundary_return = 2 * outer_half_width - control_half_width - SEPARATION / 2 - 3 * width
    if steps * time_step >= earliest_boundary_return:
        raise ValueError("outer boundary could contaminate the control volume")
    side = 2 * half_cells + 1
    plane = side * side
    size = side**3
    source_list = configured_sources(resolution, width)
    source_map = {index: (left, right) for index, left, right in source_list}
    interior = [
        ix * plane + iy * side + iz
        for ix in range(1, side - 1)
        for iy in range(1, side - 1)
        for iz in range(1, side - 1)
    ]
    previous = [0.0] * size
    current = [0.0] * size
    initial = None
    last = None
    integrated = [0.0, 0.0, 0.0, 0.0]
    previous_rates = None
    peak_storage = 0.0
    for step in range(steps + 1):
        time = step * time_step
        drive_left = temporal_drive(time, 0.0)
        drive_right = temporal_drive(time, phase)
        following = _advance(previous, current, interior, side, spacing, time_step)
        for index, left, right in source_list:
            following[index] += time_step**2 * (left * drive_left + right * drive_right)
        values = diagnostics(
            previous, current, following, source_map, drive_left, drive_right,
            spacing, time_step, half_cells, control_cells,
        )
        if initial is None:
            initial = values
        if previous_rates is not None:
            for slot, rate in enumerate(values[2:]):
                integrated[slot] += 0.5 * time_step * (previous_rates[slot] + rate)
        previous_rates = values[2:]
        peak_storage = max(peak_storage, abs(values[1]))
        last = values
        previous, current = current, following
    if initial is None or last is None:
        raise RuntimeError("time-domain run produced no diagnostics")
    energy_change = last[0] - initial[0]
    momentum_change = last[1] - initial[1]
    energy_residual = energy_change + integrated[0] - integrated[2]
    momentum_residual = momentum_change + integrated[1] + integrated[3]
    return {
        "status": STATUS,
        "resolution": resolution,
        "phase": phase,
        "source_width": width,
        "source_spacing": spacing * SOURCE_FACTORS.get(resolution, 1),
        "spacing": spacing,
        "time_step": time_step,
        "steps": steps,
        "control_half_width": control_half_width,
        "outer_half_width": outer_half_width,
        "earliest_boundary_return": earliest_boundary_return,
        "source_nodes": len(source_list),
        "energy_change": energy_change,
        "momentum_change": momentum_change,
        "integrated_energy_flux": integrated[0],
        "integrated_momentum_flux": integrated[1],
        "integrated_source_work": integrated[2],
        "integrated_source_force": integrated[3],
        "energy_residual": energy_residual,
        "momentum_residual": momentum_residual,
        "peak_abs_stored_momentum": peak_storage,
    }


def assess(run: dict) -> list[str]:
    """Fail closed on malformed or non-conserving qualification output."""
    required = {
        "status", "resolution", "phase", "source_width", "source_spacing", "spacing", "time_step", "steps",
        "control_half_width", "outer_half_width", "earliest_boundary_return",
        "source_nodes", "energy_change", "momentum_change",
        "integrated_energy_flux", "integrated_momentum_flux",
        "integrated_source_work", "integrated_source_force",
        "energy_residual", "momentum_residual", "peak_abs_stored_momentum",
    }
    if not isinstance(run, dict) or set(run) != required:
        return ["incomplete time-domain 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 time-domain measurement"]
    if not isinstance(run["resolution"], str) or run["resolution"] not in CONFIGS or run["phase"] not in (-math.pi / 2, 0.0, math.pi / 2) or run["source_width"] not in WIDTHS:
        return ["unregistered time-domain configuration"]
    failures = []
    spacing, half_cells, control_cells, steps, time_ratio = CONFIGS[run["resolution"]]
    if run["status"] != STATUS or run["source_nodes"] != len(configured_sources(run["resolution"], run["source_width"])):
        failures.append("qualification status or source cohort drifted")
    if run["peak_abs_stored_momentum"] < 0:
        failures.append("stored momentum magnitude is negative")
    expected = {
        "source_spacing": spacing * SOURCE_FACTORS.get(run["resolution"], 1),
        "spacing": spacing,
        "time_step": time_ratio * spacing,
        "steps": steps,
        "control_half_width": control_cells * spacing,
        "outer_half_width": half_cells * spacing,
        "earliest_boundary_return": 2 * half_cells * spacing - control_cells * spacing - SEPARATION / 2 - 3 * run["source_width"],
    }
    if any(not math.isclose(run[key], value, rel_tol=0, abs_tol=1e-12) for key, value in expected.items()):
        failures.append("declared grid or time configuration drifted")
    if run["steps"] * run["time_step"] >= run["earliest_boundary_return"]:
        failures.append("outer boundary could contaminate control volume")
    energy_balance = run["energy_change"] + run["integrated_energy_flux"] - run["integrated_source_work"]
    momentum_balance = run["momentum_change"] + run["integrated_momentum_flux"] + run["integrated_source_force"]
    if abs(run["energy_residual"] - energy_balance) > 1e-12:
        failures.append("energy residual is inconsistent")
    if abs(run["momentum_residual"] - momentum_balance) > 1e-12:
        failures.append("momentum residual is inconsistent")
    if abs(energy_balance) > 0.006:
        failures.append("energy control-volume balance failed")
    if abs(momentum_balance) > 0.0005:
        failures.append("momentum control-volume balance failed")
    if run["phase"] != 0 and run["peak_abs_stored_momentum"] < 0.01:
        failures.append("near-field momentum storage was not observed")
    if run["phase"] == -math.pi / 2 and run["integrated_source_force"] >= 0:
        failures.append("forward phase source reaction has wrong sign")
    if run["phase"] == math.pi / 2 and run["integrated_source_force"] <= 0:
        failures.append("reverse phase source reaction has wrong sign")
    if run["phase"] == 0 and abs(run["integrated_source_force"]) > 1e-8:
        failures.append("in-phase null source reaction failed")
    return failures


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--resolution", choices=CONFIGS, default="medium")
    parser.add_argument("--phase", choices=("forward", "zero", "reverse"), default="forward")
    parser.add_argument("--width", type=float, choices=WIDTHS, default=WIDTH)
    args = parser.parse_args()
    phase = {"forward": -math.pi / 2, "zero": 0.0, "reverse": math.pi / 2}[args.phase]
    run = simulate(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()
