#!/usr/bin/env python3
"""Exact, standard-library-only examples for PIT, isolation and Tutte certificates.
All public graph inputs are simple undirected graphs on range(n). Prime-field
routines require a certified prime; this pedagogical implementation accepts
only the explicit verified primes below. Random seeds make examples replayable,
not a replacement for the independent-uniform-coin hypothesis in the proofs.
"""
import sys
sys.dont_write_bytecode = True
from fractions import Fraction
from itertools import combinations, permutations, product
from pathlib import Path
from random import Random
import json

PRIMES = (2, 3, 5, 17, 257, 65537)

def require(ok, reason):
    if not ok:
        raise ValueError(reason)

def check(ok, reason):
    if not ok:
        raise AssertionError(reason)

def prime(p):
    require(type(p) is int and p in PRIMES, 'choose a listed certified prime')

def circuit_validate(gates, variables, output):
    require(type(variables) is int and variables >= 0, 'variable count')
    require(type(output) is int and 0 <= output < len(gates), 'output gate')
    degree = []
    for i, g in enumerate(gates):
        require(isinstance(g, (tuple, list)) and len(g) >= 1, 'gate record')
        if g[0] == 'const':
            require(len(g) == 2 and type(g[1]) is int, 'integer constant')
            degree.append(0)
        elif g[0] == 'var':
            require(len(g) == 2 and type(g[1]) is int and 0 <= g[1] < variables, 'variable')
            degree.append(1)
        else:
            require(len(g) == 3 and g[0] in ('add', 'sub', 'mul'), 'gate operation')
            require(all(type(x) is int and 0 <= x < i for x in g[1:]), 'topological predecessor')
            a, b = (degree[x] for x in g[1:])
            degree.append(a + b if g[0] == 'mul' else max(a, b))
    return degree

def evaluate(gates, point, p):
    """Already validated DAG; retain all gate values as a replayable trace."""
    values = []
    for g in gates:
        if g[0] == 'const': value = g[1]
        elif g[0] == 'var': value = point[g[1]]
        elif g[0] == 'add': value = values[g[1]] + values[g[2]]
        elif g[0] == 'sub': value = values[g[1]] - values[g[2]]
        else: value = values[g[1]] * values[g[2]]
        values.append(value % p)
    return values

def pit(gates, variables, output, p, sample_bits, rounds, rng):
    prime(p)
    require(type(sample_bits) is int and sample_bits >= 0, 'sample bits')
    require(type(rounds) is int and rounds > 0, 'positive rounds')
    M = 1 << sample_bits
    require(M <= p, 'sample set must inject into the field')
    degrees = circuit_validate(gates, variables, output)
    D = degrees[output]
    bound = min(Fraction(1), Fraction(D, M)) ** rounds
    # For degree zero, evaluate the constant circuit once without random bits.
    count = 1 if D == 0 else rounds
    trials = []
    for _ in range(count):
        point = [0] * variables if D == 0 else [rng.getrandbits(sample_bits) for _ in range(variables)]
        values = evaluate(gates, point, p)
        trials.append({'point': point, 'value': values[output]})
        if values[output] != 0:
            return {'status': 'NONZERO', 'degree_bound': D, 'point': point,
                    'gate_values': values, 'trials': trials, 'miss_bound': str(bound)}
    return {'status': 'ZERO' if D == 0 else 'NOT_DETECTED', 'degree_bound': D,
            'trials': trials, 'miss_bound': str(bound)}

def graph_validate(n, edges):
    require(type(n) is int and n >= 0, 'vertex count')
    normalized = []
    for e in edges:
        require(len(e) == 2 and all(type(v) is int and 0 <= v < n for v in e), 'edge endpoints')
        u, v = sorted(e)
        require(u != v, 'no loops')
        normalized.append((u, v))
    require(len(set(normalized)) == len(normalized), 'no parallel edges')
    return tuple(sorted(normalized))

def tutte_matrix(n, edges, values):
    """Edge/value sequences are aligned. Lower entries use the SAME variable negated."""
    require(len(edges) == len(values), 'edge/value lengths')
    A = [[0] * n for _ in range(n)]
    for (u, v), x in zip(edges, values):
        A[u][v], A[v][u] = x, -x
    return A

def square_matrix(A):
    n = len(A)
    require(all(len(row) == n for row in A), 'square matrix')
    require(all(type(x) is int for row in A for x in row), 'integer matrix')
    return n

def det_mod(A, p):
    prime(p)
    n = square_matrix(A)
    a = [[x % p for x in row] for row in A]
    d = 1
    for k in range(n):
        j = next((i for i in range(k, n) if a[i][k]), None)
        if j is None: return 0
        if j != k:
            a[k], a[j] = a[j], a[k]
            d = -d
        pivot = a[k][k]
        d = d * pivot % p
        inv = pow(pivot, -1, p)
        for i in range(k + 1, n):
            factor = a[i][k] * inv % p
            a[i][k] = 0
            for j in range(k + 1, n):
                a[i][j] = (a[i][j] - factor * a[k][j]) % p
    return d % p

