#!/usr/bin/env python3
"""
gamma2p.py — KKLT-CERT restriction layer, stage 1 (library + r1 check).

Two-parameter fundamental period of the mirror of X = CP^4_{1,1,1,6,9}[18]
(h21_eff = 2 symmetric locus; the Candelas-Font-Katz-Morrison two-parameter
family) as an explicit Gamma-series from the toric data, plus the literature /
GV cross-checks (check r1) and the sign-convention pin that defines the DKMM
monomial curve.

TORIC DATA (Mori/charge vectors; basis = DKMM effective basis, index 1 =
elliptic-fiber class E [GV 540], index 2 = base class B [GV 3]):
    l^(1) = (-6; 0,0,0, 2,3,1)     [fiber]
    l^(2) = ( 0; 1,1,1, 0,0,-3)    [base]
Gamma-series coefficient (nu = n + rho):
    A(nu) = Gamma(6 nu1 + 1) / [ Gamma(nu2+1)^3 Gamma(2 nu1+1) Gamma(3 nu1+1)
                                 Gamma(nu1 - 3 nu2 + 1) ]
    w(z; rho) = sum_{n1,n2 >= 0} A(n+rho) z1^{n1+rho1} z2^{n2+rho2}
    w0 = w(z;0):  c(n) = (6n1)! / [ (n2!)^3 (2n1)! (3n1)! (n1-3n2)! ],
nonzero iff n1 >= 3n2 (cone sector).  rho-DERIVATIVE periods receive
contributions from the out-of-cone sector n1 < 3n2 (1/Gamma vanishes at
nonpositive integers but its rho-derivatives do not) — handled EXACTLY below
via truncated rho-jets with symbolic (gamma, zeta2, zeta3) coefficients.
gamma-cancellation (CY condition sum_i l_i = 0) is asserted per term.

r1 checks in __main__:
  (c1) fiber sub-series c(n,0) = (6n)!/((3n)!(2n)!n!) = 1,60,13860,4084080,...
       (classic E8-elliptic literature constants, hardcoded).
  (c2) CFKM PF operators annihilate the series — exact recurrence identities
         L1 = th1(th1-3th2) - 12 z1 (6th1+1)(6th1+5)
         L2 = th2^3 - z2 (th1-3th2)(th1-3th2-1)(th1-3th2-2)
       checked for all n in a window (these are the literature PF system).
  (c3) full GV extraction (HKTY double-log, holomorphic limit) through
       degree (6,6); match vs the pinned independent table ([C] 2112.13863
       "CY 39", as pinned in pilot/repro_w0.py) through (4,4); local-P2
       base-degree numbers (0,5)=1695, (0,6)=-17064 vs classic literature.
       This simultaneously pins the sign convention sigma (z_a = sigma_a *
       ztilde_a) and the global epsilon of the double-log normalization.

Resources: pure python, single-thread.
"""
from fractions import Fraction as Fr
from itertools import product
import json, os, sys

TR = 3  # rho-jet truncation (total degree); need <=3 (triple logs)

# ---------------- symbol coefficients: {(g,z2,z3): Fr} = gamma^g zeta2^z2 zeta3^z3
def sadd(a, b):
    out = dict(a)
    for k, v in b.items():
        w = out.get(k, Fr(0)) + v
        if w: out[k] = w
        elif k in out: del out[k]
    return out

def smul(a, b):
    out = {}
    for ka, va in a.items():
        for kb, vb in b.items():
            k = (ka[0]+kb[0], ka[1]+kb[1], ka[2]+kb[2])
            w = out.get(k, Fr(0)) + va*vb
            if w: out[k] = w
            elif k in out: del out[k]
    return out

def sscale(a, r):
    return {k: v*r for k, v in a.items() if v*r != 0}

S_ONE = {(0,0,0): Fr(1)}
S_G   = {(1,0,0): Fr(1)}
S_Z2  = {(0,1,0): Fr(1)}
S_Z3  = {(0,0,1): Fr(1)}

def s_rat(a):
    """rational part; assert no gamma."""
    for k in a:
        assert k[0] == 0, f"gamma survived: {a}"
    return a.get((0,0,0), Fr(0))

# ---------------- rho-jets: {(i,j): symdict}, i+j <= TR
def jmul(A, B):
    out = {}
    for (i, j), ca in A.items():
        for (k, l), cb in B.items():
            if i+k+j+l <= TR:
                key = (i+k, j+l)
                out[key] = sadd(out.get(key, {}), smul(ca, cb))
    return {k: v for k, v in out.items() if v}

