"""Proves a lower bound on the area of the Mandelbrot set M, in ball arithmetic.

The theorem. A hyperbolic component W of M of period n has a center c, where 0 is a
periodic point of f_c(z) = z^2 + c of exact period n. Its multiplier map is a conformal
isomorphism from W onto the unit disk (Douady and Hubbard), so the inverse map
    phi(t) = c + a_1 t + a_2 t^2 + ...
is univalent on the disk, and the area theorem gives Area(W) = pi * sum_k k |a_k|^2.
Every partial sum is a lower bound. Distinct hyperbolic components are disjoint subsets
of M, so lower bounds summed over distinct components bound Area(M) from below.

What this program proves, for each component proposed in data/components.txt.gz (a period
n and an approximate center written in decimal):
  1. A closed disk D of radius r holds exactly one root of G_n(c) = f_c^n(0): with
     Y = 1/G_n'(m), the map T(c) = c - Y G_n(c) sends D into itself and is a contraction
     on D, so it has exactly one fixed point there.
  2. G_d has no zero in D for any proper divisor d of n, so 0 has exact period n at that
     root, which is therefore the center of a hyperbolic component of period n.
  3. Enclosures of 1/a_1 = lambda'(c) = 2^n G_n'(c) prod_{0<j<n} f_c^j(0), and for
     larger components of a_1..a_K from the Taylor series of the multiplier, which bound
     the component's area from below.
Proposed centers of one period whose disks overlap might be one center; only the first
of them counts. A component whose disk lies in the open upper half-plane also counts
for its mirror image, which is a distinct component of the same area, unless that image
overlaps a counted disk. Areas are rounded down to multiples of 2^-96 and summed exactly.

Every number comes from python-flint's ball arithmetic (FLINT's arb and acb), whose
results enclose the exact values. Nothing here relies on floating point being right.
"""
import gzip
import json
import math
import os
import sys
from fractions import Fraction
from multiprocessing import Pool

from flint import acb, acb_series, arb, ctx

DATA = os.path.join("data", "components.txt.gz")
RESULTS = os.path.join("results", "R1.json")
GRID_BITS = 96             # each component's bound is rounded down to a multiple of 2^-96
SERIES_TERMS = 8           # a_1..a_8 for components of period <= SERIES_MAX_PERIOD ...
SERIES_MAX_PERIOD = 64
SERIES_MIN_AREA = 1e-9     # ... whose first-term bound exceeds this; others keep pi |a_1|^2


# ---------------------------------------------------------------- ball arithmetic pieces

def box(m, r):
    """The closed square of half-width r around the exact point m; it contains the disk D(m, r)."""
    return acb(arb(m.real, r), arb(m.imag, r))


def orbit(c, n):
    """Enclosures of f_c^j(0) for j = 1..n, and of the derivative of f_c^n(0) in c."""
    z, dz, zs = acb(0), acb(0), []
    for _ in range(n):
        dz = 2 * z * dz + 1
        z = z * z + c
        zs.append(z)
    return zs, dz


def refine(m, n):
    """Newton's method on G_n at midpoints. This only finds a good disk; it proves nothing."""
    tol = arb(2) ** -(ctx.prec - 16)
    for _ in range(60):
        zs, dz = orbit(m, n)
        if dz.contains(0):
            return None
        step = (zs[-1] / dz).mid()
        m = (m - step).mid()
        if abs(step).upper() <= tol * (1 + abs(m).upper()):
            break
    return m


def certify(m, n, r):
    """If the disk D(m, r) holds exactly one center of exact period n, the orbit over its square."""
    zs_m, dg_m = orbit(m, n)
    if dg_m.contains(0):
        return None
    y = (1 / dg_m).mid()                       # an exact complex number
    b = box(m, r)
    zs_b, dg_b = orbit(b, n)
    lip = abs(1 - y * dg_b).upper()            # |T'| <= lip on the square, so on D(m, r)
    if not lip < 1:
        return None
    # |T(c) - m| <= |T(m) - m| + lip |c - m| <= |Y G(m)| + lip r for every c in D(m, r)
    if not abs(y * zs_m[-1]).upper() + lip * r < r:
        return None
    for d in range(1, n):
        if n % d == 0 and zs_b[d - 1].contains(0):
            return None
    return b, zs_b, dg_b


