"""Reproduce NumberDB T443, Hastings-McLeod Painleve-II q(s), at 100 digits.

    python -m pip install numberdb python-flint
    python generate.py --sample -6 -4 0 2
    python generate.py --compute-json entries.json
    python generate.py --check-fredholm
    python generate.py --publish

The last command publishes to the website and requires NUMBERDB_API_KEY.
All other commands are offline.  The arguments are exactly the existing T443
selection: every reduced rational of denominator at most four in [-6,2].

The Hastings-McLeod solution is integrated backward from a positive boundary
where q and q' are replaced by Airy Ai values.  Two runs vary working
precision, Taylor-series order, maximum step and right boundary:

    (230 decimal digits, 180 series terms, step <= 1/8, boundary 40)
    (280 decimal digits, 240 series terms, step <= 1/12, boundary 48)

The stored values are taken from the stronger run only after both runs agree
to relative error below 1e-112 at every table entry.  During integration the
Hamiltonian identity J = q'^2 - s*q^2 - q^4 is checked against the propagated
J = integral_s^infinity q(x)^2 dx.  A separate Airy-kernel Fredholm resolvent
Nyström computation checks selected endpoint and ordinary values.

These checks are numerical agreement tests, not rigorous enclosures.  Arb is
used for multiprecision arithmetic, but step results are advanced by midpoints
and the nonlinear boundary error and Taylor remainders are not enclosed.
"""

import argparse
import functools
import json
import os
from decimal import Decimal, localcontext
from fractions import Fraction
from pathlib import Path

import flint
from flint import acb, arb, arb_mat, arb_poly, arb_series, ctx
import numberdb


DIGITS = 100
TABLE = os.environ.get("NUMBERDB_TABLE", "T443")

STANDARD_CONFIGURATION = (230, 180, 8, 40)
STRONG_CONFIGURATION = (280, 240, 12, 48)
CONFIGURATIONS = (STANDARD_CONFIGURATION, STRONG_CONFIGURATION)

AGREEMENT_TOLERANCE = arb("1e-112")
HAMILTONIAN_TOLERANCE = arb("1e-112")
FREDHOLM_TOLERANCE = arb("1e-95")

DEFAULT_FREDHOLM_ARGUMENTS = ("-6", "-4", "0", "2")


def arguments():
    return sorted({
        Fraction(numerator, denominator)
        for denominator in range(1, 5)
        for numerator in range(-6 * denominator, 2 * denominator + 1)
    })


def parse_fraction(text):
    return Fraction(str(text))


def initial_state(boundary):
    value, derivative, _, _ = boundary.airy()
    integral = derivative**2 - boundary * value**2
    moment = (2 * boundary**2 * value**2 - 2 * boundary * derivative**2 - value * derivative) / 3
    tail = acb.integral(
        lambda point, analytic: point.airy_ai(),
        boundary,
        100,
        abs_tol=arb("1e-210"),
        rel_tol=arb("1e-190"),
    )
    if not tail.is_finite() or not tail.imag.contains(0) or not tail.real > 0:
        raise ArithmeticError("Invalid Airy tail integral")
    if not tail.real.rad() < arb("1e-200"):
        raise ArithmeticError("Airy tail integral is insufficiently accurate")
    return [value.mid(), derivative.mid(), tail.real.mid(), integral.mid(), moment.mid()]


def taylor_step(position, state, increment, order):
    value, derivative, tail, integral, moment = state
    series = arb_series([value, derivative], prec=order)
    variable = arb_series([position, 1], prec=order)
    for precision in range(4, order + 1, 2):
        ctx.cap = precision
        series = arb_series(series.coeffs(), prec=precision)
        series = (
            arb_series([value, derivative], prec=precision)
            + (variable * series + 2 * series**3).integral().integral()
        )
    ctx.cap = order
    polynomial = arb_poly(series.coeffs())
    square_integral = (series * series).integral()
    result = [
        polynomial(increment),
        polynomial.derivative()(increment),
        tail - polynomial.integral()(increment),
        integral - arb_poly(square_integral.coeffs())(increment),
        moment - integral * increment + arb_poly(square_integral.integral().coeffs())(increment),
    ]
    if not all(value.is_finite() for value in result):
        raise ArithmeticError("Non-finite Taylor result")
    return [value.mid() for value in result]


