"""Values of the Buchstab function omega(u) -- numberdb.org/T434

The Buchstab function is defined by omega(u) = 1/u on 1 <= u <= 2
and (u omega(u))' = omega(u - 1) for u > 2.  This table stores values
at every two-decimal argument 1.00 <= u <= 8.00.

Run it with SageMath:

    $ sage -pip install numberdb          # once
    $ sage -python generate.py            # check the table against this code
    $ sage -python generate.py --publish  # fill the draft, with NUMBERDB_API_KEY set

The computation uses h(u) = u omega(u).  On each unit interval h is carried
as a midpoint Taylor series with Sage real-ball coefficients, and the tail of
the product by 1/(u - 1) is bounded by a geometric majorant.
"""

import os
import sys

import numberdb.sage as numberdb
from sage.rings.rational_field import QQ
from sage.rings.real_arb import RealBallField


TABLE = os.environ.get("NUMBERDB_TABLE") or "T434"
MIN_CENTS = 100
MAX_CENTS = 800
TAYLOR_DEGREE = 320
WORKING_GUARD = 128


def _key_from_stdin():
    if os.environ.get("NUMBERDB_KEY_FROM_STDIN") != "1":
        return
    token = sys.stdin.read().strip()
    if "=" in token and token.split("=", 1)[0].isupper():
        token = token.split("=", 1)[1].strip().strip("'\"")
    if token:
        os.environ["NUMBERDB_API_KEY"] = token


def _field(digits):
    return RealBallField(numberdb.bits(digits, losing=WORKING_GUARD))


def _eval_poly(coeffs, y):
    total = coeffs[-1].parent()(0)
    for coeff in reversed(coeffs):
        total = total * y + coeff
    return total


def _upper_abs(x):
    return x.parent()(abs(x).upper())


def _with_error(x, radius):
    if radius <= 0:
        return x
    return x.add_error(radius)


class BuchstabTaylorModel:
    def __init__(self, digits, degree=TAYLOR_DEGREE):
        self.field = _field(digits)
        self.degree = degree
        self.coeffs = {}
        self.errors = {}
        self.omega_coeffs = {}
        self.omega_errors = {}
        self._build()

    def _build(self):
        r = self.field(QQ(1) / QQ(2))
        zero = self.field(0)
        one = self.field(1)

        self.coeffs[1] = [one] + [zero] * self.degree
        self.errors[1] = zero

        for k in range(2, 8):
            previous = self.coeffs[k - 1]
            previous_error = self.errors[k - 1]
            a = self.field(QQ(2 * k - 1) / QQ(2))

            derivative = []
            prior = zero
            for n in range(self.degree):
                coeff = (previous[n] - prior) / a
                derivative.append(coeff)
                prior = coeff

            poly_tail = self._product_tail(previous, a, self.degree)
            inherited_tail = previous_error / (a - r)
            omega_error = poly_tail + inherited_tail
            self.omega_coeffs[k - 1] = derivative
            self.omega_errors[k - 1] = omega_error

            current = [zero] * (self.degree + 1)
            for n, coeff in enumerate(derivative):
                current[n + 1] = coeff / self.field(n + 1)

            left_value = self.h(k - 1, r)
            correction = zero
            minus_r = -r
            power = one
            for n in range(1, self.degree + 1):
                power *= minus_r
                correction += current[n] * power
            current[0] = left_value - correction

            self.coeffs[k] = current
            self.errors[k] = previous_error + omega_error

    def _product_tail(self, coeffs, a, omitted_from):
        r = self.field(QQ(1) / QQ(2))
        q = r / a
        factor = self.field(1) / (a * (self.field(1) - q))
        tail = self.field(0)
        r_power = self.field(1)
        for i, coeff in enumerate(coeffs):
            if i:
                r_power *= r
            first_omitted = max(0, omitted_from - i)
            tail += _upper_abs(coeff) * r_power * factor * (q ** first_omitted)
        return tail

    def interval_for(self, u):
        if not (QQ(1) <= u <= QQ(8)):
            raise ValueError("u is outside the computed range: %s" % (u,))
        if u <= QQ(2):
            return 1
        if u.denominator() == 1:
            return int(u) - 1
        return int(u.floor())

    def h(self, interval, y):
        value = _eval_poly(self.coeffs[interval], self.field(y))
        return _with_error(value, self.errors[interval])

    def omega(self, u):
        u = QQ(u)
        if u <= QQ(2):
            return QQ(1) / u
        interval = self.interval_for(u)
        center = QQ(2 * interval + 1) / QQ(2)
        y = u - center
        return self.h(interval, y) / self.field(u)

    def integral_omega_to(self, x):
        x = QQ(x)
        if x <= QQ(1):
            return self.field(0)
        total = self.field(0)
        for interval in range(1, int(x.floor()) + 1):
            if interval >= 7:
                break
            left = QQ(interval)
            right = min(x, QQ(interval + 1))
            if right <= left:
                continue
            total += self._integral_omega_on_interval(interval, left, right)
            if right == x:
                break
        return total

    def _integral_omega_on_interval(self, interval, left, right):
        center = QQ(2 * interval + 1) / QQ(2)
        y0 = self.field(left - center)
        y1 = self.field(right - center)
        total = self.field(0)
        for n, coeff in enumerate(self.omega_coeffs[interval]):
            total += coeff * (y1 ** (n + 1) - y0 ** (n + 1)) / self.field(n + 1)
        length = self.field(right - left)
        return _with_error(total, self.omega_errors[interval] * length)