def det_exact(A, stats=None):
    """Row-pivoted rational elimination. Fraction reduces after each operation."""
    n = square_matrix(A)
    a = [[Fraction(x) for x in row] for row in A]
    d = Fraction(1)
    def observe(x):
        if stats is not None:
            stats['max_bits'] = max(stats.get('max_bits', 0), abs(x.numerator).bit_length(), x.denominator.bit_length())
    for row in a:
        for x in row: observe(x)
    for k in range(n):
        j = next((i for i in range(k, n) if a[i][k]), None)
        if j is None: return 0
        if j != k:
            a[k], a[j] = a[j], a[k]
            d = -d
        pivot = a[k][k]
        d *= pivot
        observe(d)
        for i in range(k + 1, n):
            factor = a[i][k] / pivot
            observe(factor)
            a[i][k] = Fraction(0)
            for j in range(k + 1, n):
                a[i][j] -= factor * a[k][j]
                observe(a[i][j])
    check(d.denominator == 1, 'integer determinant exactness')
    return d.numerator

def valuation2(x):
    require(type(x) is int, 'integer valuation')
    if x == 0: return None  # infinity, never equal to a finite extraction target
    x = abs(x)
    return (x & -x).bit_length() - 1

def perfect_certificate(n, edges, candidate):
    E = set(edges)
    used = set()
    for e in candidate:
        if e not in E or e[0] in used or e[1] in used: return False
        used.update(e)
    return len(used) == n and len(candidate) * 2 == n

def extract_weighted(n, edges, weights):
    """Weights correspond to the sorted canonical edge table, not input order."""
    E = graph_validate(n, edges)
    require(len(weights) == len(E) and all(type(w) is int and w > 0 for w in weights), 'positive integer weights')
    if n == 0:
        return {'status': 'MATCHING', 'matching': [], 'determinant': 1, 'valuation': 0, 'minors': [], 'max_bits': 1}
    if n % 2:
        return {'status': 'NO', 'reason': 'odd vertex count'}
    A = tutte_matrix(n, E, [1 << w for w in weights])
    stats = {'max_bits': 0}
    d = det_exact(A, stats)
    v = valuation2(d)
    if v is None:
        return {'status': 'RETRY', 'determinant': 0, 'valuation': None, 'minors': [], **stats}
    candidate = []
    traces = []
    for e, w in zip(E, weights):
        keep = [i for i in range(n) if i not in e]
        minor = det_exact([[A[i][j] for j in keep] for i in keep], stats)
        mv = valuation2(minor)
        selected = mv is not None and mv + 2 * w == v
        traces.append({'edge': list(e), 'weight': w, 'minor': minor, 'valuation': mv, 'selected': selected})
        if selected: candidate.append(e)
    ok = perfect_certificate(n, E, candidate)
    return {'status': 'MATCHING' if ok else 'RETRY', 'matching': [list(e) for e in candidate] if ok else None,
            'candidate': [list(e) for e in candidate], 'determinant': d, 'valuation': v, 'minors': traces, **stats}

def search_matching(n, edges, rounds, rng):
    E = graph_validate(n, edges)
    require(type(rounds) is int and rounds > 0, 'positive search budget')
    if n == 0: return {'status': 'MATCHING', 'matching': [], 'attempts': 0}
    if n % 2: return {'status': 'NO', 'reason': 'odd vertex count', 'attempts': 0}
    if not E: return {'status': 'NO', 'reason': 'nonempty edgeless graph', 'attempts': 0}
    R = 2 * len(E)
    history = []
    for _ in range(rounds):
        weights = [rng.randrange(1, R + 1) for _ in E]
        result = extract_weighted(n, E, weights)
        history.append({'weights': weights, 'status': result['status']})
        if result['status'] == 'MATCHING':
            return {'status': 'MATCHING', 'matching': result['matching'], 'attempts': len(history), 'history': history}
    return {'status': 'UNKNOWN', 'attempts': rounds, 'history': history, 'conditional_miss_bound': str(Fraction(1, 2 ** rounds))}

# Independent small-instance oracles below deliberately do not use elimination.
def det_leibniz(A):
    n = len(A)
    total = 0
    for perm in permutations(range(n)):
        term = -1 if sum(perm[i] > perm[j] for i in range(n) for j in range(i + 1, n)) % 2 else 1
        for i in range(n): term *= A[i][perm[i]]
        total += term
    return total

