#!/usr/bin/env python3
"""Generate independent real Ei/E1 reference fixtures.

Ei values and E1 values through x=2 use 120-digit Decimal power series.
E1 values above x=2 use 1024-point Gauss-Legendre quadrature of its integral
after the substitution t=x+u/(1-u). Both methods are independent of the
runtime series and continued fraction in MathBase.SpecialFunctions.
"""

from __future__ import annotations

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

ROOT = Path(__file__).resolve().parent.parent
OUTPUT = ROOT / "tests" / "ExponentialIntegralReference.inc"
PRECISION = 120
GENERATOR_VERSION = 1
QUADRATURE_ORDER = 1024
EULER_GAMMA = Decimal(
    "0.5772156649015328606065120900824024310421593359399235988057672348848677267776646709369470632917467495146314472498070"
)
EI_ARGUMENTS = (
    "-100", "-32", "-10", "-5", "-2.000001", "-2", "-1.999999",
    "-1", "-0.3725074107813666", "-0.1", "-1e-12", "-1e-100",
    "-1e-308", "1e-308", "1e-100", "1e-12", "0.1", "0.3725074107813666",
    "1", "1.999999", "2", "2.000001", "5", "10", "32", "100",
)
E1_ARGUMENTS = (
    "1e-308", "1e-100", "1e-12", "0.000001", "0.01", "0.1", "0.5",
    "1", "1.999999", "2", "2.000001", "3", "5", "10", "32", "100",
)


def decimal_ei_positive(x: Decimal) -> Decimal:
    with localcontext() as context:
        context.prec = PRECISION
        total = Decimal(0)
        term = Decimal(1)
        for n in range(1, 2000):
            term *= x / n
            addend = term / n
            total += addend
            if abs(addend) < Decimal("1e-115"):
                break
        else:
            raise RuntimeError("Ei decimal series did not converge")
        return +(EULER_GAMMA + x.ln() + total)


def decimal_e1_series(x: Decimal) -> Decimal:
    with localcontext() as context:
        context.prec = PRECISION
        total = Decimal(0)
        term = Decimal(1)
        for n in range(1, 1000):
            term *= -x / n
            addend = term / n
            total -= addend
            if abs(addend) < Decimal("1e-115"):
                break
        else:
            raise RuntimeError("E1 decimal series did not converge")
        return +(-EULER_GAMMA - x.ln() + total)


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


NODES, WEIGHTS = gauss_legendre(QUADRATURE_ORDER)


def e1_quadrature(x: float) -> float:
    total = 0.0
    exp_minus_x = math.exp(-x)
    for node, weight in zip(NODES, WEIGHTS):
        t = 0.5 * (node + 1.0)
        one_minus_t = 1.0 - t
        u = t / one_minus_t
        total += 0.5 * weight * math.exp(-u) / (x + u) / (one_minus_t * one_minus_t)
    return exp_minus_x * total


def e1_reference(x_text: str) -> Decimal:
    x = Decimal(x_text)
    if x <= 2:
        return decimal_e1_series(x)
    return Decimal(format(e1_quadrature(float(x)), ".17g"))


def ei_reference(x_text: str) -> Decimal:
    x = Decimal(x_text)
    if x == 0:
        raise ValueError("Ei reference input cannot be zero")
    if x > 0:
        return decimal_ei_positive(x)
    return -e1_reference(str(-x))


def pascal(value: Decimal) -> str:
    return format(float(value), ".17g")


def render() -> str:
    lines = [
        "{ Generated by tools/generate_exponential_integral_data.py v1.",
        f"  Ei/E1 series precision: {PRECISION} decimal digits; E1 quadrature: {QUADRATURE_ORDER}-point Gauss-Legendre. }}",
        "const",
        f"  EiReferenceCount = {len(EI_ARGUMENTS)};",
        "  EiReferences: array[0..EiReferenceCount - 1] of record",
        "    X, Value: Double;",
        "  end = (",
    ]
    for index, x_text in enumerate(EI_ARGUMENTS):
        comma = "," if index + 1 < len(EI_ARGUMENTS) else ""
        lines.append(f"    (X: {x_text}; Value: {pascal(ei_reference(x_text))}){comma}")
    lines.extend(["  );", ""])
    lines.extend([
        f"  E1ReferenceCount = {len(E1_ARGUMENTS)};",
        "  E1References: array[0..E1ReferenceCount - 1] of record",
        "    X, Value: Double;",
        "  end = (",
    ])
    for index, x_text in enumerate(E1_ARGUMENTS):
        comma = "," if index + 1 < len(E1_ARGUMENTS) else ""
        lines.append(f"    (X: {x_text}; Value: {pascal(e1_reference(x_text))}){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("Exponential-integral 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())
