"""Independent retarded-field source-reaction check for PF-001 pulses.

The cross-source reaction is calculated directly from retarded pair fields,
without advancing or sampling the finite-difference field. Self-force pairs
cancel for each identical source drive. This remains model qualification.
"""

from __future__ import annotations

import argparse
import json
import math

WIDTH = 0.25
WIDTHS = (0.2, 0.25)
SEPARATION = math.pi / 2
DURATION = 4.0
STATUS = "retarded-source qualification only; not a canonical result or physical validation"
CONFIGS = {
    "coarse": (0.5, 14, 0.2, 38),
    "medium": (0.4, 18, 0.16, 47),
    "medium_half_dt": (0.4, 18, 0.08, 94),
    "medium_quarter_dt": (0.4, 18, 0.04, 188),
    "medium_wider": (0.4, 18, 0.16, 47),
    "medium_outer_wide": (0.4, 21, 0.16, 47),
    "fine": (1 / 3, 21, 2 / 15, 56),
    "medium_late": (0.4, 26, 0.16, 88),
    "medium_late_half_dt": (0.4, 26, 0.08, 176),
    "medium_late_outer_wide": (0.4, 30, 0.16, 88),
    "fine_late": (1 / 3, 32, 2 / 15, 105),
    "fine_late_half_dt": (1 / 3, 32, 1 / 15, 210),
    "frozen_half_late": (0.2, 52, 0.08, 176),
    "frozen_half_late_outer_wide": (0.2, 60, 0.08, 176),
    "frozen_half_late_half_dt": (0.2, 52, 0.04, 352),
    "frozen_third_late": (0.4 / 3, 78, 0.04, 352),
}
SOURCE_SPACING = {
    "frozen_half_late": 0.4, "frozen_half_late_outer_wide": 0.4,
    "frozen_half_late_half_dt": 0.4, "frozen_third_late": 0.4,
}


def source_grid(resolution: str) -> tuple[float, int]:
    spacing, half_cells, _, _ = CONFIGS[resolution]
    source_spacing = SOURCE_SPACING.get(resolution, spacing)
    return source_spacing, round(half_cells * spacing / source_spacing)


def profile(spacing: float, half_cells: int, center_z: float, width: float = WIDTH) -> list[tuple[float, float, float, float]]:
    if width not in WIDTHS:
        raise ValueError("unregistered retarded source width")
    nodes = []
    for ix in range(-half_cells, half_cells + 1):
        x = ix * spacing
        if abs(x) > 3 * width:
            continue
        for iy in range(-half_cells, half_cells + 1):
            y = iy * spacing
            if abs(y) > 3 * width:
                continue
            for iz in range(-half_cells, half_cells + 1):
                z = iz * spacing
                dz = z - center_z
                if abs(dz) > 3 * width:
                    continue
                weight = math.exp(-0.5 * (x * x + y * y + dz * dz) / width**2)
                nodes.append((x, y, z, weight))
    total = sum(node[3] for node in nodes)
    if total <= 0:
        raise ValueError("oracle source profile vanished")
    return [(x, y, z, weight / total) for x, y, z, weight in nodes]


def pulse_and_derivative(time: float, phase: float) -> tuple[float, float]:
    if time <= 0 or time >= DURATION:
        return 0.0, 0.0
    angle = math.pi * time / DURATION
    envelope = math.sin(angle) ** 2
    envelope_derivative = math.pi / DURATION * math.sin(2 * angle)
    carrier = time + phase
    return (
        envelope * math.cos(carrier),
        envelope_derivative * math.cos(carrier) - envelope * math.sin(carrier),
    )