def first_term_area(n, zs_b, dg_b):
    """pi |a_1|^2, a lower bound on the component's area.

    |1/a_1| is a product of n moduli, multiplied as real balls: a product of n complex
    rectangles would widen by up to sqrt(2) a factor, as squaring does."""
    modulus = abs(dg_b) * arb(2) ** n
    for z in zs_b[:-1]:
        modulus *= abs(z)
    bound = arb.pi() / modulus ** 2
    return bound if bound.is_finite() else arb(0)


def series_area(b, n, terms):
    """pi * sum_{k <= terms} k |a_k|^2, a lower bound on the component's area, or None."""
    saved = ctx.cap
    ctx.cap = terms + 1
    try:
        c = acb_series([b, 1])                 # the parameter c + t, t a formal variable
        z = acb_series([0])
        for _ in range(terms):                 # z <- f^n(z) fixes one more coefficient of the
            w = z                              # attracting periodic point each pass, since its
            for _ in range(n):                 # multiplier vanishes at the center
                w = w * w + c
            z = w
        lam, w = acb_series([1]), z
        for _ in range(n):                     # the multiplier, prod 2 f^j(z)
            lam = lam * (2 * w)
            w = w * w + c
        co = list(lam.coeffs()) + [acb(0)] * (terms + 1)
        if not co[0].contains(0) or co[1].contains(0):
            return None
        # the multiplier is exactly 0 at the center; invert t -> lambda(t)
        a = list(acb_series([acb(0)] + co[1:terms + 1]).reversion().coeffs()) + [acb(0)] * (terms + 1)
        total = arb(0)
        for k in range(1, terms + 1):
            total += k * abs(a[k]) ** 2
        total *= arb.pi()
        return total if total.is_finite() else None
    finally:
        ctx.cap = saved


