#!/usr/bin/env python3
"""Generate high-precision Gauss 2F1 reference fixtures with Decimal series."""

from __future__ import annotations

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

ROOT = Path(__file__).resolve().parent.parent
OUTPUT = ROOT / "tests" / "HypergeometricReference.inc"
PRECISION = 160
GENERATOR_VERSION = 1
ARGUMENTS = (
    ("0", "3", "2", "0.75"),
    ("1", "1", "2", "-0.75"),
    ("1", "1", "2", "0.5"),
    ("2", "3", "4", "0.75"),
    ("2.5", "3.25", "4", "0.75"),
    ("-1.25", "2.5", "3.5", "-0.75"),
    ("-2", "3", "4", "0.75"),
    ("-2.0000001", "1.5", "4", "0.75"),
    ("-2.5", "-1.25", "0.5", "-0.75"),
    ("-2.5", "-1.25", "32", "0.75"),
    ("-16", "16", "0.5", "0.75"),
    ("16", "16", "0.5", "0.75"),
    ("16", "-16", "32", "0.75"),
    ("-16", "-16", "32", "-0.75"),
    ("-1", "8", "2", "0.249999"),
    ("-1", "8", "2", "0.25"),
    ("-1", "8", "2", "0.250001"),
    ("1e-12", "2.5", "3", "0.75"),
    ("2.5", "3.25", "4", "1e-12"),
    ("1", "2", "3", "-0.75"),
    ("1", "2", "3", "0.75"),
    ("16", "-3.5", "0.5", "-0.75"),
    ("-15.5", "-16", "1", "0.75"),
    ("15.5", "16", "32", "-0.75"),
    ("-0.5", "7.75", "0.5", "0.75"),
)


def gauss_series(a_text: str, b_text: str, c_text: str, x_text: str) -> Decimal:
    with localcontext() as context:
        context.prec = PRECISION
        a, b, c, x = map(Decimal, (a_text, b_text, c_text, x_text))
        if x == 0:
            return Decimal(1)
        total = Decimal(1)
        term = Decimal(1)
        for n in range(1, 100_000):
            term *= (a + n - 1) * (b + n - 1) * x
            term /= (c + n - 1) * n
            total += term
            if term == 0 or (
                n >= 128
                and abs(term)
                <= max(Decimal(1), abs(total)) * Decimal("1e-145")
            ):
                return +total
        raise RuntimeError(f"Decimal Gauss series did not converge: {a}, {b}, {c}, {x}")


def render() -> str:
    lines = [
        f"{{ Generated by tools/generate_hypergeometric_data.py v{GENERATOR_VERSION}.",
        f"  Gauss 2F1 reference series precision: {PRECISION} decimal digits. }}",
        "const",
        f"  HypergeometricReferenceCount = {len(ARGUMENTS)};",
        "  HypergeometricReferences: array[0..HypergeometricReferenceCount - 1] of record",
        "    A, B, C, X, Value: Double;",
        "  end = (",
    ]
    for index, arguments in enumerate(ARGUMENTS):
        values = (*arguments, format(float(gauss_series(*arguments)), ".17g"))
        comma = "," if index + 1 < len(ARGUMENTS) else ""
        lines.append(
            "    (A: {}; B: {}; C: {}; X: {}; Value: {}){}".format(*values, 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("Hypergeometric 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())