class BuchstabFunctionValues(numberdb.Generator):
    table = TABLE
    parameters = ("u",)
    type = "R"
    digits = 100
    rigour = "proven"

    def __init__(self):
        self._models = {}

    def enumerate(self):
        for cents in range(MIN_CENTS, MAX_CENTS + 1):
            yield {"u": str(QQ(cents) / QQ(100))}

    def _model(self, digits):
        if digits not in self._models:
            self._models[digits] = BuchstabTaylorModel(digits)
        return self._models[digits]

    def value(self, params, digits):
        return self._model(digits).omega(QQ(params["u"]))


def _as_ball(field, value):
    if callable(getattr(value, "parent", None)) and "Ball" in type(value).__name__:
        return value
    return field(value)


def self_check(digits=100):
    generator = BuchstabFunctionValues()
    model = generator._model(digits)
    field = model.field

    for cents in range(100, 201):
        u = QQ(cents) / QQ(100)
        got = generator.value({"u": str(u)}, digits)
        expected = QQ(1) / u
        if got != expected:
            raise AssertionError("initial interval failed at u=%s" % (u,))

    for cents in range(201, 301):
        u = QQ(cents) / QQ(100)
        got = _as_ball(field, generator.value({"u": str(u)}, digits))
        expected = (field(1) + field(u - 1).log()) / field(u)
        if not (got - expected).contains_zero():
            raise AssertionError("closed form failed at u=%s" % (u,))

    for cents in range(200, MAX_CENTS + 1):
        u = QQ(cents) / QQ(100)
        got_h = _as_ball(field, generator.value({"u": str(u)}, digits)) * field(u)
        expected_h = field(1) + model.integral_omega_to(u - 1)
        if not (got_h - expected_h).contains_zero():
            raise AssertionError("integral equation failed at u=%s" % (u,))

    omega_8 = _as_ball(field, generator.value({"u": "8"}, digits))
    limit = (-field.euler_constant()).exp()
    difference = omega_8 - limit
    if not (difference.upper() < 0
            and abs(difference).upper() < field("1e-9").upper()):
        raise AssertionError("limit check failed: omega(8)-e^-gamma = %s" % difference)

    print("self-check passed for Buchstab values")
    print("omega(8) - e^-gamma = %s" % difference)


if __name__ == "__main__":
    _key_from_stdin()
    generator = BuchstabFunctionValues()

    if os.environ.get("NUMBERDB_SELF_CHECK") == "1":
        self_check(generator.digits)
    elif os.environ.get("NUMBERDB_PUBLISH") == "1" or "--publish" in sys.argv:
        print(generator.publish(message="Buchstab function values in ball arithmetic"))
    else:
        report = generator.verify(sample=None)
        print(report)
        sys.exit(0 if report.ok else 1)
