#!/usr/bin/env python3
"""Exact, standard-library finite change-detection certificates.

Run with --output to keep author/reviewer outputs separate. Finite checks do
not replace the finite-block absorption and stopped-martingale proofs.
"""
from fractions import Fraction as F
from itertools import product
from collections import Counter, deque
from pathlib import Path
import argparse
import json

COUNTS = Counter()

def check(ok, name):
    COUNTS[name] += 1
    if not ok:
        raise ArithmeticError(name)

def bernoulli_factors(p0, p1):
    p0, p1 = F(p0), F(p1)
    if not 0 < p0 < 1 or not 0 <= p1 <= 1:
        raise ValueError('need 0 < p0 < 1 and 0 <= p1 <= 1')
    return ((1-p1)/(1-p0), p1/p0)

def statistics(factors, r=0):
    """(unfloored suffix maximum V, floored W=exp(C), headstarted SR)."""
    r = F(r)
    if r < 0:
        raise ValueError('negative headstart')
    v, w, sr = F(1), F(1), r
    out = []
    for a in factors:
        a = F(a)
        if a < 0:
            raise ValueError('negative likelihood factor')
        v = max(1, v)*a
        w = max(1, w*a)
        sr = (1+sr)*a
        out.append((v, w, sr))
    return out

def suffix_statistics(factors, r=0):
    """Direct quadratic-time definition, without the recursive update."""
    out = []
    for n in range(1, len(factors)+1):
        suffixes = []
        for k in range(n):
            s = F(1)
            for a in factors[k:n]:
                s *= a
            suffixes.append(s)
        whole = F(1)
        for a in factors[:n]:
            whole *= a
        v = max(suffixes)
        out.append((v, max(F(1), v), sum(suffixes, F(0))+F(r)*whole))
    return out

def geometric_mixture(factors):
    """pi_k=2^-k, with the unstarted tail T_n=2^-n retained."""
    b, tail = F(1), F(1)
    out = []
    for a in factors:
        next_tail = tail/2
        b = a*(b-tail+next_tail)+next_tail
        tail = next_tail
        out.append(b)
    return out

def direct_mixture(factors):
    out = []
    for n in range(1, len(factors)+1):
        b = F(1, 2**n)
        for k in range(n):
            term = F(1, 2**(k+1))
            for a in factors[k:n]:
                term *= a
            b += term
        out.append(b)
    return out

def solve(a, b):
    """Square rational elimination; a singular system is explicitly refused."""
    n = len(a)
    if len(b) != n or any(len(row) != n for row in a):
        raise ValueError('square system required')
    m = [[F(x) for x in row]+[F(y)] for row, y in zip(a, b)]
    for j in range(n):
        pivot = next((i for i in range(j, n) if m[i][j]), None)
        if pivot is None:
            raise ValueError('singular system; no finite inverse certificate')
        m[j], m[pivot] = m[pivot], m[j]
        z = m[j][j]
        m[j] = [x/z for x in m[j]]
        for i in range(n):
            if i != j:
                z = m[i][j]
                m[i] = [x-z*y for x, y in zip(m[i], m[j])]
    return [row[-1] for row in m]

def waiting_means(q):
    n = len(q)
    if any(len(row) != n or any(x < 0 for x in row) or sum(row) > 1 for row in q):
        raise ValueError('substochastic square matrix required')
    return solve([[F(i == j)-q[i][j] for j in range(n)] for i in range(n)], [1]*n)

def survival(q, n, start=0):
    if n < 0 or not isinstance(n, int):
        raise ValueError('nonnegative integer horizon required')
    v = [F(i == start) for i in range(len(q))]
    for _ in range(n):
        v = [sum((v[i]*q[i][j] for i in range(len(q))), F(0)) for j in range(len(q))]
    return sum(v, F(0))

def cusum_chain(p, height):
    p = F(p)
    if not 0 < p <= 1 or not isinstance(height, int) or height < 1:
        raise ValueError('positive success probability and integer height required')
    q = [[F(0) for _ in range(height)] for _ in range(height)]
    for i in range(height):
        q[i][max(0, i-1)] += 1-p
        if i+1 < height:
            q[i][i+1] += p
    return q

def reset_chain(p, multiplier, threshold, r=0):
    """Exact reachable SR states for success multiplier c>1, failure zero.

    There are finitely many successes before each threshold crossing, and a
    failure returns to zero. Breadth-first closure includes a general r<A.
    """
    p, c, a, r = map(F, (p, multiplier, threshold, r))
    if not 0 < p <= 1 or c <= 1 or not 0 <= r < a:
        raise ValueError('need p>0, c>1 and 0<=r<A')
    states, pending = [r], deque([r])
    while pending:
        x = pending.popleft()
        for y in [F(0), c*(1+x)]:
            if y < a and y not in states:
                states.append(y)
                pending.append(y)
    q = [[F(0) for _ in states] for _ in states]
    for i, x in enumerate(states):
        for prob, y in [(1-p, F(0)), (p, c*(1+x))]:
            if y < a:
                q[i][states.index(y)] += prob
    return states, q