def pairs(spacing: float, half_cells: int, width: float = WIDTH) -> list[tuple[float, float]]:
    left = profile(spacing, half_cells, -SEPARATION / 2, width)
    right = profile(spacing, half_cells, SEPARATION / 2, width)
    interactions = []
    for x, y, z, left_weight in left:
        for xx, yy, zz, right_weight in right:
            radius = math.dist((x, y, z), (xx, yy, zz))
            if radius <= 0:
                raise ValueError("distinct source profiles overlap at a singular point")
            interactions.append(((z - zz) * left_weight * right_weight / (4 * math.pi * radius), radius))
    return interactions


def reaction_rate(time: float, phase: float, interactions: list[tuple[float, float]]) -> float:
    left_now, _ = pulse_and_derivative(time, 0.0)
    right_now, _ = pulse_and_derivative(time, phase)
    result = 0.0
    for factor, radius in interactions:
        right_past, right_derivative = pulse_and_derivative(time - radius, phase)
        left_past, left_derivative = pulse_and_derivative(time - radius, 0.0)
        right_gradient = -(right_derivative / radius + right_past / radius**2)
        left_gradient = -(left_derivative / radius + left_past / radius**2)
        result += factor * (left_now * right_gradient - right_now * left_gradient)
    return result


def evaluate(resolution: str, phase: float, width: float = WIDTH) -> dict:
    if resolution not in CONFIGS or phase not in (-math.pi / 2, 0.0, math.pi / 2) or width not in WIDTHS:
        raise ValueError("unregistered oracle configuration")
    spacing, half_cells, time_step, steps = CONFIGS[resolution]
    source_spacing, source_half_cells = source_grid(resolution)
    interactions = pairs(source_spacing, source_half_cells, width)
    previous = None
    impulse = 0.0
    peak_rate = 0.0
    for step in range(steps + 1):
        force = reaction_rate(step * time_step, phase, interactions)
        peak_rate = max(peak_rate, abs(force))
        if previous is not None:
            impulse += 0.5 * time_step * (previous + force)
        previous = force
    return {
        "status": STATUS,
        "resolution": resolution,
        "phase": phase,
        "source_width": width,
        "source_spacing": source_spacing,
        "spacing": spacing,
        "time_step": time_step,
        "steps": steps,
        "pair_count": len(interactions),
        "integrated_source_force": impulse,
        "peak_abs_source_force_rate": peak_rate,
    }


def assess(run: dict) -> list[str]:
    required = {
        "status", "resolution", "phase", "source_width", "source_spacing", "spacing", "time_step", "steps",
        "pair_count", "integrated_source_force", "peak_abs_source_force_rate",
    }
    if not isinstance(run, dict) or set(run) != required:
        return ["incomplete retarded-source record"]
    numbers = required - {"status", "resolution"}
    if any(type(run[key]) not in (int, float) or not math.isfinite(run[key]) for key in numbers):
        return ["non-finite retarded-source 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 retarded-source configuration"]
    spacing, half_cells, time_step, steps = CONFIGS[run["resolution"]]
    source_spacing, source_half_cells = source_grid(run["resolution"])
    expected = {
        "source_spacing": source_spacing,
        "spacing": spacing,
        "time_step": time_step,
        "steps": steps,
        "pair_count": len(pairs(source_spacing, source_half_cells, run["source_width"])),
    }
    failures = []
    if run["status"] != STATUS or any(
        not math.isclose(run[key], value, rel_tol=0, abs_tol=1e-12)
        for key, value in expected.items()
    ):
        failures.append("retarded-source cohort or configuration drifted")
    if run["peak_abs_source_force_rate"] < 0:
        failures.append("source-force rate magnitude is negative")
    if run["phase"] == -math.pi / 2 and run["integrated_source_force"] >= 0:
        failures.append("forward retarded reaction has wrong sign")
    if run["phase"] == math.pi / 2 and run["integrated_source_force"] <= 0:
        failures.append("reverse retarded reaction has wrong sign")
    if run["phase"] == 0 and abs(run["integrated_source_force"]) > 1e-8:
        failures.append("in-phase retarded reaction did not vanish")
    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 = 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()
