"""Checks verify.py against values found another way, and that it refuses what it should.

Independent areas come from the boundary of each component: the points c where the
attracting cycle's multiplier is e^(i theta) are found by Newton's method in two
unknowns (c and a point z of the cycle), with mpmath at 40 digits, and the enclosed area
is (1/2) Im of the integral of conj(c) dc around them, by the trapezoidal rule. That
shares nothing with verify.py's Taylor series except the definition of a component.
"""
import math
import sys

import mpmath as mp
from flint import acb, arb

import verify

mp.mp.dps = 40


def boundary_area(center, n, points=256):
    """Area of the hyperbolic component with this center, from its boundary curve.

    The boundary is c(t) where the attracting cycle's multiplier is e^(i t). Each point is
    found by Newton's method on (f_c^n(z) - z, (f_c^n)'(z) - multiplier) in (c, z),
    following the multiplier from 0 out to the circle and then around it, at the midpoints
    t = 2 pi (j + 1/2) / points, which avoid the root at t = 0. The area is
    (1/2) Im of the integral of conj(c) c'(t) dt, with c'(t) from the implicit function
    theorem, by the trapezoidal rule, which converges geometrically for this periodic
    analytic integrand."""

    def solve(c, z, target):
        for _ in range(80):
            w, dw_dz, dw_dc = z, mp.mpc(1), mp.mpc(0)
            d, dd_dz, dd_dc = mp.mpc(1), mp.mpc(0), mp.mpc(0)
            for _ in range(n):
                d, dd_dz, dd_dc = 2 * w * d, 2 * (dw_dz * d + w * dd_dz), 2 * (dw_dc * d + w * dd_dc)
                w, dw_dz, dw_dc = w * w + c, 2 * w * dw_dz, 2 * w * dw_dc + 1
            f1, f2 = w - z, d - target
            a11, a12, a21, a22 = dw_dc, dw_dz - 1, dd_dc, dd_dz
            det = a11 * a22 - a12 * a21
            dc = (f1 * a22 - a12 * f2) / det
            dz = (a11 * f2 - a21 * f1) / det
            c, z = c - dc, z - dz
            if abs(dc) + abs(dz) < mp.mpf(10) ** -36:
                break
        # dc/d(multiplier): the Jacobian applied to the multiplier's own derivative (0, -1)
        return c, z, -a12 / det

    c, z = mp.mpc(center), mp.mpc(0)
    first = mp.expjpi(mp.mpf(1) / points)
    for step in range(1, 33):                       # walk the multiplier out to the circle
        c, z, _ = solve(c, z, first * mp.mpf(step) / 32)
    total = mp.mpf(0)
    for j in range(points):
        t = 2 * mp.pi * (j + mp.mpf(1) / 2) / points
        mult = mp.expj(t)
        c, z, dc_dmult = solve(c, z, mult)
        total += mp.im(mp.conj(c) * dc_dmult * 1j * mult)
    return total / 2 * (2 * mp.pi / points)


def newton_center(re, im, n):
    """The center near (re, im), by Newton's method in mpmath, independently of verify.py."""
    c = mp.mpc(re, im)
    for _ in range(100):
        z, dz = mp.mpc(0), mp.mpc(0)
        for _ in range(n):
            dz = 2 * z * dz + 1
            z = z * z + c
        c -= z / dz
    return c


def check(name, n, re, im, terms, exact=None):
    got = verify.prove(n, re, im, terms)
    assert got is not None, f"{name} was not proven"
    box, area = got
    center = newton_center(re, im, n)
    inside = box.real.contains(arb(mp.nstr(center.real, 45))) and box.imag.contains(arb(mp.nstr(center.imag, 45)))
    assert inside, f"{name}: the proven square misses the center Newton's method finds"
    reference = exact if exact is not None else boundary_area(mp.mpc(re, im), n)
    low = mp.mpf(area.lower().str(40, radius=False))
    assert low <= reference * (1 + mp.mpf(10) ** -30), f"{name}: {low} exceeds {reference}"
    gap = (reference - low) / reference
    print(f"{name:52s} n={n}  K={terms}  bound {mp.nstr(low, 16):>22s}  reference {mp.nstr(reference, 16):>22s}  gap {mp.nstr(gap, 3)}")
    return gap


def main():
    pi = mp.pi
    # exact areas: the main cardioid phi(t) = t/2 - t^2/4, the period-2 disk phi(t) = -1 + t/4
    assert check("main cardioid", 1, "0", "0", 8, 3 * pi / 8) < mp.mpf(10) ** -20
    assert abs(check("main cardioid, first term only", 1, "0", "0", 1, 3 * pi / 8) - mp.mpf(1) / 3) < mp.mpf(10) ** -20
    assert check("period-2 disk", 2, "-1", "0", 8, pi / 16) < mp.mpf(10) ** -20
    assert check("period-2 disk, first term only", 2, "-1", "0", 1, pi / 16) < mp.mpf(10) ** -20

    # components checked against their boundary curves: (name, period, center, round)
    cases = [
        ("1/3 bulb of the cardioid", 3, "-0.12256116687665362", "0.74486176661974424", True),
        ("period-3 cardioid on the real axis", 3, "-1.7548776662466927", "0", False),
        ("1/4 bulb of the cardioid", 4, "0.28227139076691387", "0.53006061757852531", True),
        ("1/2 bulb of the 1/3 bulb", 6, "-0.11341865594172322", "0.86056947245544906", True),
        ("period-4 cardioid off the axis", 4, "-0.15652016683375508", "1.0322471089228318", False),
        ("period-5 cardioid", 5, "-1.9854115692183978", "0", False),
    ]
    for name, n, re, im, round_ in cases:
        assert check(name, n, re, im, 8) < mp.mpf(10) ** -12
        gap = check(name + ", first term only", n, re, im, 1)
        # a round satellite loses well under 1% to the first term, a cardioid about a third
        assert (gap < mp.mpf("0.01")) if round_ else (mp.mpf("0.25") < gap < mp.mpf("0.4")), gap

    # what must be refused: a center under a multiple of its period
    assert verify.prove(6, "-1.7548776662466927", "0", 1) is None, "period 6 accepted for a period-3 center"
    assert verify.prove(4, "-1", "0", 1) is None, "period 4 accepted for the period-2 center"
    assert verify.prove(2, "0", "0", 1) is None, "period 2 accepted for the period-1 center"

    # rounding: bounds go down onto the 2^-96 grid, squares widen outward
    third = arb(1) / 3
    assert verify.floor_grid(third) == (2 ** verify.GRID_BITS) // 3
    assert verify.floor_grid(arb(0)) == 0
    assert verify.outward(arb(1).lower(), -1) < 1 < verify.outward(arb(1).lower(), 1)

    # the overlap rule keeps the first of the squares of one period that meet
    boxes = [(3, 0, 2, 0, 2), (3, 1, 3, 1, 3), (4, 1, 3, 1, 3), (3, 5, 6, 5, 6), (3, 2.5, 4, 2.5, 4)]
    assert verify.overlapping(boxes) == {1}, verify.overlapping(boxes)
    print("all checks passed")
    return 0


if __name__ == "__main__":
    sys.exit(main())