def weighted_paths(n, p):
    for bits in product((0, 1), repeat=n):
        k = sum(bits)
        yield bits, F(p)**k*(1-F(p))**(n-k)

def first_cross(values, threshold):
    return next((i+1 for i, x in enumerate(values) if x >= threshold), None)

def run():
    COUNTS.clear()
    path_cases = 0
    # Explicit suffixes, recursion and retained prior tail on zero/one factors.
    for n in range(7):
        for factors in product((F(0), F(1, 2), F(1), F(2)), repeat=n):
            for r in (F(0), F(2, 3)):
                trace = statistics(factors, r)
                check(trace == suffix_statistics(factors, r), 'suffix versus recursive statistics')
                for v, w, sr in trace:
                    check(w == max(1, v) and sr >= v, 'floor identity and sum-max ordering')
            check(geometric_mixture(factors) == direct_mixture(factors), 'proper prior tail direct versus recursive')
            path_cases += 1
    for p0, p1 in [(F(1,3),F(2,3)), (F(1,2),F(1)), (F(3,5),F(1)), (F(2,5),F(2,5)), (F(2,3),F(0))]:
        lf = bernoulli_factors(p0, p1)
        check((1-p0)*lf[0]+p0*lf[1] == 1, 'correct conditional likelihood budget')
        for n in range(1, 8):
            e_sr, e_mix = F(0), F(0)
            for bits, weight in weighted_paths(n, p0):
                fs = [lf[b] for b in bits]
                e_sr += weight*statistics(fs)[-1][2]
                e_mix += weight*geometric_mixture(fs)[-1]
            check(e_sr == n and e_mix == 1, 'fixed-time total budgets')
        for n in range(6):
            for bits, _ in weighted_paths(n, p0):
                fs = [lf[b] for b in bits]
                old = geometric_mixture(fs)[-1] if fs else F(1)
                new = [(geometric_mixture(fs+[l]))[-1] for l in lf]
                check((1-p0)*new[0]+p0*new[1] == old, 'proper prior conditional martingale')
        for r in (F(0), F(1,2), F(2)):
            for a in (F(3), F(6), F(13,2)):
                for n in range(1, 7):
                    er, et, alarm = F(0), F(0), F(0)
                    for bits, weight in weighted_paths(n, p0):
                        vals = [x[2] for x in statistics([lf[b] for b in bits], r)]
                        tau = first_cross(vals, a)
                        stopped = tau or n
                        er += weight*vals[stopped-1]
                        et += weight*stopped
                        alarm += weight*(tau is not None)
                    check(er == r+et, 'bounded stopped martingale identity')
                    check(a*alarm <= r+et <= r+n, 'finite-horizon alarm budget')
    chain_cases = 0
    for p in (F(1,4),F(1,3),F(1,2),F(2,3),F(3,4),F(1)):
        for h in range(1, 6):
            q = cusum_chain(p, h)
            means = waiting_means(q)
            for i in range(h):
                check(means[i] == 1+sum(q[i][j]*means[j] for j in range(h)), 'reflected first-step mean certificate')
                check(means[i] >= h-i, 'minimum possible reflected waiting')
                check(survival(q,h,i) <= 1-p**h, 'uniform favorable-block absorption')
                if i:
                    check(means[i] <= means[i-1], 'monotone initial-state delay')
            for n in range(8):
                for start in range(h):
                    alive = F(0)
                    for bits, weight in weighted_paths(n, p):
                        s, hit = start, False
                        for bit in bits:
                            s = max(0, s+(1 if bit else -1))
                            if s >= h:
                                hit = True
                                break
                        alive += weight*(not hit)
                    check(alive == survival(q,n,start), 'reflected path enumeration versus matrix powers')
            chain_cases += 1
    reset_cases = 0
    for p in (F(1,3),F(1,2),F(3,5),F(1)):
        for c in (F(3,2),F(5,3),F(2),F(3)):
            for a in (F(2),F(6),F(13,2),F(14)):
                for r in (F(0),F(1,2)):
                    states, q = reset_chain(p,c,a,r)
                    means = waiting_means(q)
                    for i in range(len(states)):
                        check(means[i] == 1+sum(q[i][j]*means[j] for j in range(len(states))), 'reset first-step mean certificate')
                    for n in range(7):
                        alive = F(0)
                        for bits, weight in weighted_paths(n,p):
                            vals = [x[2] for x in statistics([c if b else F(0) for b in bits],r)]
                            alive += weight*(first_cross(vals,a) is None)
                        check(alive == survival(q,n), 'reset path enumeration versus matrix powers')
                    if r == 0:
                        value, b = F(0), 0
                        while value < a:
                            value = c*(1+value)
                            b += 1
                        check(means[0] == sum((p**(-j) for j in range(1,b+1)),F(0)), 'independent consecutive-run mean formula')
                    if p*c == 1:
                        check(means[0] >= a-r, 'correct-model exact ARL bound')
                    reset_cases += 1
    null = waiting_means(cusum_chain(F(1,3),3))
    changed = waiting_means(cusum_chain(F(2,3),3))
    check(null == [33,30,21] and changed == [F(51,8),F(39,8),F(21,8)], 'terminal CUSUM exact means')
    alarm10 = 1-survival(cusum_chain(F(1,3),3),10)
    check(alarm10 == F(13609,59049), 'terminal horizon ten probability')
    direct_alarm = F(0)
    for bits, weight in weighted_paths(10,F(1,3)):
        values = [t[0] for t in statistics([F(2) if b else F(1,2) for b in bits])]
        direct_alarm += weight*(first_cross(values,8) is not None)
    check(direct_alarm == alarm10, 'terminal full ten-step path enumeration')
    fs = [F(2) if b else F(1,2) for b in [0,0,1,0,1,1,0,1,1]]
    check([x[0] for x in statistics(fs)] == [F(1,2),F(1,2),2,1,2,4,2,4,8], 'terminal late-change trace')
    check(geometric_mixture([F(2)]*3) == [F(3,2),F(11,4),F(43,8)], 'terminal all-success prior trace')
    check(waiting_means(cusum_chain(F(1,3),4)) == [78,75,66,45], 'terminal changed threshold null means')
    check(waiting_means(cusum_chain(F(2,3),4)) == [F(147,16),F(123,16),F(87,16),F(45,16)], 'terminal changed threshold delay means')
    reset_targets = []
    for p,c,a,r,target in [(F(1,2),2,6,0,6),(F(1,2),2,6,2,4),(F(1,2),2,F(13,2),0,14),(F(3,5),2,6,0,F(40,9)),(F(3,5),F(5,3),6,0,F(245,27))]:
        states,q = reset_chain(p,c,a,r)
        actual = waiting_means(q)[0]
        check(actual == target, 'terminal headstart threshold and misspecification')
        reset_targets.append({'p':str(p),'multiplier':str(c),'A':str(a),'r':str(r),'mean':str(actual),'states':[str(x) for x in states]})
    stopped_e = F(0)
    for bits,weight in weighted_paths(2,F(1,2)):
        sr = [x[2] for x in statistics([F(2*b) for b in bits])]
        tau = 1 if bits[0] else 2
        stopped_e += weight*sr[tau-1]/tau
    check(stopped_e == F(5,4), 'normalized SR bounded-stop counterexample')
    check([x[2] for x in statistics([1]*9)] == list(range(1,10)), 'identical model SR immigration')
    check(geometric_mixture([1]*9) == [1]*9, 'identical model proper mixture')
    bad = [lambda:bernoulli_factors(0,F(1,2)),lambda:bernoulli_factors(1,F(1,2)),lambda:bernoulli_factors(F(1,2),2),lambda:statistics([-1]),lambda:statistics([1],-1),lambda:cusum_chain(0,3),lambda:cusum_chain(F(1,2),0),lambda:reset_chain(F(1,2),1,6),lambda:reset_chain(F(1,2),2,6,6),lambda:reset_chain(F(1,2),2,6,-1),lambda:waiting_means([[F(1)]]),lambda:waiting_means([[F(-1)]]),lambda:survival([[F(1,2)]],-1)]
    for f in bad:
        try:
            f()
        except ValueError:
            check(True, 'invalid input refusal')
        else:
            check(False, 'invalid input refusal')
    return {'status':'PASS','checks':sum(COUNTS.values()),'groups':dict(COUNTS),'explicit_factor_paths':path_cases,'reflected_chains':chain_cases,'reset_chains':reset_cases,'cusum_null_means':[str(x) for x in null],'cusum_changed_means':[str(x) for x in changed],'cusum_alarm_by_10':str(alarm10),'reset_targets':reset_targets,'normalized_stopped_mean':str(stopped_e),'scope':'Exact finite certificates, not simulation or a continuous-state numerical solver.'}

if __name__ == '__main__':
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument('--output', required=True, type=Path)
    args = ap.parse_args()
    result = run()
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text(json.dumps(result, ensure_ascii=False, indent=2)+'\n')
    print(json.dumps(result, ensure_ascii=False))