def compute_configuration(configuration, selected=None):
    working_digits, order, step_denominator, boundary = configuration
    old_precision, old_cap = ctx.prec, ctx.cap
    try:
        ctx.dps, ctx.cap = working_digits, order
        selected_arguments = arguments() if selected is None else sorted(parse_fraction(item) for item in selected)
        targets = [(arb(argument.numerator) / argument.denominator, str(argument))
                   for argument in selected_arguments]
        targets.sort(key=lambda target: target[0], reverse=True)
        position = arb(boundary)
        state = initial_state(position)
        maximum_step = arb(1) / step_denominator
        samples = {}
        max_hamiltonian_abs = arb(0)

        for target, argument in targets:
            while position - target > maximum_step:
                state = taylor_step(position, state, -maximum_step, order)
                position -= maximum_step
            increment = (target - position).mid()
            if not increment.is_zero():
                state = taylor_step(position, state, increment, order)
            position = target.mid()
            value, derivative, tail, integral, moment = state
            identity = derivative**2 - position * value**2 - value**4
            hamiltonian_abs = abs(identity - integral)
            max_hamiltonian_abs = max(max_hamiltonian_abs, hamiltonian_abs)
            if not hamiltonian_abs < HAMILTONIAN_TOLERANCE:
                raise ArithmeticError(f"Hamiltonian control failed at {argument}")
            if not value.is_finite() or not value > 0:
                raise ArithmeticError(f"Invalid q(s) at {argument}")
            samples[argument] = value.mid().str(working_digits - 5, radius=False)
        return samples, {"max_hamiltonian_abs": max_hamiltonian_abs.str(20, radius=False)}
    finally:
        ctx.prec, ctx.cap = old_precision, old_cap


def compare_configurations(selected=None):
    first, first_info = compute_configuration(STANDARD_CONFIGURATION, selected=selected)
    second, second_info = compute_configuration(STRONG_CONFIGURATION, selected=selected)
    if first.keys() != second.keys():
        raise ArithmeticError("The two runs cover different arguments")
    old_precision = ctx.prec
    try:
        ctx.dps = 320
        worst_relative = arb(0)
        worst_absolute = arb(0)
        worst_argument = None
        for argument, text in second.items():
            standard, strong = arb(first[argument]), arb(text)
            relative = abs((standard - strong) / strong)
            absolute = abs(standard - strong)
            if not relative < AGREEMENT_TOLERANCE:
                raise ArithmeticError(f"Insufficient numerical agreement at {argument}")
            if relative > worst_relative:
                worst_relative = relative
                worst_absolute = absolute
                worst_argument = argument
        summary = {
            "entries": len(second),
            "standard": {
                "working_digits": STANDARD_CONFIGURATION[0],
                "series_terms": STANDARD_CONFIGURATION[1],
                "max_step": f"1/{STANDARD_CONFIGURATION[2]}",
                "right_boundary": STANDARD_CONFIGURATION[3],
                **first_info,
            },
            "strong": {
                "working_digits": STRONG_CONFIGURATION[0],
                "series_terms": STRONG_CONFIGURATION[1],
                "max_step": f"1/{STRONG_CONFIGURATION[2]}",
                "right_boundary": STRONG_CONFIGURATION[3],
                **second_info,
            },
            "worst_relative_difference": worst_relative.str(20, radius=False),
            "worst_absolute_difference": worst_absolute.str(20, radius=False),
            "worst_argument": worst_argument,
        }
        return second, summary
    finally:
        ctx.prec = old_precision


@functools.lru_cache(maxsize=1)
def checked_result():
    values, summary = compare_configurations()
    with localcontext() as decimal_context:
        decimal_context.prec = 320
        rounded = {argument: format(Decimal(text), ".99e") for argument, text in values.items()}
    return rounded, summary


def airy_kernel(x, y):
    ai_x, aip_x, _, _ = x.airy()
    ai_y, aip_y, _, _ = y.airy()
    if x == y:
        return aip_x**2 - x * ai_x**2
    return (ai_x * aip_y - aip_x * ai_y) / (x - y)


def fredholm_q(argument, nodes=260, length=52, dps=210):
    old_precision = ctx.prec
    try:
        ctx.dps = dps
        fraction = parse_fraction(argument)
        s = (arb(fraction.numerator) / fraction.denominator).mid()
        quadrature = [arb.legendre_p_root(nodes, index, weight=True) for index in range(nodes)]
        abscissas = [(s + length * (point + 1) / 2).mid() for point, weight in quadrature]
        roots = [(length * weight / 2).sqrt().mid() for point, weight in quadrature]

        matrix = arb_mat(nodes, nodes)
        rhs = arb_mat(nodes, 1)
        for row in range(nodes):
            rhs[row, 0] = (roots[row] * abscissas[row].airy()[0]).mid()
            for column in range(row, nodes):
                value = (roots[row] * roots[column] * airy_kernel(abscissas[row], abscissas[column])).mid()
                matrix[row, column] = matrix[column, row] = -value
            matrix[row, row] += 1

        solution = matrix.solve(rhs)
        value = s.airy()[0].mid()
        for column in range(nodes):
            value += (airy_kernel(s, abscissas[column]) * roots[column] * solution[column, 0]).mid()
        if not value.is_finite() or not value > 0:
            raise ArithmeticError(f"Invalid Fredholm q(s) at {argument}")
        return value
    finally:
        ctx.prec = old_precision


