"""Generate the first-return lemmas of proofs/SB3.lean (the `trace_*` and `p*` lemmas).

The construction is the bump rule `bump` below (the same rule as `Bump` in the Lean file). For
each region of the section {0 y z} this script runs the successor map on sample points, splits
each orbit segment into single steps and "runs" (three steps that add a constant 0/1 vector),
fits the start of every run and its number of rounds as integer affine forms in m/2, y and z,
checks the fitted description on every point of the region for many m, and prints it as a Lean
lemma whose proof steps through the same segment. Lean checks the lemmas; nothing here is part
of the proof.

Usage: python3 code/gen_traces.py > generated.lean   (Python 3.10 or later, no packages)
"""
import collections
from fractions import Fraction


# The construction ------------------------------------------------------------------------

def bump(m, y, z):
    """Whether the arc from x y z bumps x (to x+1 mod m) rather than saves it. k = m//2 + 1."""
    k = m // 2 + 1
    if m % 2 == 1 and m >= 7 and z == 1 and y >= k + 1:
        return 1
    if y >= k and z < k:
        return 0
    if y == 0 and z != 0:
        return 0
    if z == 1 and y >= 2:
        return 0
    return 1


def step(m, v):
    x, y, z = v
    return (y, z, (x + bump(m, y, z)) % m)


# The first-return map on the section {0 y z}, in closed form ------------------------------

def R_even(m, y, z):
    k = m // 2 + 1
    if y == 0: return (0, 1) if z == 0 else (z, 0)
    if z == 0: return (1, y) if y <= k - 1 else (0, (y + 1) % m)
    if 2 <= y <= k - 1 and z == 1: return (y + 1, 1)
    if y >= k and 1 <= z <= k - 1: return (y + 1, z) if y < m - 1 else (0, z + 1)
    if y == 1:
        if 1 <= z <= k - 2: return (m - 1, m - z)
        if k - 1 <= z <= m - 2: return (k - 1, z + 1)
        return (2, 1)
    if 2 <= y <= k - 1:
        if 2 <= z <= k - 2: return (m - y, m - z) if y <= z else (m + 1 - y, m + 1 - z)
        return (m + 1 - z, y)
    return (m - z, y)


def R_odd(m, y, z):
    k = m // 2 + 1
    if y == 0: return (0, 1) if z == 0 else (z, 0)
    if z == 0: return (1, y) if y <= k - 1 else (0, (y + 1) % m)
    if 2 <= y <= k - 1 and z == 1: return (y + 1, 1)
    if y == k and 1 <= z <= k - 1: return (k + 1, z)
    if k + 1 <= y and z == 1:
        if y <= m - 3: return (k, y + 2)
        if y == m - 2: return (m - 1, k)
        return (2, 1)
    if k + 1 <= y and 2 <= z <= k - 1: return (y + 1, z) if y < m - 1 else (0, z + 1)
    if y == 1:
        if 1 <= z <= k - 2: return (m - 1, m - z)
        if z in (k - 1, k): return (k, z + 2)
        return (m + 1 - z, 2)
    if 2 <= y <= k - 1:
        if 2 <= z <= k - 1: return (m - y, m - z) if y <= z else (m + 1 - y, m + 1 - z)
        if z == k: return (k, y)
        return (m + 1 - z, y + 1) if y <= k - 2 else (m - z, k)
    if z == k: return (k - 1, y)
    if y < m - 1: return (m - z, y + 1)
    return (z + 1, 1) if z < m - 1 else (0, 2)


# Segments: single steps and runs ----------------------------------------------------------

def trace(m, y, z):
    v = (0, y, z)
    out = [v]
    while True:
        v = step(m, v)
        out.append(v)
        if v[0] == 0:
            return out


def segment(t):
    segs = []
    p = 0
    T = len(t) - 1

    def inc(a, b):
        return tuple(b[i] - a[i] for i in range(3))

    while p < T:
        if p + 6 <= T:
            d = inc(t[p], t[p + 3])
            if all(x in (0, 1) for x in d) and inc(t[p + 3], t[p + 6]) == d:
                r = 2
                while p + 3 * (r + 1) <= T and inc(t[p + 3 * r], t[p + 3 * (r + 1)]) == d:
                    r += 1
                segs.append(("r", d, r, p))
                p += 3 * r
                continue
        segs.append(("s", p))
        p += 1
    return segs