def all_matchings(n, edges):
    E = set(edges)
    def visit(vertices):
        if not vertices:
            yield ()
            return
        u = min(vertices)
        for v in sorted(vertices - {u}):
            e = (u, v)
            if e in E:
                for rest in visit(vertices - {u, v}): yield (e,) + rest
    return list(visit(set(range(n))))

def sparse_polynomial(gates, variables, p):
    polynomials = []
    zero_exp = (0,) * variables
    for g in gates:
        if g[0] == 'const': z = {zero_exp: g[1] % p}
        elif g[0] == 'var': z = {tuple(int(i == g[1]) for i in range(variables)): 1}
        elif g[0] in ('add', 'sub'):
            z = dict(polynomials[g[1]])
            sign = 1 if g[0] == 'add' else -1
            for mon, c in polynomials[g[2]].items(): z[mon] = (z.get(mon, 0) + sign * c) % p
        else:
            z = {}
            for a, ca in polynomials[g[1]].items():
                for b, cb in polynomials[g[2]].items():
                    mon = tuple(x + y for x, y in zip(a, b))
                    z[mon] = (z.get(mon, 0) + ca * cb) % p
        polynomials.append({mon: c for mon, c in z.items() if c})
    return polynomials

def sparse_value(poly, point, p):
    total = 0
    for mon, c in poly.items():
        term = c
        for x, e in zip(point, mon): term = term * pow(x, e, p) % p
        total += term
    return total % p

