"""Transient source-quadrature audit with analytically integrated pulse pairs.

This measures source sampling separately from propagation, surface and time
discretization. It does not assign a canonical PF-001 result.
"""

import argparse
import json
import math

WIDTHS = (0.2, 0.25)
SEPARATION = math.pi / 2
DURATION = 4.0
STATUS = "transient source-sampling audit only; not a canonical result"
SETTINGS = (
    ("lattice", 0.4), ("lattice", 0.2), ("lattice", 0.1), ("lattice", 0.05),
    ("full", 9), ("full", 17), ("full", 33),
    ("half", 9), ("half", 17), ("half", 33),
    ("simpson", 33), ("simpson", 65),
)
# Chosen before this audit, exploratory relative reaction error allocation.
SAMPLING_LIMIT = 0.001
REFERENCE_LIMIT = 1e-5


def axis(width, rule, setting, center=0.0):
    if width not in WIDTHS or (rule, setting) not in SETTINGS:
        raise ValueError("unregistered source quadrature")
    if center not in (0.0, -SEPARATION / 2, SEPARATION / 2):
        raise ValueError("unregistered source center")
    if rule == "lattice":
        half = math.ceil((abs(center) + 3 * width) / setting)
        entries = [(i, i * setting) for i in range(-half, half + 1)
                   if abs(i * setting - center) <= 3 * width]
    else:
        entries = [(i, center + 3 * width * (2 * i / (setting - 1) - 1))
                   for i in range(setting)]
    nodes = []
    for i, position in entries:
        weight = math.exp(-0.5 * ((position - center) / width)**2)
        if rule == "half" and i in (0, setting - 1):
            weight *= 0.5
        if rule == "simpson":
            weight *= 1 if i in (0, setting - 1) else (4 if i % 2 else 2)
        nodes.append((i, position, weight))
    total = sum(node[2] for node in nodes)
    return [(i, position, weight / total) for i, position, weight in nodes]


def differences(left, right):
    """Exact tensor-product pair aggregation, keyed by integer node offsets."""
    groups = {}
    for i, x, weight in left:
        for j, y, other in right:
            delta, accumulated = groups.get(i - j, (x - y, 0.0))
            groups[i - j] = (delta, accumulated + weight * other)
    return list(groups.values())


def pulse_overlap(radius):
    """Integral of the two sin-squared envelopes and its radius derivative."""
    if not math.isfinite(radius) or radius < 0:
        raise ValueError("invalid pair distance")
    if radius >= DURATION:
        return 0.0, 0.0
    b = 2 * math.pi / DURATION
    angle = b * radius
    length = DURATION - radius
    overlap = (length * (2 + math.cos(angle)) + 3 * math.sin(angle) / b) / 8
    derivative = (-2 + 2 * math.cos(angle) - length * b * math.sin(angle)) / 8
    return overlap, derivative


def pair_impulse(dx, dy, dz, phase=-math.pi / 2):
    radius = math.sqrt(dx * dx + dy * dy + dz * dz)
    if radius <= 0 or phase not in (-math.pi / 2, 0.0, math.pi / 2):
        raise ValueError("singular pair or unregistered phase")
    overlap, derivative = pulse_overlap(radius)
    return math.sin(phase) * dz / (4 * math.pi) * (
        (math.cos(radius) * overlap + math.sin(radius) * derivative) / radius**2
        - math.sin(radius) * overlap / radius**3
    )


def measure(width, rule, setting):
    transverse = axis(width, rule, setting)
    left = axis(width, rule, setting, -SEPARATION / 2)
    right = axis(width, rule, setting, SEPARATION / 2)
    xy = differences(transverse, transverse)
    zz = differences(left, right)
    reaction = math.fsum(wx * wy * wz * pair_impulse(x, y, z)
                         for x, wx in xy for y, wy in xy for z, wz in zz)
    return {"width": width, "rule": rule, "setting": setting,
            "source_nodes": len(transverse)**2 * (len(left) + len(right)),
            "pair_groups": len(xy)**2 * len(zz), "integrated_source_force": reaction}


def evaluate():
    return {"status": STATUS, "rows": [measure(width, rule, setting)
            for width in WIDTHS for rule, setting in SETTINGS]}


def assess(run):
    if not isinstance(run, dict) or set(run) != {"status", "rows"}:
        return ["incomplete source-sampling record"]
    if run["status"] != STATUS or not isinstance(run["rows"], list) or len(run["rows"]) != len(WIDTHS) * len(SETTINGS):
        return ["source-sampling cohort drifted"]
    expected = ((width, rule, setting) for width in WIDTHS for rule, setting in SETTINGS)
    for row, (width, rule, setting) in zip(run["rows"], expected):
        if not isinstance(row, dict) or set(row) != {
            "width", "rule", "setting", "source_nodes", "pair_groups", "integrated_source_force",
        }:
            return ["malformed source-sampling row"]
        if any(type(row[key]) not in (int, float) or not math.isfinite(row[key])
               for key in row if key != "rule"):
            return ["non-finite source-sampling measurement"]
        if (row["width"], row["rule"], row["setting"]) != (width, rule, setting):
            return ["source-sampling identity drifted"]
        transverse = axis(width, rule, setting)
        left = axis(width, rule, setting, -SEPARATION / 2)
        right = axis(width, rule, setting, SEPARATION / 2)
        if row["source_nodes"] != len(transverse)**2 * (len(left) + len(right)) or row["pair_groups"] != len(differences(transverse, transverse))**2 * len(differences(left, right)):
            return ["source-sampling node cohort drifted"]
        if row["integrated_source_force"] >= -1e-8:
            return ["source-sampling reaction has wrong sign or vanished"]
        if abs(row["integrated_source_force"] - measure(width, rule, setting)["integrated_source_force"]) > 1e-12:
            return ["source-sampling measurement drifted"]
    failures = []
    for offset, width in enumerate(WIDTHS):
        rows = run["rows"][offset * len(SETTINGS):(offset + 1) * len(SETTINGS)]
        values = {(row["rule"], row["setting"]): row["integrated_source_force"] for row in rows}
        reference = values["simpson", 65]
        relative = lambda value: abs(value - reference) / abs(reference)
        if relative(values["simpson", 33]) > REFERENCE_LIMIT:
            failures.append(f"width {width}: reference source quadrature is unresolved")
        if relative(values["half", 33]) > SAMPLING_LIMIT:
            failures.append(f"width {width}: endpoint-aware sampling exceeds allocation")
        if relative(values["lattice", 0.4]) > SAMPLING_LIMIT:
            failures.append(f"width {width}: original lattice exceeds sampling allocation")
    return failures


def main():
    argparse.ArgumentParser(description=__doc__).parse_args()
    run = evaluate()
    failures = assess(run)
    print(json.dumps({**run, "qualification_failures": failures}, indent=2, allow_nan=False))
    if failures:
        raise SystemExit(1)


if __name__ == "__main__":
    main()
