#!/usr/bin/env python3
"""Generate independent Legendre elliptic-integral fixtures.

The references use 1024-point Gauss-Legendre quadrature of the defining
theta integrals. This is intentionally independent of the Carlson RF/RD
duplication algorithm in MathBase.SpecialFunctions.
"""

from __future__ import annotations

import argparse
import math
from pathlib import Path

ROOT = Path(__file__).resolve().parent.parent
OUTPUT = ROOT / "tests" / "EllipticReference.inc"
GENERATOR_VERSION = 1
QUADRATURE_ORDER = 1024
PAIRS = (
    (0.0, 0.0), (0.1, 0.0), (0.7, 0.0), (math.pi / 2, 0.0),
    (-math.pi / 2, 0.0), (0.2, 1e-12), (1.2, 0.01), (math.pi / 2, 0.01),
    (0.4, 0.1), (1.4, 0.1), (math.pi / 2, 0.1), (-1.4, 0.1),
    (0.8, 0.5), (1.5, 0.5), (math.pi / 2, 0.5), (-math.pi / 2, 0.5),
    (1.2, 0.9), (math.pi / 2, 0.9), (1.55, 0.99), (math.pi / 2, 0.99),
    (-1.55, 0.99), (1.4, 0.9999),
)
COMPLETE_PARAMETERS = (0.0, 1e-12, 0.01, 0.1, 0.5, 0.9, 0.99)


def gauss_legendre(order: int) -> tuple[list[float], list[float]]:
    nodes = [0.0] * order
    weights = [0.0] * order
    half = (order + 1) // 2
    for i in range(half):
        x = math.cos(math.pi * (i + 0.75) / (order + 0.5))
        for _ in range(32):
            p0, p1 = 1.0, x
            for degree in range(2, order + 1):
                p0, p1 = p1, ((2 * degree - 1) * x * p1 - (degree - 1) * p0) / degree
            derivative = order * (x * p1 - p0) / (x * x - 1.0)
            step = p1 / derivative
            x -= step
            if abs(step) <= 2e-16:
                break
        weight = 2.0 / ((1.0 - x * x) * derivative * derivative)
        nodes[i], nodes[order - i - 1] = -x, x
        weights[i] = weights[order - i - 1] = weight
    return nodes, weights


NODES, WEIGHTS = gauss_legendre(QUADRATURE_ORDER)


def integrate(phi: float, m: float) -> tuple[float, float]:
    if phi == 0.0:
        return phi, math.sin(phi)
    midpoint = phi / 2.0
    scale = phi / 2.0
    f_total = e_total = 0.0
    for node, weight in zip(NODES, WEIGHTS):
        theta = midpoint + scale * node
        radicand = 1.0 - m * math.sin(theta) ** 2
        factor = scale * weight
        f_total += factor / math.sqrt(radicand)
        e_total += factor * math.sqrt(radicand)
    return f_total, e_total


def pascal(value: float) -> str:
    return f"{value:.17g}"


def render() -> str:
    lines = [
        "{ Generated by tools/generate_elliptic_data.py v1.",
        f"  {QUADRATURE_ORDER}-point Gauss-Legendre quadrature of DLMF 19.2.4/5. }}",
        "const",
        f"  EllipticReferenceCount = {len(PAIRS)};",
        "  EllipticReferences: array[0..EllipticReferenceCount - 1] of record",
        "    Phi, M, F, E: Double;",
        "  end = (",
    ]
    for i, (phi, m) in enumerate(PAIRS):
        f_value, e_value = integrate(phi, m)
        comma = "," if i + 1 < len(PAIRS) else ""
        lines.append(
            f"    (Phi: {pascal(phi)}; M: {pascal(m)}; F: {pascal(f_value)}; "
            f"E: {pascal(e_value)}){comma}"
        )
    lines.extend(["  );", ""])
    lines.extend([
        f"  EllipticCompleteCount = {len(COMPLETE_PARAMETERS)};",
        "  EllipticCompleteReferences: array[0..EllipticCompleteCount - 1] of record",
        "    M, K, E: Double;",
        "  end = (",
    ])
    for i, m in enumerate(COMPLETE_PARAMETERS):
        k_value, e_value = integrate(math.pi / 2.0, m)
        comma = "," if i + 1 < len(COMPLETE_PARAMETERS) else ""
        lines.append(
            f"    (M: {pascal(m)}; K: {pascal(k_value)}; E: {pascal(e_value)}){comma}"
        )
    lines.extend(["  );", ""])
    return "\n".join(lines)


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--check", action="store_true")
    args = parser.parse_args()
    contents = render()
    if args.check:
        if not OUTPUT.exists() or OUTPUT.read_text(encoding="utf-8") != contents:
            print(f"{OUTPUT.relative_to(ROOT)} is stale; run this generator")
            return 1
        print("Elliptic reference corpus is current")
        return 0
    OUTPUT.write_text(contents, encoding="utf-8", newline="\n")
    print(f"Wrote {OUTPUT.relative_to(ROOT)}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