def structure(m, y, z):
    t = trace(m, y, z)
    out = []
    for s in segment(t):
        if s[0] == "s":
            if out and out[-1][0] == "s":
                out[-1] = ("s", out[-1][1] + 1, out[-1][2])
            else:
                out.append(("s", 1, s[1]))
        else:
            out.append(("r", s[1], s[2], s[3]))
    return t, out


def signature_of(m, y, z):
    return tuple((seg[0], seg[1]) for seg in structure(m, y, z)[1])


# Integer affine forms in (m//2, y, z, 1) ---------------------------------------------------

def feats(m, y, z):
    return (m // 2, y, z, 1)


def solve_affine(rows, vals):
    n = len(rows[0])
    A = [[Fraction(x) for x in r] + [Fraction(v)] for r, v in zip(rows, vals)]
    piv = []
    r = 0
    for c in range(n):
        p = next((i for i in range(r, len(A)) if A[i][c] != 0), None)
        if p is None:
            continue
        A[r], A[p] = A[p], A[r]
        inv = 1 / A[r][c]
        A[r] = [x * inv for x in A[r]]
        for i in range(len(A)):
            if i != r and A[i][c] != 0:
                f = A[i][c]
                A[i] = [a - f * b for a, b in zip(A[i], A[r])]
        piv.append(c)
        r += 1
    for i in range(r, len(A)):
        if A[i][n] != 0:
            return None
    sol = [Fraction(0)] * n
    for i, c in enumerate(piv):
        sol[c] = A[i][n]
    return sol


def fit(samples, getter):
    sol = solve_affine([feats(*s) for s in samples], [getter(s) for s in samples])
    if sol is None or any(x.denominator != 1 for x in sol):
        return None
    sol = [int(x) for x in sol]
    for s in samples:
        if sum(a * b for a, b in zip(sol, feats(*s))) != getter(s):
            return None
    return tuple(sol)


def ev(form, m, y, z):
    return sum(a * b for a, b in zip(form, feats(m, y, z)))


def infer(samples):
    info = {smp: structure(*smp) for smp in samples}
    first = info[samples[0]][1]
    shape = [(seg[0], seg[1]) for seg in first]
    for smp in samples:
        if [(seg[0], seg[1]) for seg in info[smp][1]] != shape:
            return None
    sym = []
    for idx, seg in enumerate(first):
        if seg[0] == "s":
            sym.append(("s", seg[1]))
            continue
        forms = []
        for c in range(3):
            f = fit(samples, lambda smp, c=c: info[smp][0][info[smp][1][idx][3]][c])
            if f is None:
                return None
            forms.append(f)
        rf = fit(samples, lambda smp: info[smp][1][idx][2])
        if rf is None:
            return None
        sym.append(("r", seg[1], tuple(forms), rf))
    return sym


def validate(sym, m, y, z, target):
    v = (0, y, z)
    for seg in sym:
        if seg[0] == "s":
            for _ in range(seg[1]):
                v = step(m, v)
        else:
            _, d, forms, rf = seg
            n = ev(rf, m, y, z)
            if n < 0:
                return False
            S0 = tuple(ev(f, m, y, z) for f in forms)
            if S0 != v:
                return False
            for j in range(n):
                Sj = tuple(S0[i] + d[i] * j for i in range(3))
                if any(c < 0 or c >= m for c in Sj):
                    return False
                if step(m, step(m, step(m, Sj))) != tuple(S0[i] + d[i] * (j + 1) for i in range(3)):
                    return False
            v = tuple(S0[i] + d[i] * n for i in range(3))
    return v == target


# Regions of the section, and the Lean they become ----------------------------------------

EVEN = [
 ("E01", "y = 0 ∧ z = 0", lambda m,K,y,z: y == 0 and z == 0, ["y = 0", "z = 0"], ("0", "1")),
 ("E02", "", lambda m,K,y,z: y == 0 and 1 <= z <= m-1, ["y = 0", "1 ≤ z", "z ≤ m - 1"], ("z", "0")),
 ("E03", "", lambda m,K,y,z: 1 <= y <= K and z == 0, ["1 ≤ y", "y ≤ m / 2", "z = 0"], ("1", "y")),
 ("E04", "", lambda m,K,y,z: K+1 <= y <= m-2 and z == 0, ["m / 2 + 1 ≤ y", "y ≤ m - 2", "z = 0"], ("0", "y + 1")),
 ("E04w", "", lambda m,K,y,z: y == m-1 and z == 0, ["y = m - 1", "z = 0"], ("0", "0")),
 ("E05", "", lambda m,K,y,z: 2 <= y <= K and z == 1, ["2 ≤ y", "y ≤ m / 2", "z = 1"], ("y + 1", "1")),
 ("E06", "", lambda m,K,y,z: K+1 <= y <= m-2 and 1 <= z <= K, ["m / 2 + 1 ≤ y", "y ≤ m - 2", "1 ≤ z", "z ≤ m / 2"], ("y + 1", "z")),
 ("E06w", "", lambda m,K,y,z: y == m-1 and 1 <= z <= K, ["y = m - 1", "1 ≤ z", "z ≤ m / 2"], ("0", "z + 1")),
 ("E07a", "", lambda m,K,y,z: y == 1 and z == 1, ["y = 1", "z = 1"], ("m - 1", "m - 1")),
 ("E07b", "", lambda m,K,y,z: y == 1 and 2 <= z <= K-1, ["y = 1", "2 ≤ z", "z ≤ m / 2 - 1"], ("m - 1", "m - z")),
 ("E08", "", lambda m,K,y,z: y == 1 and K <= z <= m-2, ["y = 1", "m / 2 ≤ z", "z ≤ m - 2"], ("m / 2", "z + 1")),
 ("E09", "", lambda m,K,y,z: y == 1 and z == m-1, ["y = 1", "z = m - 1"], ("2", "1")),
 ("E10", "", lambda m,K,y,z: 2 <= y <= z <= K-1, ["2 ≤ y", "y ≤ z", "z ≤ m / 2 - 1"], ("m - y", "m - z")),
 ("E11", "", lambda m,K,y,z: 2 <= z < y <= K and z <= K-1, ["2 ≤ z", "z < y", "y ≤ m / 2"], ("m + 1 - y", "m + 1 - z")),
 ("E12", "", lambda m,K,y,z: 2 <= y <= K and K <= z <= m-1, ["2 ≤ y", "y ≤ m / 2", "m / 2 ≤ z", "z ≤ m - 1"], ("m + 1 - z", "y")),
 ("E13", "", lambda m,K,y,z: K+1 <= y <= m-1 and K+1 <= z <= m-1, ["m / 2 + 1 ≤ y", "y ≤ m - 1", "m / 2 + 1 ≤ z", "z ≤ m - 1"], ("m - z", "y")),
]
ODD = [
 ("O01", "", lambda m,K,y,z: y == 0 and z == 0, ["y = 0", "z = 0"], ("0", "1")),
 ("O02", "", lambda m,K,y,z: y == 0 and 1 <= z <= m-1, ["y = 0", "1 ≤ z", "z ≤ m - 1"], ("z", "0")),
 ("O03", "", lambda m,K,y,z: 1 <= y <= K and z == 0, ["1 ≤ y", "y ≤ m / 2", "z = 0"], ("1", "y")),
 ("O04", "", lambda m,K,y,z: K+1 <= y <= m-2 and z == 0, ["m / 2 + 1 ≤ y", "y ≤ m - 2", "z = 0"], ("0", "y + 1")),
 ("O04w", "", lambda m,K,y,z: y == m-1 and z == 0, ["y = m - 1", "z = 0"], ("0", "0")),
 ("O05", "", lambda m,K,y,z: 2 <= y <= K and z == 1, ["2 ≤ y", "y ≤ m / 2", "z = 1"], ("y + 1", "1")),
 ("O06", "", lambda m,K,y,z: y == K+1 and 1 <= z <= K, ["y = m / 2 + 1", "1 ≤ z", "z ≤ m / 2"], ("m / 2 + 2", "z")),
 ("O07", "", lambda m,K,y,z: K+2 <= y <= m-3 and z == 1, ["m / 2 + 2 ≤ y", "y ≤ m - 3", "z = 1"], ("m / 2 + 1", "y + 2")),
 ("O08", "", lambda m,K,y,z: y == m-2 and z == 1, ["y = m - 2", "z = 1"], ("m - 1", "m / 2 + 1")),
 ("O09", "", lambda m,K,y,z: y == m-1 and z == 1, ["y = m - 1", "z = 1"], ("2", "1")),
 ("O10", "", lambda m,K,y,z: K+2 <= y <= m-2 and 2 <= z <= K, ["m / 2 + 2 ≤ y", "y ≤ m - 2", "2 ≤ z", "z ≤ m / 2"], ("y + 1", "z")),
 ("O10w", "", lambda m,K,y,z: y == m-1 and 2 <= z <= K, ["y = m - 1", "2 ≤ z", "z ≤ m / 2"], ("0", "z + 1")),
 ("O11a", "", lambda m,K,y,z: y == 1 and z == 1, ["y = 1", "z = 1"], ("m - 1", "m - 1")),
 ("O11b", "", lambda m,K,y,z: y == 1 and 2 <= z <= K-1, ["y = 1", "2 ≤ z", "z ≤ m / 2 - 1"], ("m - 1", "m - z")),
 ("O12", "", lambda m,K,y,z: y == 1 and K <= z <= K+1, ["y = 1", "m / 2 ≤ z", "z ≤ m / 2 + 1"], ("m / 2 + 1", "z + 2")),
 ("O13", "", lambda m,K,y,z: y == 1 and K+2 <= z <= m-1, ["y = 1", "m / 2 + 2 ≤ z", "z ≤ m - 1"], ("m + 1 - z", "2")),
 ("O14", "", lambda m,K,y,z: 2 <= y <= z <= K, ["2 ≤ y", "y ≤ z", "z ≤ m / 2"], ("m - y", "m - z")),
 ("O15", "", lambda m,K,y,z: 2 <= z < y <= K, ["2 ≤ z", "z < y", "y ≤ m / 2"], ("m + 1 - y", "m + 1 - z")),
 ("O16", "", lambda m,K,y,z: 2 <= y <= K and z == K+1, ["2 ≤ y", "y ≤ m / 2", "z = m / 2 + 1"], ("m / 2 + 1", "y")),
 ("O17", "", lambda m,K,y,z: 2 <= y <= K-1 and K+2 <= z <= m-1, ["2 ≤ y", "y ≤ m / 2 - 1", "m / 2 + 2 ≤ z", "z ≤ m - 1"], ("m + 1 - z", "y + 1")),
 ("O18", "", lambda m,K,y,z: y == K and K+2 <= z <= m-1, ["y = m / 2", "m / 2 + 2 ≤ z", "z ≤ m - 1"], ("m - z", "m / 2 + 1")),
 ("O19", "", lambda m,K,y,z: K+1 <= y <= m-1 and z == K+1, ["m / 2 + 1 ≤ y", "y ≤ m - 1", "z = m / 2 + 1"], ("m / 2", "y")),
 ("O20", "", lambda m,K,y,z: K+1 <= y <= m-2 and K+2 <= z <= m-1, ["m / 2 + 1 ≤ y", "y ≤ m - 2", "m / 2 + 2 ≤ z", "z ≤ m - 1"], ("m - z", "y + 1")),
 ("O21", "", lambda m,K,y,z: y == m-1 and K+2 <= z <= m-2, ["y = m - 1", "m / 2 + 2 ≤ z", "z ≤ m - 2"], ("z + 1", "1")),
 ("O22", "", lambda m,K,y,z: y == m-1 and z == m-1, ["y = m - 1", "z = m - 1"], ("0", "2")),
]

def term(coef, name):
    if coef == 1: return name
    return "%d * %s" % (coef, name)

def render(form):
    """Affine form (aK, by, cz, d) -> Lean Nat expression, as (P) - (N) when needed."""
    names = ("(m / 2)", "y", "z")
    pos, neg = [], []
    for c, nm in zip(form[:3], names):
        if c > 0: pos.append(term(c, nm))
        elif c < 0: neg.append(term(-c, nm))
    d = form[3]
    if d > 0: pos.append(str(d))
    elif d < 0: neg.append(str(-d))
    P = " + ".join(pos) if pos else "0"
    if not neg: return P
    return "(%s) - (%s)" % (P, " + ".join(neg))

def gen_class(name, parity, pred, hyps, target, mfit, mchk):
    Rf = R_even if parity == 0 else R_odd
    cells = [(m, y, z) for m in mchk for y in range(m) for z in range(m) if pred(m, m // 2, y, z)]
    fitcells = [c for c in cells if c[0] in mfit]
    bysig = collections.defaultdict(list)
    for c in fitcells: bysig[signature_of(*c)].append(c)
    sig, smp = max(bysig.items(), key=lambda kv: len(kv[1]))
    sym = infer(smp)
    assert sym is not None, name
    bad = [c for c in cells if not validate(sym, c[0], c[1], c[2], (0,) + Rf(*c))]
    assert not bad, (name, bad[:5])
    lines = []
    par = "m % 2 = 0" if parity == 0 else "m % 2 = 1"
    hs = ["(hpar : %s)" % par, "(hm : 12 ≤ m)"] + ["(h%d : %s)" % (i, h) for i, h in enumerate(hyps)]
    lines.append("theorem trace_%s (m y z : Nat) %s :" % (name, " ".join(hs)))
    lines.append("    Reach m (0, y, z) (0, %s, %s) := by" % target)
    for seg in sym:
        if seg[0] == 's':
            lines.append("  " + "; ".join(["sbstep"] * seg[1]))
        else:
            _, d, forms, rf = seg
            comps = []
            for i in range(3):
                b = render(forms[i])
                comps.append("(%s) + j" % b if d[i] == 1 else "(%s)" % b)
            lines.append("  refine Reach.run (fun j => (%s, %s, %s)) (%s) (triple_eq (by omega) (by omega) (by omega)) ?_ ?_" % (comps[0], comps[1], comps[2], render(rf)))
            lines.append("  · intro j hj; (try dsimp only); sbstep; sbstep; sbstep; sbdone")
            lines.append("  try dsimp only")
    lines.append("  sbdone")
    return "\n".join(lines)

def gen_wrappers():
    out = []
    for parity, table in ((0, EVEN), (1, ODD)):
        par = "m % 2 = 0" if parity == 0 else "m % 2 = 1"
        for (name, _, pred, hyps, target) in table:
            conj = " ∧ ".join("(%s)" % h for h in hyps)
            n = len(hyps)
            pat = "⟨" + ", ".join("h%d" % i for i in range(n)) + "⟩" if n > 1 else "h0"
            out.append(
                "theorem p%s (m y z y' z' : Nat) (hpar : %s) (hm : 12 ≤ m) (H : %s)\n"
                "    (e1 : %s = y') (e2 : %s = z') : Reach m (0, y, z) (0, y', z') := by\n"
                "  obtain %s := H\n"
                "  subst e1; subst e2\n"
                "  exact trace_%s m y z hpar hm %s" % (name, par, conj, target[0], target[1], pat, name, " ".join("h%d" % i for i in range(n))))
    return "\n\n".join(out) + "\n"


if __name__ == "__main__":
    out = []
    for (name, _, pred, hyps, target) in EVEN:
        out.append(gen_class(name, 0, pred, hyps, target, [30, 32, 34, 36, 38, 40], list(range(12, 41, 2))))
    for (name, _, pred, hyps, target) in ODD:
        out.append(gen_class(name, 1, pred, hyps, target, [31, 33, 35, 37, 39, 41], list(range(13, 42, 2))))
    print("\n\n".join(out))
    print()
    print(gen_wrappers(), end="")
