"""Exact replay of every certified constant printed in the paper.

Two-site part: sympy rationals, independent of the workspace producers.
Sixteen-site part: reads the rational boxes written by slope_certificate.py.
Deadline part: reads the validated enclosures written by deadline_validated.py.
Every assertion is an exact rational comparison.  Writes data/paper_checks.json.
"""
import json
from fractions import Fraction as F
from math import comb
from pathlib import Path
import sympy as s


def calculate():
    R = s.Rational
    
    ST = [(0, 0), (0, 1), (0, 2), (1, 0), (1, 1), (2, 0)]   # UU UR RR AU AR AA
    b = R(1, 10)
    out = {}
    
    # ---------------------------------------------------------------- two-site source
    Q = s.zeros(6); D = s.zeros(6); pairs = []
    for i, (a, r) in enumerate(ST):
        u = 2 - a - r
        for dest, rate in (((a + 1, r), u * (R(1, 100) + R(a, 2))), ((a, r + 1), u * (R(1, 100) + R(r, 2))),
                           ((a - 1, r), a * (R(1, 100) + R(r, 4))), ((a, r - 1), r * (R(1, 100) + R(a, 4)))):
            if rate:
                j = ST.index(dest); Q[i, j] += rate; Q[i, i] -= rate
        row = []
        for x in range(a + 1):
            for y in range(r + 1):
                p = R(comb(a, x) * comb(r, y), 2 ** (a + r))
                D[i, ST.index((x, y))] += p; row.append((ST.index((x, y)), ST.index((a - x, r - y)), p))
        pairs.append(row)
    d = s.Matrix([R(1, 100) if a > r else R(3, 10) for a, r in ST])
    Hm = s.diag(*[b + x for x in d]) - Q
    Rv = Hm.inv(); assert min(Rv) >= 0
    joint = lambda z: s.Matrix([sum(p * z[j] * z[k] for j, k, p in row) for row in pairs])
    sq = lambda v: v.applyfunc(lambda x: x * x)
    FJ = lambda z: Rv * (d + b * joint(z)); FI = lambda z: Rv * (d + b * sq(D * z))
    
    qIbar = s.Matrix([R(n, 10 ** 6) for n in (948115, 983103, 987779, 378814, 775706, 335634)])
    qJbar = s.Matrix([R(94, 100), R(982, 1000), R(987, 1000), R(27, 100), R(735, 1000), R(235, 1000)])
    resI = Hm * qIbar - d - b * sq(D * qIbar); resJ = Hm * qJbar - d - b * joint(qJbar)
    assert min(resI) == R(435828959, 40000000000000) and min(resJ) >= 0 and max(qIbar) < R(988, 1000)
    z = s.zeros(6, 1)
    for _ in range(8):
        z = FI(z).applyfunc(lambda x: R(s.floor(x * 10 ** 6), 10 ** 6))
    assert z[5] == R(313291, 10 ** 6)
    w = s.Matrix([R(n, 1000) for n in (2148, 2011, 1997, 4880, 2882, 4895)])
    Jbar = Rv * (2 * b) * s.diag(*list(D * qIbar)) * D
    kappa = max((Jbar * w)[i] / w[i] for i in range(6)); assert kappa < R(796, 1000)
    out["two_site_coarse"] = dict(min_residual_I=str(min(resI)), qI_AA_lower=str(z[5]), kappa=float(kappa))
    
    # tight boxes: 250 downward iterates, resolvent-weighted supersolutions
    lo = {}; up = {}
    for lab, Fm in (("J", FJ), ("I", FI)):
        l = s.zeros(6, 1)
        for _ in range(250):
            l = Fm(l).applyfunc(lambda x: R(s.floor(x * 10 ** 12), 10 ** 12))
        if lab == "I":
            der = Rv * (2 * b) * s.diag(*list(D * l)) * D
        else:
            der = s.zeros(6)
            for i, row in enumerate(pairs):
                for j, k, p in row:
                    der[i, j] += b * p * l[k]; der[i, k] += b * p * l[j]
            der = Rv * der
        v = (s.eye(6) - der).inv() * s.ones(6, 1)
        u = l + s.Matrix([R(s.ceiling(x * 10 ** 6), 10 ** 12) for x in v])
        res = Hm * u - d - b * (sq(D * u) if lab == "I" else joint(u))
        assert min(res) > 0 and max(u) < 1 and all(l[i] <= u[i] for i in range(6))
        lo[lab], up[lab] = l, u
    T = Rv * b
    dl = (T * (sq(D * lo["J"]) - joint(up["J"]))).applyfunc(lambda x: max(0, x))
    du = T * (sq(D * up["J"]) - joint(lo["J"]))
    Bl = T * s.diag(*list(D * (lo["I"] + lo["J"]))) * D; Bu = T * s.diag(*list(D * (up["I"] + up["J"]))) * D
    assert max((Bu * w)[i] / w[i] for i in range(6)) < 1          # positive-vector contraction for B_+
    hl = (s.eye(6) - Bl).inv() * dl; hu = (s.eye(6) - Bu).inv() * du
    assert R(1144144, 10 ** 7) < hl[5] and hu[5] < R(1144331, 10 ** 7)
    assert R(2207216, 10 ** 7) < lo['J'][5] and up['J'][5] < R(2207253, 10 ** 7)
    assert R(3351447, 10 ** 7) < lo['I'][5] and up['I'][5] < R(3351497, 10 ** 7) and max(up['I']) < R(9875805, 10 ** 7)
    assert max(max(up[m][i] - lo[m][i] for i in range(6)) for m in 'JI') < R(5, 10 ** 6)
    assert min(hl) > R(246, 10 ** 5) and max(hl) > R(12347, 10 ** 5) and z[5] - qJbar[5] > R(78, 1000)
    assert all(hl[i] > 0 for i in range(6))                        # strict separation at every state
    out["two_site_tight"] = dict(h_AA=[float(hl[5]), float(hu[5])], qJ_AA=[float(lo["J"][5]), float(up["J"][5])],
                                 qI_AA=[float(lo["I"][5]), float(up["I"][5])], qI_max_upper=float(max(up["I"])),
                                 h_all_lower=[float(x) for x in hl], h_all_upper=[float(x) for x in hu])
    
    return out, lo, up, hl, hu, Q, D, pairs, d, Hm
