#!/usr/bin/env python3
"""Generate high-precision Jacobi references with Decimal AGM arithmetic."""

from __future__ import annotations

import argparse
from decimal import Decimal, localcontext
from pathlib import Path

ROOT = Path(__file__).resolve().parent.parent
OUTPUT = ROOT / "tests" / "JacobiEllipticReference.inc"
GENERATOR_VERSION = 1
PRECISION = 100
PI = Decimal("3.141592653589793238462643383279502884197169399375105820974944592307816406286")
PAIRS = (
    (0.1, 1e-12), (0.7, 0.01), (1.2, 0.1), (0.8, 0.4225),
    (-0.8, 0.4225), (1.4, 0.5), (1.85, 0.5), (-1.85, 0.5),
    (2.0, 0.5), (3.0, 0.5), (4.0, 0.5), (7.0, 0.5), (-7.0, 0.5),
    (12.0, 0.9), (20.0, 0.99), (25.0, 0.9999), (-25.0, 0.9999),
    (50.0, 0.1), (80.0, 0.75), (99.0, 0.5), (-100.0, 0.25),
)


def sin_decimal(value: Decimal) -> Decimal:
    value %= 2 * PI
    if value > PI:
        value -= 2 * PI
    term = value
    total = term
    index = 1
    while abs(term) > Decimal("1e-95"):
        term *= -(value * value) / Decimal((2 * index) * (2 * index + 1))
        total += term
        index += 1
    return total


def asin_decimal(value: Decimal) -> Decimal:
    term = value
    total = term
    square = value * value
    index = 0
    while abs(term) > Decimal("1e-95"):
        numerator = Decimal((2 * index + 1) ** 2)
        denominator = Decimal(2 * (index + 1) * (2 * index + 3))
        term *= square * numerator / denominator
        total += term
        index += 1
    return total


def values(u: float, m: float) -> tuple[float, float, float]:
    with localcontext() as context:
        context.prec = PRECISION
        argument, parameter = Decimal(repr(u)), Decimal(repr(m))
        a_values = [Decimal(1)]
        c_values = [Decimal(0)]
        b_value = (1 - parameter).sqrt()
        for _ in range(1, 64):
            c_value = (a_values[-1] - b_value) / 2
            a_value = (a_values[-1] + b_value) / 2
            c_values.append(c_value)
            a_values.append(a_value)
            b_value = (a_values[-2] * b_value).sqrt()
            if c_value <= Decimal("1e-90") * a_value:
                break
        else:
            raise RuntimeError(f"Decimal AGM did not converge for m={m}")
        count = len(a_values) - 1
        phi = Decimal(2**count) * a_values[count] * argument
        for index in range(count, 0, -1):
            ratio = c_values[index] / a_values[index]
            phi = (phi + asin_decimal(ratio * sin_decimal(phi))) / 2
        sn = sin_decimal(phi)
        cn = sin_decimal(PI / 2 - phi)
        dn = (1 - parameter * sn * sn).sqrt()
        return float(sn), float(cn), float(dn)


def render() -> str:
    lines = [
        f"{{ Generated by tools/generate_jacobi_elliptic_data.py v{GENERATOR_VERSION}.",
        f"  Decimal descending AGM and inverse-sine series: {PRECISION} digits. }}",
        "const",
        f"  JacobiEllipticReferenceCount = {len(PAIRS)};",
        "  JacobiEllipticReferences: array[0..JacobiEllipticReferenceCount - 1] of record",
        "    U, M, SN, CN, DN: Double;",
        "  end = (",
    ]
    for index, (u, m) in enumerate(PAIRS):
        sn, cn, dn = values(u, m)
        comma = "," if index + 1 < len(PAIRS) else ""
        lines.append(
            f"    (U: {u:.17g}; M: {m:.17g}; SN: {sn:.17g}; "
            f"CN: {cn:.17g}; DN: {dn:.17g}){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("Jacobi 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())