def prove(n, re, im, terms=1):
    """Proves one proposed component: (square around its center, area lower bound), or None.

    Arb's complex balls are rectangles, and squaring one can widen it by up to sqrt(2)
    beyond the true image, so along s steps of an orbit radii can grow about 2^(s/2) more
    than derivatives alone would make them. The precision therefore grows with the number
    of steps, and the disk's radius is 2^-(precision/2), far below any component's size."""
    steps = n * (terms + 1) if terms > 1 else n
    base = 128 + 64 * -(-steps // 64)
    for prec in (base, 2 * base):
        ctx.prec = prec
        m = refine(acb(re, im), n)
        if m is None:
            continue
        r = arb(2) ** -(prec // 2)
        got = certify(m, n, r)
        if got is None:
            continue
        b, zs_b, dg_b = got
        area = first_term_area(n, zs_b, dg_b)
        if terms > 1:
            s = series_area(b, n, terms)
            if s is not None and s.lower() > area.lower():
                area = s
        return b, area
    return None


def terms_for(n, first_area):
    """How many Taylor coefficients to enclose. The first term already holds nearly all of a
    small or round component's area; long series over long orbits need too much precision."""
    if n <= SERIES_MAX_PERIOD and first_area > SERIES_MIN_AREA:
        return SERIES_TERMS
    return 1


def floor_grid(x):
    """The largest multiple of 2^-GRID_BITS not above the lower end of the ball x, as an integer."""
    lo = x.lower()
    if not lo > 0:
        return 0
    m, e = lo.man_exp()
    m, e = int(m), int(e) + GRID_BITS
    return m << e if e >= 0 else m >> -e     # >> rounds toward -infinity


def outward(x, direction):
    """The exact endpoint x as a double, rounded away from the box (direction -1 or +1).

    Disks are far smaller than doubles can resolve, so the squares compared for overlaps are
    widened to doubles. That can only make distinct components look like one, which drops
    area from the bound; it can never make one component count twice."""
    m, e = x.man_exp()
    value = Fraction(int(m)) * (Fraction(2) ** int(e))
    return math.nextafter(float(value), direction * math.inf)


def prove_chunk(chunk):
    out = []
    for index, n, re, im in chunk:
        got = prove(n, re, im, 1)
        if got is not None:
            terms = terms_for(n, float(got[1].lower()))
            if terms > 1:
                better = prove(n, re, im, terms)
                if better is not None and better[1].lower() > got[1].lower():
                    got = better
        if got is None:
            out.append((index, n, None, None))
            continue
        b, area = got
        edges = (outward(b.real.lower(), -1), outward(b.real.upper(), 1),
                 outward(b.imag.lower(), -1), outward(b.imag.upper(), 1))
        out.append((index, n, edges, floor_grid(area)))
    return out


# ---------------------------------------------------------------- the whole proof

def read_components(path):
    rows = []
    with gzip.open(path, "rt") as f:
        for line in f:
            line = line.strip()
            if not line or line.startswith("#"):
                continue
            period, re, im = line.split()[:3]
            rows.append((len(rows), int(period), re, im))
    return rows


def overlapping(boxes):
    """Indexes of boxes (period, re_lo, re_hi, im_lo, im_hi, ...) that meet an earlier box of the same period."""
    order = sorted(range(len(boxes)), key=lambda i: (boxes[i][0], boxes[i][1], i))
    dropped, active = set(), []
    for i in order:
        n, re_lo, re_hi, im_lo, im_hi = boxes[i][:5]
        active = [j for j in active if boxes[j][0] == n and boxes[j][2] >= re_lo]
        if any(boxes[j][3] <= im_hi and im_lo <= boxes[j][4] for j in active):
            dropped.add(i)
            continue
        active.append(i)
    return dropped


def main():
    rows = read_components(DATA)
    size = 500
    chunks = [rows[i:i + size] for i in range(0, len(rows), size)]
    with Pool(os.cpu_count() or 1) as pool:
        proven = [r for part in pool.imap(prove_chunk, chunks) for r in part]
    proven.sort()

    boxes, failed = [], 0
    for index, n, edges, area in proven:
        if edges is None:
            failed += 1
            continue
        re_lo, re_hi, im_lo, im_hi = edges
        if im_hi < 0:
            failed += 1                        # the data lists upper half-plane centers only
            continue
        boxes.append((n, re_lo, re_hi, im_lo, im_hi, area, False))
        if im_lo > 0:                          # its mirror image, a distinct component
            boxes.append((n, re_lo, re_hi, -im_hi, -im_lo, area, True))
    dropped = overlapping(boxes)
    kept = [b for i, b in enumerate(boxes) if i not in dropped]

    total = sum(b[5] for b in kept)            # an integer number of 2^-GRID_BITS units
    bound = Fraction(total, 2 ** GRID_BITS)
    by_period = {}
    for b in kept:
        by_period.setdefault(b[0], 0)
        by_period[b[0]] += b[5]

    def down(x, digits):
        return math.floor(x * 10 ** digits) / 10 ** digits

    results = {
        "components_listed": len(rows),
        "components_proven": len(rows) - failed,
        "components_failed": failed,
        "disks_overlapping_an_earlier_one": sum(1 for i in dropped if not boxes[i][6]),
        "mirror_images_counted": sum(1 for b in kept if b[6]),
        "components_counted": len(kept),
        "largest_period": max(b[0] for b in kept),
        "area_lower_bound": down(bound, 10),
        "area_lower_bound_numerator": str(total),
        "area_lower_bound_denominator_log2": GRID_BITS,
        "periods_1_to_2_area_lower_bound": down(Fraction(sum(v for k, v in by_period.items() if k <= 2), 2 ** GRID_BITS), 10),
        "periods_3_and_up_area_lower_bound": down(Fraction(sum(v for k, v in by_period.items() if k >= 3), 2 ** GRID_BITS), 10),
    }
    os.makedirs("results", exist_ok=True)
    with open(RESULTS, "w") as f:
        json.dump(results, f, indent=2)
        f.write("\n")
    print(json.dumps(results, indent=2))
    return 0 if failed == 0 else 1


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