def main():
    rng = Random(7102026)
    counters = {}
    # Circuit representation is checked against independently expanded coefficients.
    evaluations = 0
    for _ in range(1200):
        variables = rng.randrange(1, 4)
        gates = [('var', i) for i in range(variables)] + [('const', rng.randrange(-5, 6))]
        for j in range(rng.randrange(1, 12)):
            gates.append((rng.choice(('add', 'sub', 'mul')), rng.randrange(len(gates)), rng.randrange(len(gates))))
        degree = circuit_validate(gates, variables, len(gates) - 1)
        poly = sparse_polynomial(gates, variables, 5)[-1]
        check(not poly or max(map(sum, poly)) <= degree[-1], 'syntactic degree bound')
        zeros = 0
        for point in product(range(5), repeat=variables):
            val = evaluate(gates, point, 5)[-1]
            check(val == sparse_value(poly, point, 5), 'DAG vs coefficients')
            zeros += val == 0
            evaluations += 1
        if poly: check(zeros * 5 <= degree[-1] * 5 ** variables, 'Schwartz-Zippel count')
        result = pit(gates, variables, len(gates)-1, 17, 4, 3, rng)
        if result['status'] == 'NONZERO': check(result['gate_values'][-1] != 0, 'PIT certificate')
    counters['circuit_instances'] = 1200
    counters['exact_circuit_point_evaluations'] = evaluations
    # Every nonempty family on a three-element universe, every weight in 1..6.
    bad_families = 0
    weight_assignments = 0
    max_failures = 0
    for mask in range(1, 1 << 8):
        family = [s for s in range(8) if mask >> s & 1]
        failures = 0
        for weights in product(range(1, 7), repeat=3):
            scores = [sum(weights[i] for i in range(3) if s >> i & 1) for s in family]
            best = min(scores)
            failures += scores.count(best) > 1
            weight_assignments += 1
        check(failures <= 108, 'isolation bound over fixed family')
        max_failures = max(max_failures, failures)
        bad_families += failures > 0
    counters.update(isolation_families=255, isolation_weight_assignments=weight_assignments,
                    isolation_max_failed_assignments=max_failures, isolation_families_with_ties=bad_families)
    determinant_checks = 0
    for n in range(6):
        for _ in range(100):
            A = [[rng.randrange(-7, 8) for _ in range(n)] for _ in range(n)]
            exact = det_leibniz(A)
            check(det_exact(A) == exact, 'rational vs permutation determinant')
            for p in (2, 3, 5, 17): check(det_mod(A, p) == exact % p, 'modular determinant')
            determinant_checks += 1
    counters['arbitrary_integer_determinants'] = determinant_checks
    graphs = trials = isolated = successful = 0
    for n in range(6):
        possible = list(combinations(range(n), 2))
        for mask in range(1 << len(possible)):
            edges = [e for i, e in enumerate(possible) if mask >> i & 1]
            matches = all_matchings(n, edges)
            for _ in range(3):
                weights = [rng.randrange(1, 2 * max(1, len(edges)) + 1) for e in edges]
                result = extract_weighted(n, edges, weights)
                if result['status'] == 'MATCHING':
                    successful += 1
                    check(perfect_certificate(n, edges, [tuple(e) for e in result['matching']]), 'explicit returned certificate')
                if matches:
                    wm = dict(zip(edges, weights))
                    scores = [sum(wm[e] for e in mat) for mat in matches]
                    if scores.count(min(scores)) == 1:
                        isolated += 1
                        expected = matches[scores.index(min(scores))]
                        check(result['status'] == 'MATCHING' and set(map(tuple, result['matching'])) == set(expected), 'isolated extraction')
                else: check(result['status'] != 'MATCHING', 'false certificate')
                values = [rng.randrange(17) for e in edges]
                det = det_mod(tutte_matrix(n, edges, values), 17)
                if det: check(bool(matches), 'finite-field nonzero existence certificate')
                trials += 1
            graphs += 1
    for _ in range(600):
        n = rng.choice((6, 8))
        edges = [e for e in combinations(range(n), 2) if rng.random() < .45]
        matches = all_matchings(n, edges)
        weights = [rng.randrange(1, 2 * max(1, len(edges)) + 1) for e in edges]
        result = extract_weighted(n, edges, weights)
        if result['status'] == 'MATCHING':
            successful += 1
            check(perfect_certificate(n, edges, list(map(tuple, result['matching']))), 'large small-graph witness')
        if matches:
            wm = dict(zip(edges, weights))
            scores = [sum(wm[e] for e in mat) for mat in matches]
            if scores.count(min(scores)) == 1:
                isolated += 1
                check(result['status'] == 'MATCHING' and set(map(tuple, result['matching'])) == set(matches[scores.index(min(scores))]), 'large isolated extraction')
        else: check(result['status'] != 'MATCHING', 'large false certificate')
        graphs += 1
        trials += 1
    counters.update(graph_instances=graphs, weighted_extractions=trials, isolated_instances=isolated, returned_matchings=successful)
    edges = list(combinations(range(4), 2))
    main_trace = extract_weighted(4, edges, [1, 3, 5, 6, 4, 2])
    check(main_trace['determinant'] == 3717184 and main_trace['valuation'] == 6, 'hand determinant')
    ties = extract_weighted(4, edges, [1] * 6)
    check(ties['status'] == 'RETRY' and len(ties['candidate']) == 6, 'validate nonisolated result')
    cancelling_edges = [(0,1),(0,2),(1,3),(2,3)]
    cancellation = extract_weighted(4, cancelling_edges, [1]*4)
    check(cancellation['determinant'] == 0 and len(all_matchings(4,cancelling_edges)) == 2, 'zero substitution is not absence')
    # F4 = F2[a]/(a^2+a+1): 0,1,a,a+1 encoded as two bits.
    def f4_mul(x,y):
        raw = (x if y & 1 else 0) ^ ((x << 1) if y & 2 else 0)
        return raw ^ 7 if raw & 4 else raw
    f4_values = [f4_mul(x,x) ^ x for x in range(4)]
    check(f4_values == [0,0,1,1], 'same-characteristic extension detects X^2-X')
    # Difference (x+y)^2 - (x^2+2xy+y^2), plus a perturbed version minus x.
    gates = [('var',0),('var',1),('add',0,1),('mul',2,2),('mul',0,0),('mul',0,1),('const',2),('mul',5,6),('mul',1,1),('add',4,7),('add',9,8),('sub',3,10)]
    zero = pit(gates,2,11,17,4,4,Random(7))
    changed = pit(gates+[('sub',11,0)],2,12,17,4,4,Random(7))
    check(zero['status']=='NOT_DETECTED' and changed['status']=='NONZERO','PIT hand example')
    invalid = [lambda: circuit_validate([('add',0,0)],0,0),lambda: circuit_validate([('var',1)],1,0),
               lambda: pit([('const',0)],0,0,4,1,1,rng),lambda: pit([('const',0)],0,0,3,2,1,rng),
               lambda: pit([('const',0)],0,0,5,1,0,rng),lambda: graph_validate(2,[(0,1),(1,0)]),
               lambda: graph_validate(1,[(0,0)]),lambda: extract_weighted(2,[(0,1)],[0]),
               lambda: det_exact([[1,2]]),lambda: search_matching(0,[],0,rng)]
    for fn in invalid:
        try: fn()
        except ValueError: pass
        else: raise AssertionError('invalid input was accepted')
    counters['invalid_interfaces_rejected'] = len(invalid)
    result = {'status':'PASS', **counters, 'examples':{'identity':zero,'perturbed_identity':changed,'unique_K4':main_trace,
              'tied_K4':ties,'cancelling_K22':cancellation,'F4_values_of_X2_minus_X':f4_values,
              'bounded_search': search_matching(4,edges,8,Random(11))},
              'probability_scope':'Exact exhaustive finite spaces verify examples; general guarantees are proved in the pages. Fixed PRNG seeds are replayability only.'}
    output=Path(__file__).with_name('algorithms-algebraic-certificates-results.json')
    output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps(result,ensure_ascii=False,indent=2))

if __name__ == '__main__': main()