def _rounded_values_from_texts(values):
    with localcontext() as decimal_context:
        decimal_context.prec = 320
        return {argument: format(Decimal(text), ".99e") for argument, text in values.items()}


def check_fredholm(arguments_to_check=DEFAULT_FREDHOLM_ARGUMENTS, reference_values=None):
    if reference_values is None:
        values, _summary = compare_configurations(selected=arguments_to_check)
        reference_values = _rounded_values_from_texts(values)
    old_precision = ctx.prec
    try:
        ctx.dps = 250
        controls = {}
        worst_relative = arb(0)
        worst_argument = None
        for argument in arguments_to_check:
            argument = str(parse_fraction(argument))
            value = arb(reference_values[argument])
            control = fredholm_q(argument)
            relative = abs((control - value) / value)
            controls[argument] = {
                "taylor_value": reference_values[argument],
                "fredholm_value": control.str(120, radius=False),
                "relative_difference": relative.str(20, radius=False),
            }
            if not relative < FREDHOLM_TOLERANCE:
                raise ArithmeticError(f"Fredholm control disagrees at {argument}")
            if relative > worst_relative:
                worst_relative = relative
                worst_argument = argument
        return {
            "nodes": 260,
            "length": 52,
            "working_digits": 210,
            "tolerance": FREDHOLM_TOLERANCE.str(10, radius=False),
            "worst_relative_difference": worst_relative.str(20, radius=False),
            "worst_argument": worst_argument,
            "controls": controls,
        }
    finally:
        ctx.prec = old_precision


class HastingsMcLeodPainleveIIValues(numberdb.Generator):
    table = TABLE
    parameters = ("s",)
    type = "R"
    digits = DIGITS
    rigour = "heuristic (agreement-checked)"

    def enumerate(self):
        for argument in arguments():
            yield {"s": str(argument)}

    def value(self, params, digits):
        if digits > DIGITS:
            raise ValueError("This calibration supports at most 100 significant digits")
        values, _summary = checked_result()
        return {"number": values[str(parse_fraction(params["s"]))], "digits": DIGITS}

    def environment(self):
        return {**super().environment(), "python-flint": flint.__version__}


def offline_entries():
    generator = HastingsMcLeodPainleveIIValues()
    return [{"params": params, **generator.value(params, DIGITS)} for params in generator.enumerate()]


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--sample", nargs="*", help="run both Taylor configurations only at these s values")
    parser.add_argument("--compute-json", type=Path)
    parser.add_argument("--summary-json", type=Path)
    parser.add_argument("--check-fredholm", action="store_true")
    parser.add_argument("--fredholm-json", type=Path)
    parser.add_argument("--fredholm-arguments", nargs="*", default=list(DEFAULT_FREDHOLM_ARGUMENTS))
    parser.add_argument("--publish", action="store_true")
    options = parser.parse_args()

    if options.sample is not None:
        values, summary = compare_configurations(selected=options.sample)
        print(json.dumps({"values": values, "summary": summary}, indent=2))
        return

    generator = HastingsMcLeodPainleveIIValues()
    if options.publish:
        print(generator.publish(
            message="Refine Hastings-McLeod Painleve-II q(s) values to 100 significant digits",
            assisted_by=os.environ.get("NUMBERDB_ASSISTED_BY", ""),
        ))
        return

    values = None
    summary = None
    needs_full_taylor = (
        options.compute_json
        or options.summary_json
        or (not options.check_fredholm and not options.fredholm_json)
    )
    if needs_full_taylor:
        values, summary = checked_result()
    if options.compute_json:
        records = offline_entries()
        options.compute_json.write_text(json.dumps(records, indent=2) + "\n")
        print(f"Computed {len(records)} entries at {DIGITS} significant digits")
    if options.summary_json:
        options.summary_json.write_text(json.dumps(summary, indent=2) + "\n")
        print(f"Wrote Taylor summary to {options.summary_json}")
    if options.check_fredholm or options.fredholm_json:
        fredholm = check_fredholm(options.fredholm_arguments, reference_values=values)
        fredholm_record = {"taylor": summary, "fredholm": fredholm}
        if options.fredholm_json:
            options.fredholm_json.write_text(json.dumps(fredholm_record, indent=2) + "\n")
            print(f"Wrote Fredholm controls to {options.fredholm_json}")
        if options.check_fredholm:
            print(json.dumps(fredholm_record, indent=2))
    if values is not None and not options.compute_json and not options.summary_json and not options.check_fredholm and not options.fredholm_json:
        print(json.dumps(summary, indent=2))
        print(f"Offline check completed for {len(values)} entries at {DIGITS} significant digits")


if __name__ == "__main__":
    main()