def jadd(A, B):
    out = dict(A)
    for k, v in B.items():
        out[k] = sadd(out.get(k, {}), v)
    return {k: v for k, v in out.items() if v}

def jscale(A, r):
    return {k: sscale(v, r) for k, v in A.items() if sscale(v, r)}

def jconst(r):
    return {(0,0): {(0,0,0): Fr(r)}} if r != 0 else {}

def jlin(v, c=Fr(1)):
    """c*(v1 rho1 + v2 rho2)"""
    out = {}
    if v[0]*c: out[(1,0)] = {(0,0,0): Fr(v[0])*c}
    if v[1]*c: out[(0,1)] = {(0,0,0): Fr(v[1])*c}
    return out

def jinv_unit(A):
    """1/A for jet with nonzero pure-rational constant term."""
    c0 = A.get((0,0), {}).get((0,0,0), Fr(0))
    assert c0 != 0 and list(A.get((0,0), {})) == [(0,0,0)], "non-unit jet"
    r = jscale(jadd(A, jconst(-c0)), Fr(1)/c0)   # (A - c0)/c0, zero const
    # 1/A = (1/c0)(1 - r + r^2 - r^3)
    out = jconst(1); pw = jconst(1)
    for k in range(1, TR+1):
        pw = jmul(pw, r)
        out = jadd(out, jscale(pw, Fr((-1)**k)))
    return jscale(out, Fr(1)/c0)

def invgamma1(v):
    """jet of 1/Gamma(1+x), x = v.rho:  exp(g x - z2 x^2/2 + z3 x^3/3 - ...)"""
    x = jlin(v); x2 = jmul(x, x); x3 = jmul(x2, x)
    out = jconst(1)
    out = jadd(out, {k: smul(c, S_G) for k, c in x.items()})
    t2 = sadd(sscale(smul(S_G, S_G), Fr(1,2)), sscale(S_Z2, Fr(-1,2)))
    out = jadd(out, {k: smul(c, t2) for k, c in x2.items()})
    t3 = sadd(sadd(sscale(smul(smul(S_G,S_G),S_G), Fr(1,6)),
                   sscale(smul(S_G,S_Z2), Fr(-1,2))), sscale(S_Z3, Fr(1,3)))
    out = jadd(out, {k: smul(c, t3) for k, c in x3.items()})
    return out

def recip_gamma_tables(v, amin, amax):
    """R[a] = jet of 1/Gamma(1+a+x), x=v.rho, for a in [amin, amax] (incremental)."""
    R = {0: invgamma1(v)}
    x = jlin(v)
    for a in range(0, amax):
        # 1/G(1+(a+1)+x) = 1/G(1+a+x) * 1/(a+1+x)
        inv = jconst(Fr(1, a+1))
        pw = jconst(1)
        for k in range(1, TR+1):
            pw = jmul(pw, x)
            inv = jadd(inv, jscale(pw, Fr((-1)**k, (a+1)**(k+1))))
        R[a+1] = jmul(R[a], inv)
    for a in range(0, amin, -1):
        # Gamma(1+a+x) = (a+x) Gamma(a+x)  =>  1/G(1+(a-1)+x) = (x+a) * 1/G(1+a+x)
        R[a-1] = jmul(R[a], jadd(x, jconst(a)))
    return R

class GammaJets:
    """A(n+rho) jets for the P[1,1,1,6,9] family, exact, both sectors."""
    def __init__(self, n1max, n2max):
        self.R6 = recip_gamma_tables((6,0), 0, 6*n1max)   # 1/G(1+6nu1)
        self.R2 = recip_gamma_tables((2,0), 0, 2*n1max)
        self.R3 = recip_gamma_tables((3,0), 0, 3*n1max)
        self.Rb = recip_gamma_tables((0,1), 0, n2max)     # 1/G(1+nu2)
        self.Rm = recip_gamma_tables((1,-3), -3*n2max, n1max)  # 1/G(1+nu1-3nu2)
    def A(self, n1, n2):
        num = jinv_unit(self.R6[6*n1])                     # Gamma(1+6nu1)
        J = jmul(num, self.R2[2*n1])
        J = jmul(J, self.R3[3*n1])
        b = self.Rb[n2]
        J = jmul(J, jmul(b, jmul(b, b)))
        J = jmul(J, self.Rm[n1-3*n2])
        for k, c in J.items():        # CY gamma-cancellation assert
            for mono in c: assert mono[0] == 0, f"gamma at n=({n1},{n2})"
        return J

