#!/usr/bin/env python3
"""Exact finite spin laws, influence budgets, and replayable map certificates.

Only Python's standard library is required. Run with --output PATH.
Finite enumeration checks the displayed models; the accompanying pages prove
universal contraction and the almost-sure CFTP theorem. No random simulation
or floating-point output is used as evidence for exact sampling.
"""
from fractions import Fraction as F
from itertools import product
from collections import Counter
from pathlib import Path
import argparse
import json

CHECKS = Counter()

def check(value, label):
    if not value:
        raise RuntimeError(label)
    CHECKS[label] += 1


def spins(n):
    return list(product((-1, 1), repeat=n))


class Ising:
    """Rational factors r=e^(2J)>0, s=e^(2h)>0 on a finite simple graph."""
    def __init__(self, n, edges, ratios, fields):
        if not isinstance(n, int) or n < 1 or len(fields) != n or len(edges) != len(ratios):
            raise ValueError('nonempty finite graph and matching inputs required')
        if any(not isinstance(i, int) or not isinstance(j, int) or not 0 <= i < j < n for i, j in edges):
            raise ValueError('edges must be ordered distinct vertices')
        if len(set(edges)) != len(edges):
            raise ValueError('parallel edges are not part of this input contract')
        self.n, self.edges = n, list(edges)
        self.r, self.s = list(map(F, ratios)), list(map(F, fields))
        if any(v <= 0 for v in self.r + self.s):
            raise ValueError('strictly positive finite factors required')
        self.states = spins(n)
        self.w = {x: self.weight(x) for x in self.states}
        self.Z = sum(self.w.values())
        self.pi = {x: w/self.Z for x, w in self.w.items()}

    def weight(self, x):
        v = F(1)
        for (i, j), r in zip(self.edges, self.r):
            if x[i] == x[j]:
                v *= r
        for a, s in zip(x, self.s):
            if a == 1:
                v *= s
        return v

    def positive(self, x, i):
        odds = self.s[i]
        for (j, k), r in zip(self.edges, self.r):
            if i == j:
                odds *= r**x[k]
            elif i == k:
                odds *= r**x[j]
        return odds/(1+odds)

    def update(self, x, i, u):
        if not 0 <= i < self.n or not 0 <= u <= 1:
            raise ValueError('invalid update record')
        y = list(x)
        y[i] = 1 if u < self.positive(x, i) else -1
        return tuple(y)

    def push(self, mu, i):
        out = {x: F(0) for x in self.states}
        for x, mass in mu.items():
            p = self.positive(x, i)
            for value, prob in [(1, p), (-1, 1-p)]:
                y = list(x)
                y[i] = value
                out[tuple(y)] += mass*prob
        return out

    def influence(self):
        C = [[F(0) for _ in range(self.n)] for _ in range(self.n)]
        for i in range(self.n):
            for j in range(self.n):
                if i == j:
                    continue
                for x in self.states:
                    y = list(x)
                    y[j] *= -1
                    C[i][j] = max(C[i][j], abs(self.positive(x, i)-self.positive(tuple(y), i)))
        return C


def tv(mu, nu):
    return sum(abs(mu[x]-nu[x]) for x in mu)/2


def apply_budget(C, b, q):
    if len(q) != len(b) or any(v < 0 for v in q) or sum(q) != 1:
        raise ValueError('scan must be a probability vector')
    return [(1-q[i])*b[i]+q[i]*sum(C[i][j]*b[j] for j in range(len(b))) for i in range(len(b))]


def scan_model(model, order, initial):
    C = model.influence()
    b = [F(1)]*model.n
    mu = {x: F(x == initial) for x in model.states}
    for i in order:
        q = [F(j == i) for j in range(model.n)]
        b = apply_budget(C, b, q)
        mu = model.push(mu, i)
    return b, mu


def marginal(mu, indices):
    out = Counter()
    for x, mass in mu.items():
        out[tuple(x[i] for i in indices)] += mass
    return out


def replay(model, records, initial):
    trace = [initial]
    for i, u in records:
        trace.append(model.update(trace[-1], i, u))
    return trace


