"""Reproduce the 165 Tracy-Widom densities in NumberDB T451 at 100 digits.

    python -m pip install numberdb python-flint
    python generate.py --compute-json entries.json
    python generate.py --check-fredholm
    python generate.py
    python generate.py --publish

The last two commands verify against / publish to the website. Publishing
requires NUMBERDB_API_KEY; identify assisted runs with NUMBERDB_ASSISTED_BY.
No SageMath installation is required. Tested with python-flint 0.6.0.

The arguments remain beta=1,2,4 and every reduced rational of denominator
at most four in [-6,3]. This file changes precision, not the selection.

Painleve II and its three tail integrals are integrated by Taylor series.
Two runs vary working precision, series order, maximum step and the right
boundary simultaneously: (230 digits, 180 terms, 1/8, 40) and
(260 digits, 220 terms, 1/10, 44). At the right boundary q and q' use Airy
values; the integrals use their Airy tails rather than zero. At every target
the Hamiltonian identity J=q'^2-x*q^2-q^4 is checked to absolute 1e-110.
Every density must agree to relative 1e-110 before 100 digits are written.

The calibration's worst relative difference was below 7.1e-145 at beta=4,
s=-6. Independent Fredholm determinant derivatives check all three laws at
zero and beta=4 at both range endpoints. --check-fredholm repeats them.

Arb is used for efficient multiprecision arithmetic, but step results are
replaced by their midpoints. Neither the Taylor remainder nor the nonlinear
right-boundary error is enclosed. Agreement after changing all integration
controls is evidence, not proof: the rigour is heuristic (agreement-checked).
This must not be relabelled proven merely because Arb is used internally.
"""

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
CONFIGURATIONS = ((230, 180, 8, 40), (260, 220, 10, 44))


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


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-190'), rel_tol=arb('1e-170'))
    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-185'):
        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):
    working_digits, order, step_denominator, boundary = configuration
    old_precision, old_cap = ctx.prec, ctx.cap
    try:
        ctx.dps, ctx.cap = working_digits, order
        targets = []
        for argument in arguments():
            point = arb(argument.numerator) / argument.denominator
            targets.extend([(point, 'plain', str(argument)),
                            (point * arb(2).sqrt(), 'sqrt2', str(argument))])
        targets.sort(key=lambda target: target[0], reverse=True)
        position = arb(boundary)
        state = initial_state(position)
        maximum_step = arb(1) / step_denominator
        samples = {}
        for target, scaling, 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
            if not abs(identity - integral) < arb('1e-110'):
                raise ArithmeticError(f'Hamiltonian control failed at {scaling} {argument}')
            factor = (-moment / 2).exp()
            if scaling == 'plain':
                densities = {'1': factor * (-tail / 2).exp() * (integral + value) / 2,
                             '2': factor**2 * integral}
            else:
                densities = {'4': factor * (integral * (tail / 2).cosh() - value * (tail / 2).sinh()) / arb(2).sqrt()}
            for beta, density in densities.items():
                if not density.is_finite() or not density > 0:
                    raise ArithmeticError(f'Invalid density at {beta} {argument}')
                samples[beta + ',' + argument] = density.mid().str(working_digits - 5, radius=False)
        return samples
    finally:
        ctx.prec, ctx.cap = old_precision, old_cap


@functools.lru_cache(maxsize=1)
def checked_values():
    first, second = [compute_configuration(configuration) for configuration in CONFIGURATIONS]
    old_precision = ctx.prec
    try:
        ctx.dps = 280
        if first.keys() != second.keys():
            raise ArithmeticError('The two runs cover different arguments')
        for identity, text in second.items():
            value = arb(text)
            if not abs((arb(first[identity]) - value) / value) < arb('1e-110'):
                raise ArithmeticError(f'Insufficient numerical agreement at {identity}')
        with localcontext() as decimal_context:
            decimal_context.prec = 280
            return {identity: format(Decimal(text), '.99e') for identity, text in second.items()}
    finally:
        ctx.prec = old_precision


def fredholm_densities(argument, nodes=280, length=48):
    quadrature = [arb.legendre_p_root(nodes, index, weight=True) for index in range(nodes)]
    abscissas = [length * (point + 1) / 2 for point, weight in quadrature]
    roots = [(length * weight / 2).sqrt() for point, weight in quadrature]
    kernel, derivative = arb_mat(nodes, nodes), arb_mat(nodes, nodes)
    for row in range(nodes):
        for column in range(row, nodes):
            value, slope, _, _ = (abscissas[row] + abscissas[column] + argument).airy()
            weight = roots[row] * roots[column]
            kernel[row, column] = kernel[column, row] = weight * value
            derivative[row, column] = derivative[column, row] = weight * slope
    determinants, slopes = [], []
    for sign in (-1, 1):
        matrix = sign * kernel
        for index in range(nodes):
            matrix[index, index] += 1
        determinant = matrix.det()
        determinants.append(determinant)
        slopes.append(determinant * matrix.solve(sign * derivative).trace())
    minus, plus = determinants
    minus_prime, plus_prime = slopes
    return {'1': minus_prime, '2': minus_prime * plus + minus * plus_prime,
            '4': (minus_prime + plus_prime) / arb(2).sqrt()}


def check_fredholm():
    values = checked_values()
    old_precision = ctx.prec
    try:
        ctx.dps = 210
        for argument in (0, -6, 3):
            target = arb(argument) * arb(2).sqrt()
            controls = fredholm_densities(target)
            for beta in (('1', '2', '4') if argument == 0 else ('4',)):
                identity = f'{beta},{argument}'
                value, control = arb(values[identity]), controls[beta]
                if not control.is_finite() or not control > 0:
                    raise ArithmeticError(f'Invalid Fredholm control at {identity}')
                if not abs((control - value) / value) < arb('1e-99'):
                    raise ArithmeticError(f'Fredholm control disagrees at {identity}')
    finally:
        ctx.prec = old_precision


class TracyWidomDensities(numberdb.Generator):
    table = 'T451'
    parameters = ('beta', 's')
    type = 'R'
    digits = DIGITS
    rigour = 'heuristic (agreement-checked)'

    def enumerate(self):
        for beta in ('1', '2', '4'):
            for argument in arguments():
                yield {'beta': beta, 's': str(argument)}

    def value(self, params, digits):
        if digits > DIGITS:
            raise ValueError('This calibration supports at most 100 significant digits')
        identity = str(params['beta']) + ',' + str(Fraction(params['s']))
        return {'number': checked_values()[identity], 'digits': DIGITS}

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


if __name__ == '__main__':
    parser = argparse.ArgumentParser()
    parser.add_argument('--compute-json', type=Path)
    parser.add_argument('--check-fredholm', action='store_true')
    parser.add_argument('--publish', action='store_true')
    options = parser.parse_args()
    generator = TracyWidomDensities()
    if options.check_fredholm:
        check_fredholm()
        print('Independent Fredholm checks passed')
    if options.compute_json:
        records = [{'params': params, **generator.value(params, DIGITS)} for params in generator.enumerate()]
        options.compute_json.write_text(json.dumps(records, indent=2) + '\n')
        print(f'Computed {len(records)} entries at {DIGITS} significant digits')
    elif options.publish:
        print(generator.publish(message='Refine Tracy-Widom densities to 100 significant digits',
                                assisted_by=os.environ.get('NUMBERDB_ASSISTED_BY', '')))
    elif not options.check_fredholm:
        report = generator.verify(sample=None)
        print(report)
        raise SystemExit(0 if report.ok else 1)