def cfrac(n1, n2):
    """cone-sector coefficient c(n) as exact Fraction (0 out of cone)."""
    if n1 < 3*n2: return Fr(0)
    from math import factorial
    return Fr(factorial(6*n1),
              factorial(n2)**3 * factorial(2*n1) * factorial(3*n1)
              * factorial(n1-3*n2))

# ---------------- 2-var truncated series on rectangle (N1,N2), symdict coeffs
def tmul(A, B, N1, N2):
    out = {}
    for (i, j), ca in A.items():
        for (k, l), cb in B.items():
            if i+k <= N1 and j+l <= N2:
                key = (i+k, j+l)
                out[key] = sadd(out.get(key, {}), smul(ca, cb))
    return {k: v for k, v in out.items() if v}

def tadd(A, B):
    out = dict(A)
    for k, v in B.items(): out[k] = sadd(out.get(k, {}), v)
    return {k: v for k, v in out.items() if v}

def tscale(A, r):
    return {k: sscale(v, r) for k, v in A.items() if sscale(v, r)}

def tinv_unit(A, N1, N2):
    c0 = A.get((0,0), {}).get((0,0,0), Fr(0))
    assert c0 == 1 and list(A.get((0,0), {})) == [(0,0,0)], "tinv needs unit"
    r = dict(A); del r[(0,0)]
    out = {(0,0): S_ONE}; pw = {(0,0): S_ONE}
    for k in range(1, N1+N2+1):
        pw = tmul(pw, r, N1, N2)
        if not pw: break
        out = tadd(out, tscale(pw, Fr((-1)**k)))
    return out

def texp(X, N1, N2):
    assert (0,0) not in X
    out = {(0,0): S_ONE}; pw = {(0,0): S_ONE}
    from math import factorial
    for k in range(1, N1+N2+1):
        pw = tmul(pw, X, N1, N2)
        if not pw: break
        out = tadd(out, tscale(pw, Fr(1, factorial(k))))
    return out

def tcompose(F, Z1, Z2, N1, N2):
    """F(z1,z2) with z_a = q_a * G_a(q); Z1,Z2 = the unit series G_a in q."""
    P1 = {0: {(0,0): S_ONE}}; P2 = {0: {(0,0): S_ONE}}
    for i in range(1, N1+1): P1[i] = tmul(P1[i-1], Z1, N1, N2)
    for j in range(1, N2+1): P2[j] = tmul(P2[j-1], Z2, N1, N2)
    out = {}
    for (i, j), c in F.items():
        # z1^i z2^j = q1^i q2^j G1^i G2^j
        base = tmul(P1[i], P2[j], N1-i, N2-j)
        for (k, l), s in base.items():
            key = (i+k, j+l)
            out[key] = sadd(out.get(key, {}), smul(s, c))
    return {k: v for k, v in out.items() if v}

# ---------------- HKTY double-log GV extraction
KAPPA = {(0,0,0): Fr(9), (0,0,1): Fr(3), (0,1,1): Fr(1), (1,1,1): Fr(0)}
def kap(i, j, k):
    return KAPPA[tuple(sorted((i, j, k)))]

# pinned GV table (provenance: pilot/repro_w0.py, source [C] 2112.13863
# Table "CY 39", independent of DKMM; deg<=2 cross-checked vs [L] exactly)
GV_PINNED = {
    (1,0):540, (0,1):3, (2,0):540, (1,1):-1080, (0,2):-6,
    (3,0):540, (2,1):143370, (1,2):2700, (0,3):27,
    (4,0):540, (3,1):204071184, (2,2):-574560, (1,3):-17280, (0,4):-192,
    (4,1):21772947555, (3,2):74810520, (2,3):5051970, (1,4):154440,
    (4,2):-49933059660, (3,3):-913383000, (2,4):-57879900,
    (4,3):224108858700, (3,4):13593850920, (4,4):-2953943334360,
}
# classic local-P2 (base-degree) literature values for the extension window
GV_LOCALP2_LIT = {(0,5): 1695, (0,6): -17064}