def tree_and_cycle():
    ans = {}
    for name, edges in [('tree', [(0, 1), (1, 2)]), ('triangle', [(0, 1), (0, 2), (1, 2)])]:
        m = Ising(3, edges, [3]*len(edges), [1]*3)
        corr = sum(p*x[0]*x[2] for x, p in m.pi.items())
        same = sum(p for x, p in m.pi.items() if len(set(x)) == 1)
        ans[name] = {'Z_tilde': m.Z, 'correlation_13': corr, 'all_equal': same, 'C': m.influence()}
        check((m.Z, corr, same) == ((32, F(1, 4), F(9, 16)) if name == 'tree' else (72, F(2, 3), F(3, 4))), 'displayed_tree_triangle_law')
    check(3*F(1, 4)*F(3, 4)**2+F(1, 4)**3 == F(7, 16), 'impossible_independent_cycle_edges')
    # Enumerate the root/edge bijection on several trees, including antiferromagnetic factors.
    for n in range(1, 7):
        edges = [(i//2, i) for i in range(1, n)]
        for ratio in [F(1, 3), F(1), F(3)]:
            m = Ising(n, edges, [ratio]*len(edges), [1]*n)
            images = set()
            for x in m.states:
                root_edge = (x[0],)+tuple(x[i]*x[j] for i, j in edges)
                images.add(root_edge)
                p = F(1, 2)
                for a in root_edge[1:]:
                    p *= ratio/(1+ratio) if a == 1 else 1/(1+ratio)
                check(p == m.pi[x], 'tree_independent_edge_exact_probability')
            check(len(images) == 2**n, 'tree_root_edge_bijection')
    rho = F(14, 15)
    check(3*rho**82 > F(1, 100) >= 3*rho**83, 'triangle_first_certified_step_83')
    ans['triangle_steps'] = {'rho': rho, 'steps': 83, 'bound_82': 3*rho**82, 'bound_83': 3*rho**83}
    return ans


def finite_law_checks():
    for n in range(1, 5):
        graph_options = [[], [(i, i+1) for i in range(n-1)], [(i, j) for i in range(n) for j in range(i+1, n)]]
        for edges_tuple in dict.fromkeys(map(tuple, graph_options)):
            edges = list(edges_tuple)
            for r in [F(1, 3), F(1), F(7, 3)]:
                for h in [F(1, 4), F(1), F(9)]:
                    m = Ising(n, edges, [r]*len(edges), [h if i % 2 == 0 else 1/h for i in range(n)])
                    check(sum(m.pi.values()) == 1 and all(p > 0 for p in m.pi.values()), 'positive_normalized_finite_law')
                    C = m.influence()
                    for i in range(n):
                        check(m.push(m.pi, i) == m.pi, 'individual_Gibbs_stationarity')
                        for x in m.states:
                            y = list(x); y[i] *= -1; y = tuple(y)
                            a = list(x); a[i] = 1; a = tuple(a)
                            b = list(x); b[i] = -1; b = tuple(b)
                            p = m.positive(x, i)
                            check(p == m.pi[a]/(m.pi[a]+m.pi[b]), 'local_odds_match_full_conditional')
                            pxy = p if y[i] == 1 else 1-p
                            pyx = m.positive(y, i) if x[i] == 1 else 1-m.positive(y, i)
                            check(m.pi[x]*pxy == m.pi[y]*pyx, 'single_site_detailed_balance')
                        for x in m.states:
                            for y in m.states:
                                check(abs(m.positive(x, i)-m.positive(y, i)) <= sum(C[i][j] for j in range(n) if x[j] != y[j]), 'conditional_telescoping_influence')
                    for order in [list(range(n)), list(reversed(range(n)))]:
                        for initial in [(-1,)*n, (1,)*n]:
                            budget, mu = scan_model(m, order, initial)
                            for mask in product((0, 1), repeat=n):
                                subset = [i for i in range(n) if mask[i]]
                                check(tv(marginal(mu, subset), marginal(m.pi, subset)) <= min(1, sum(budget[i] for i in subset)), 'all_subset_TV_budget_after_scan')


def star_task():
    m = Ising(4, [(0, 1), (0, 2), (0, 3)], [F(7, 3)]*3, [1, 9, 9, 9])
    C = m.influence()
    expected = [[F(0), F(2, 5), F(2, 5), F(2, 5)]]+[[F(30, 187), F(0), F(0), F(0)] for _ in range(3)]
    check(C == expected and m.Z == F(326800, 27), 'asymmetric_star_exact_influence')
    center = sum(p for x, p in m.pi.items() if x[0] == 1)
    check(center == F(35937, 40850), 'star_stationary_center')
    rows = {}
    for name, order in [('center_first', [0, 1, 2, 3]), ('leaves_first', [1, 2, 3, 0])]:
        b, mu = scan_model(m, order, (-1,)*4)
        value = sum(p for x, p in mu.items() if x[0] == 1)
        rows[name] = {'budget': b, 'joint_TV_bound': min(1, sum(b)), 'actual_joint_TV': tv(mu, m.pi), 'center_positive': value, 'center_bias': abs(value-center)}
    check(rows['center_first']['budget'] == [F(6, 5)]+[F(36, 187)]*3, 'center_first_budget')
    check(rows['leaves_first']['budget'] == [F(36, 187)]+[F(30, 187)]*3, 'leaves_first_budget')
    check(rows['center_first']['center_bias'] == F(609687, 755725), 'center_first_actual_bias')
    check(rows['leaves_first']['center_positive'] == F(279153, 363562), 'leaves_first_binomial_conditional_mixture')
    check(rows['leaves_first']['center_bias'] == F(415481886, 3712876925), 'leaves_first_actual_bias')
    rows['C'], rows['stationary_center'] = C, center
    return rows


def path_and_weight_checks():
    # Independent coordinate refresh supplies explicit edge couplings on weighted cubes.
    for n in range(1, 6):
        for weights in [[F(1)]*n, [F(i+1) for i in range(n)]]:
            q = [F(i+1, n*(n+1)//2) for i in range(n)]
            rho = 1-min(q)
            for j in range(n):
                # This coupling cost depends on the differing coordinate, not its bit value.
                # Updating that coordinate joins the pair; other updates preserve it.
                cost = sum(q[i]*(0 if i == j else weights[j]) for i in range(n))
                check(cost <= rho*weights[j], 'weighted_cube_coordinate_coupling_cost')
            check(rho == 1-q[0], 'nonuniform_refresh_slowest_coordinate')
    C = [[F(0), F(2, 5), F(2, 5), F(2, 5)]]+[[F(2, 5), F(0), F(0), F(0)] for _ in range(3)]
    v = [F(7, 4), F(1), F(1), F(1)]
    ratios = [sum(C[i][j]*v[j] for j in range(4))/v[i] for i in range(4)]
    check(ratios == [F(24, 35)]+[F(7, 10)]*3, 'weighted_star_row_ratios')
    check(max(ratios) == F(7, 10) and (1-max(ratios))/4 == F(3, 40), 'weighted_star_contraction_37_40')
    b = [F(1)]*4
    for t in range(101):
        check(all(b[i] <= F(37, 40)**t*v[i] for i in range(4)), 'weighted_star_geometric_budget')
        b = apply_budget(C, b, [F(1, 4)]*4)
    # Gluing with a zero middle marginal: explicitly reconstruct all three marginals.
    for a in range(1, 8):
        for b in range(1, 8):
            mu = [F(a, a+b), F(b, a+b)]
            nu = [F(0), F(1, 3), F(2, 3)]
            lam = [F(2, 5), F(3, 5)]
            triple = {(i,j,k): (mu[i]*nu[j])*(nu[j]*lam[k])/nu[j] if nu[j] else F(0) for i in range(2) for j in range(3) for k in range(2)}
            for i,j in product(range(2),range(3)):
                check(sum(triple[i,j,k] for k in range(2)) == mu[i]*nu[j], 'zero_middle_mass_gluing_first_marginal')
            for j,k in product(range(3),range(2)):
                check(sum(triple[i,j,k] for i in range(2)) == nu[j]*lam[k], 'zero_middle_mass_gluing_second_marginal')


def cftp_replay_task():
    m = Ising(3, [(0,1),(0,2),(1,2)], [3]*3, [1]*3)
    records = [(1,F(1,2)),(0,F(1,4)),(1,F(3,4)),(1,F(1,20)),(0,F(1,2)),(2,F(1,4)),(0,F(19,20)),(1,F(1,2))]
    rows = {}
    for T in [1,2,4,8]:
        images = [replay(m,records[-T:],x)[-1] for x in m.states]
        rows[str(T)] = {'lower': images[0], 'upper': images[-1], 'all_images': images, 'coalesced': len(set(images)) == 1}
        check(rows[str(T)]['coalesced'] == (T == 8), 'fixed_suffix_doubling_coalescence')
        for x in images:
            check(all(images[0][i] <= x[i] <= images[-1][i] for i in range(3)), 'monotone_extremes_bound_every_state')
    rows['lower_trace'] = replay(m,records,(-1,)*3)
    rows['upper_trace'] = replay(m,records,(1,)*3)
    check(rows['8']['lower'] == (-1,-1,1), 'eight_record_output')
    for i in range(3):
        for u in [F(0),F(1,20),F(1,10),F(1,2),F(9,10),F(19,20),F(1)]:
            for x in m.states:
                for y in m.states:
                    if all(a <= b for a,b in zip(x,y)):
                        check(all(a <= b for a,b in zip(m.update(x,i,u),m.update(y,i,u))), 'ferromagnetic_map_monotonicity_including_ties')
    sync = [(i,F(1,20)) for i in range(3)]
    check(all(replay(m,sync,x)[-1] == (1,1,1) for x in m.states), 'three_site_synchronizing_word')
    gamma=F(1,27000)
    check(3/gamma == 81000, 'synchronizing_word_expected_horizon_bound')
    rows['records_1_based'] = [(i+1,u) for i,u in records]
    rows['word_probability'],rows['expected_horizon_bound'] = gamma,F(81000)
    return rows


def stopping_and_counterexamples():
    # A map is represented by its values at 0 and 1.
    maps=[((0,0),F(1,2)),((1,1),F(1,4)),((1,0),F(1,4))]
    P=[[sum(p for f,p in maps if f[x] == y) for y in range(2)] for x in range(2)]
    pi=[F(3,5),F(2,5)]
    for y in range(2):
        check(sum(pi[x]*P[x][y] for x in range(2)) == pi[y], 'two_state_exact_stationarity')
    distribution={(0,1):F(1)}
    rows=[]
    for T in range(1,31):
        new=Counter()
        # Add older f to fixed endpoint composition H: the new map is H o f.
        for H,mass in distribution.items():
            for f,p in maps:
                new[(H[f[0]],H[f[1]])] += mass*p
        distribution=new
        completed_one=distribution[(1,1)]
        unresolved=sum(p for H,p in distribution.items() if H[0] != H[1])
        check(unresolved == F(1,4)**T, 'CFTP_unresolved_history_exact_mass')
        check(completed_one <= F(2,5) <= completed_one+unresolved, 'fixed_endpoint_partial_history_brackets_stationary_mass')
        rows.append({'horizon':T,'completed_one_mass':completed_one,'unresolved_mass':unresolved})
    check(F(1,4)/F(3,4) == F(1,3), 'forward_first_coalescence_and_fast_sample_selection_bias')
    check((F(1,4)+F(1,8))/(1-F(1,16)) == F(2,5), 'CFTP_even_odd_suffix_exact_geometric_sum')
    # Identity/swap grand coupling has a fully mixing kernel but no coalescent word.
    for x,y in product(range(2),repeat=2):
        check(sum(p for f,p in [((0,1),F(1,2)),((1,0),F(1,2))] if f[x] == y) == F(1,2), 'noncoalescing_maps_same_mixing_kernel')
    permutations=[(0,1),(1,0)]
    for f,g in product(permutations,repeat=2):
        composition=(f[g[0]],f[g[1]])
        check(composition in permutations and composition[0] != composition[1], 'identity_swap_composition_closure')
    # Synchronous two-site updates do not preserve the correlated target.
    m=Ising(2,[(0,1)],[3],[1,1]);out=Counter()
    for x,mass in m.pi.items():
        for y in m.states:
            p=mass
            for i in range(2):
                q=m.positive(x,i);p*=q if y[i] == 1 else 1-q
            out[y]+=p
    check(sum(out[y]*y[0]*y[1] for y in out) == F(1,8), 'parallel_old_state_updates_destroy_target_correlation')
    # Fair independent bits, state-dependent site choice.
    fair=Ising(2,[],[],[1,1]);adapt=Counter()
    for x,mass in fair.pi.items():
        i=0 if x[0] == -1 else 1
        for y,p in fair.push({x:F(1)},i).items():adapt[y]+=mass*p
    check(sum(p for x,p in adapt.items() if x[0] == 1) == F(3,4), 'state_dependent_scan_breaks_invariance')
    return {'P':P,'stationary_one':F(2,5),'forward_stop_one':F(1,3),'one_step_completed_conditional_one':F(1,3),'partial_histories':rows}


def rejected_inputs():
    tests=[lambda:Ising(0,[],[],[]),lambda:Ising(2,[(0,0)],[1],[1,1]),lambda:Ising(2,[(0,1),(0,1)],[1,1],[1,1]),lambda:Ising(1,[],[],[0]),lambda:Ising(2,[(0,1)],[-1],[1,1]),lambda:apply_budget([[F(0)]],[F(1)],[F(1,2)])]
    for f in tests:
        try:f()
        except ValueError:check(True,'invalid_model_or_scan_rejected')
        else:raise RuntimeError('invalid input accepted')


def serial(value):
    if isinstance(value,F):return str(value)
    if isinstance(value,dict):return {str(k):serial(v) for k,v in value.items()}
    if isinstance(value,(list,tuple)):return [serial(v) for v in value]
    return value


def main():
    CHECKS.clear()
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path,required=True)
    args=parser.parse_args()
    results={'tree_cycle':tree_and_cycle()}
    finite_law_checks();path_and_weight_checks()
    results['asymmetric_star']=star_task()
    results['CFTP_replay']=cftp_replay_task()
    results['stopping']=stopping_and_counterexamples()
    rejected_inputs()
    results={'status':'PASS','checks':sum(CHECKS.values()),'groups':dict(CHECKS),'arithmetic':'fractions.Fraction; explicit exception checks also run under python -O','scope':'Finite enumerations and exact certificates for the stated models; no simulation claim and no bit-cost bound.',**results}
    args.output.parent.mkdir(parents=True,exist_ok=True)
    args.output.write_text(json.dumps(serial(results),ensure_ascii=False,indent=2)+'\n')
    print(json.dumps({'status':'PASS','checks':sum(CHECKS.values()),'output':str(args.output)},ensure_ascii=False))


if __name__ == '__main__':
    main()