def build_2param(N1, N2, sig):
    """All jet series on rectangle, coefficients twisted by sigma^n."""
    GJ = GammaJets(N1, N2)
    W = {}   # W[(i,j)] = series of alpha=(i,j) jet component (i+j<=2 used here)
    for al in [(0,0),(1,0),(0,1),(2,0),(1,1),(0,2)]:
        W[al] = {}
    from math import factorial
    for n1 in range(N1+1):
        for n2 in range(N2+1):
            J = GJ.A(n1, n2)
            tw = Fr(sig[0]**n1 * sig[1]**n2)
            for al in W:
                c = J.get(al)
                if c:
                    fac = factorial(al[0])*factorial(al[1])
                    W[al][(n1,n2)] = sscale(c, tw*fac)
    return W

def gv_extract(N1, N2, sig, eps):
    """returns dict d -> n_d (rationals) via HKTY double-log, holo limit."""
    W = build_2param(N1, N2, sig)
    w0 = W[(0,0)]
    w0inv = tinv_unit(w0, N1, N2)
    S = {1: W[(1,0)], 2: W[(0,1)]}
    Sn = {a: tmul(S[a], w0inv, N1, N2) for a in (1,2)}       # S_a/w0
    SS = {(1,1): W[(2,0)], (1,2): W[(1,1)], (2,1): W[(1,1)], (2,2): W[(0,2)]}
    # P_a = (1/2) kappa_abc [ S_bc/w0 - (S_b/w0)(S_c/w0) ]   (log-free identity)
    P = {}
    for a in (1,2):
        acc = {}
        for b in (1,2):
            for c in (1,2):
                kv = kap(a-1, b-1, c-1)
                if kv == 0: continue
                t = tadd(tmul(SS[(b,c)], w0inv, N1, N2),
                         tscale(tmul(Sn[b], Sn[c], N1, N2), Fr(-1)))
                acc = tadd(acc, tscale(t, kv/2))
        P[a] = acc
    # mirror map: q_a = z_a exp(S_a/w0)  -> invert z_a = q_a G_a(q)
    E = {a: texp(Sn[a], N1, N2) for a in (1,2)}
    G1 = {(0,0): S_ONE}; G2 = {(0,0): S_ONE}
    for _ in range(N1+N2+2):
        E1q = tcompose(E[1], G1, G2, N1, N2)
        E2q = tcompose(E[2], G1, G2, N1, N2)
        G1 = tinv_unit(E1q, N1, N2)
        G2 = tinv_unit(E2q, N1, N2)
    Pq = {a: tcompose(P[a], G1, G2, N1, N2) for a in (1,2)}
    # zeta-part sanity: positive-degree coefficients must be pure rational
    for a in (1,2):
        for d, c in Pq[a].items():
            if d != (0,0):
                for mono, v in c.items():
                    assert mono == (0,0,0), f"zeta survives at q^{d}: {c}"
    # peel multicovers: coef_a(e) = eps * sum_{k: e=k d} n_d (d_a) / k^2
    nGV = {}
    degs = sorted([d for d in Pq[1] if d != (0,0)] +
                  [d for d in Pq[2] if d != (0,0) and d not in Pq[1]],
                  key=lambda d: (d[0]+d[1], d))
    seen = set()
    for e in degs:
        if e in seen: continue
        seen.add(e)
        vals = {}
        for a in (1,2):
            if e[a-1] == 0 and e == (0,0): continue
            coef = s_rat(Pq[a].get(e, {}))
            rest = Fr(0)
            k = 2
            while k*1 <= max(e):
                if e[0] % k == 0 and e[1] % k == 0:
                    d = (e[0]//k, e[1]//k)
                    if d in nGV:
                        rest += nGV[d] * Fr(d[a-1], k**2)
                k += 1
            if e[a-1] != 0:
                vals[a] = (Fr(eps)*coef - rest) / e[a-1]
            else:
                # consistency: coef must equal eps^{-1}*rest with d_a=0
                assert Fr(eps)*coef - rest == 0, f"deg-{e} a={a} inconsistency"
        if len(vals) == 2:
            assert vals[1] == vals[2], f"a=1/2 mismatch at {e}: {vals}"
        nGV[e] = list(vals.values())[0]
    return nGV, (G1, G2)

# ---------------- __main__: r1 check
if __name__ == "__main__":
    here = os.path.dirname(os.path.abspath(__file__))
    print("=== stage 1: 2-parameter Gamma-series, literature/GV cross-checks (r1) ===")
    # (c1) fiber sub-series vs classic E8-elliptic constants
    lit = [1, 60, 13860, 4084080]   # classic E8-elliptic series constants
    ours = [cfrac(n, 0) for n in range(4)]
    assert ours == lit, (ours, lit)
    print("(c1) PASS  c(n,0) = (6n)!/((3n)!(2n)!n!) =", lit, "... (classic)")
    # (c2) CFKM PF operators as exact recurrences over a window
    NCHK = 24
    ok = 0
    for n1 in range(NCHK+1):
        for n2 in range(NCHK//3+1):
            c = cfrac(n1, n2)
            if n1 >= 1:
                lhs = Fr(n1)*(n1-3*n2)*c
                rhs = Fr(12)*(6*n1-5)*(6*n1-1)*cfrac(n1-1, n2)
                assert lhs == rhs, (n1, n2, "L1")
                ok += 1
            if n2 >= 1:
                lhs = Fr(n2)**3 * c
                rhs = Fr((n1-3*n2+1))*(n1-3*n2+2)*(n1-3*n2+3)*cfrac(n1, n2-1)
                assert lhs == rhs, (n1, n2, "L2")
                ok += 1
    print(f"(c2) PASS  L1 = th1(th1-3th2)-12z1(6th1+1)(6th1+5), "
          f"L2 = th2^3 - z2(th1-3th2)(th1-3th2-1)(th1-3th2-2)")
    print(f"           annihilate the series: {ok} exact recurrence identities checked")
    # (c3) GV extraction + sign pin
    N1 = N2 = 6
    hit = None
    for sig in [(1,1), (-1,1), (1,-1), (-1,-1)]:
        for eps in (1, -1):
            try:
                nGV, _ = gv_extract(N1, N2, sig, eps)
            except AssertionError:
                continue
            if (nGV.get((1,0)) == 540 and nGV.get((0,1)) == 3
                    and nGV.get((1,1)) == -1080 and nGV.get((0,2)) == -6):
                hit = (sig, eps, nGV)
                break
        if hit: break
    assert hit, "no sign convention reproduces the pinned GV leaders"
    sig, eps, nGV = hit
    print(f"(c3) sign pin: sigma = {sig}, eps = {eps:+d}  "
          f"(z_a = sigma_a * ztilde_a; P_a q-part = eps * sum n_d d_a Li2)")
    bad = []
    for d, v in GV_PINNED.items():
        got = nGV.get(d)
        if got != v: bad.append((d, v, got))
    assert not bad, f"GV mismatches vs pinned table: {bad}"
    print(f"     PASS  all {len(GV_PINNED)} pinned GV invariants (through (4,4)) "
          f"match EXACTLY ([C] 2112.13863 'CY 39' via pilot)")
    for d, v in GV_LOCALP2_LIT.items():
        got = nGV.get(d)
        flag = "PASS" if got == v else f"FAIL got {got}"
        print(f"     {flag}  local-P2 literature {d}: {v}")
        assert got == v
    # integrality of everything extracted on the (6,6) window
    nonint = [(d, v) for d, v in nGV.items() if v.denominator != 1]
    assert not nonint, nonint
    print(f"     PASS  all {len(nGV)} extracted n_d on (6,6) window are integers")
    extra = {d: int(v) for d, v in sorted(nGV.items()) if d not in GV_PINNED}
    print("     extraction beyond pinned window (ours):",
          {d: v for d, v in list(extra.items())})
    with open(os.path.join(here, "gv_extracted.json"), "w") as f:
        json.dump({"sigma": list(sig), "eps": eps,
                   "n_d": {f"{d[0]},{d[1]}": int(v) for d, v in sorted(nGV.items())}},
                  f, indent=1)
    # the monomial curve definition (documented in README.md):
    # p = (2/5, 3/10), nu = 10 p = (4,3);  z1 = sigma1 s^4, z2 = sigma2 s^3
    with open(os.path.join(here, "curve_def.json"), "w") as f:
        json.dump({"p": ["2/5", "3/10"], "nu": [4, 3], "sigma": list(sig),
                   "eps": eps,
                   "curve": "z1 = sigma1*s^4, z2 = sigma2*s^3",
                   "flat": "T = p*tau; T_s := T1 - T2 = tau/10; q_s = e^{2 pi i tau/10}"},
                  f, indent=1)
    print("r1 CHECK: PASS   (curve_def.json + gv_extracted.json written)")
