#!/usr/bin/env python3
"""Finite exact certificates for knockoffs and conditional resampling.

Standard-library Fraction arithmetic. Checks remain active under python -O.
This program enumerates specified finite models and exact matrix identities;
the companion pages prove the general probability theorems. No simulation or
floating-point tolerance is used for an acceptance comparison.
"""
from argparse import ArgumentParser
from collections import Counter, defaultdict
from fractions import Fraction as Q
from itertools import combinations, product
from math import comb
from pathlib import Path
import json

CHECKS = Counter()

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

def matrix(rows):
    return [list(map(Q, row)) for row in rows]

def transpose(a):
    return [list(row) for row in zip(*a)]

def identity(n):
    return [[Q(i == j) for j in range(n)] for i in range(n)]

def multiply(a, b):
    if not a or not b or len(a[0]) != len(b):
        raise ValueError('matrix dimensions')
    return [[sum((x*y for x, y in zip(row, col)), Q(0))
             for col in zip(*b)] for row in a]

def add(a, b, factor=Q(1)):
    return [[x+factor*y for x, y in zip(ra, rb)] for ra, rb in zip(a, b)]

def scale(a, value):
    return [[value*x for x in row] for row in a]

def inverse(a):
    n = len(a)
    if any(len(row) != n for row in a):
        raise ValueError('square inverse input')
    aug = [list(row)+eye for row, eye in zip(a, identity(n))]
    for j in range(n):
        pivot = next((i for i in range(j, n) if aug[i][j]), None)
        if pivot is None:
            raise ValueError('singular input')
        aug[j], aug[pivot] = aug[pivot], aug[j]
        d = aug[j][j]
        aug[j] = [x/d for x in aug[j]]
        for i in range(n):
            if i != j:
                factor = aug[i][j]
                aug[i] = [x-factor*y for x, y in zip(aug[i], aug[j])]
    return [row[n:] for row in aug]

def determinant(a):
    a = [list(row) for row in a]
    value = Q(1)
    for j in range(len(a)):
        pivot = next((i for i in range(j, len(a)) if a[i][j]), None)
        if pivot is None:
            return Q(0)
        if pivot != j:
            a[j], a[pivot] = a[pivot], a[j]
            value = -value
        d = a[j][j]
        value *= d
        for i in range(j+1, len(a)):
            factor = a[i][j]/d
            for k in range(j+1, len(a)):
                a[i][k] -= factor*a[j][k]
    return value

def block_diagonal(blocks):
    nr = sum(len(b) for b in blocks)
    nc = sum(len(b[0]) for b in blocks)
    a = [[Q(0) for _ in range(nc)] for _ in range(nr)]
    row = col = 0
    for b in blocks:
        for i in range(len(b)):
            for j in range(len(b[0])):
                a[row+i][col+j] = b[i][j]
        row += len(b)
        col += len(b[0])
    return a

def select(w, q):
    """Ascending grouped-magnitude scan. None denotes the empty output."""
    w = list(map(Q, w)); q = Q(q)
    if not w or not 0 < q < 1:
        raise ValueError('p >= 1 and 0 < q < 1 required')
    groups = defaultdict(list)
    for i, value in enumerate(w):
        if value:
            groups[abs(value)].append(i)
    positive = sum(value > 0 for value in w)
    negative = sum(value < 0 for value in w)
    for t in sorted(groups):
        if Q(1+negative, max(1, positive)) <= q:
            return t, tuple(i for i, value in enumerate(w) if value >= t)
        for i in groups[t]:
            positive -= w[i] > 0
            negative -= w[i] < 0
    return None, ()

def direct_select(w, q):
    valid = [t for t in set(map(abs, w)) if t > 0 and
             Q(1+sum(x <= -t for x in w), max(1, sum(x >= t for x in w))) <= q]
    if not valid:
        return None, ()
    t = min(valid)
    return t, tuple(i for i, x in enumerate(w) if x >= t)

def false_fraction(selected, nulls):
    return Q(sum(i in nulls for i in selected), max(1, len(selected)))

def correlation_scores(x, xt, y):
    a = multiply(transpose(x), [[v] for v in y])
    b = multiply(transpose(xt), [[v] for v in y])
    return [abs(u[0])-abs(v[0]) for u, v in zip(a, b)]

def fixed_design_checks():
    x0 = matrix([[1,Q(3,5)],[0,Q(4,5)],[0,0],[0,0]])
    xt0 = matrix([[Q(9,25),Q(3,5)],[Q(12,25),0],
                  [Q(4,5),Q(12,25)],[0,Q(16,25)]])
    sigma = multiply(transpose(x0), x0)
    diagonal = scale(identity(2), Q(16,25))
    check(multiply(transpose(xt0), xt0) == sigma, 'copy_gram_exact')
    check(multiply(transpose(x0), xt0) == add(sigma, diagonal, -1), 'cross_gram_exact')
    x = block_diagonal([x0]*3); xt = block_diagonal([xt0]*3)
    y = list(map(Q, [7,Q(9,4),Q(-13,4),Q(-41,16),
                     5,Q(5,4),Q(-7,4),Q(-29,16),
                     1,Q(7,4),Q(9,4),Q(-17,16)]))
    w = correlation_scores(x, xt, y)
    check(w == [6,5,4,3,-2,1], 'six_observed_scores')
    check(select(w,Q(1,4)) == (Q(3),(0,1,2,3)), 'six_observed_selection')
    augmented = [a+b for a,b in zip(x,xt)]
    gram = multiply(transpose(augmented), augmented)
    beta = [[Q(v)] for v in [1,2,-1,0,0,0]]
    mean = multiply(transpose(augmented), multiply(x,beta))
    for mask in product([False,True],repeat=6):
        order = [i+6 if mask[i] else i for i in range(6)]
        order += [i if mask[i] else i+6 for i in range(6)]
        swapped = [[row[k] for k in order] for row in augmented]
        check(multiply(transpose(swapped),swapped)==gram,'all_pair_gram_swaps')
        new = correlation_scores([r[:6] for r in swapped], [r[6:] for r in swapped],y)
        check(new==[-v if mask[i] else v for i,v in enumerate(w)],'all_pair_score_antisymmetry')
        newmean = [mean[i] for i in order]
        check((newmean==mean)==(not any(mask[:3])),'null_mean_swap_and_nonnull_boundary')
    table=[];fdr=Q(0)
    for signs in product([-1,1],repeat=3):
        values=[Q(6),Q(5),Q(4)]+[Q(a*b) for a,b in zip([3,2,1],signs)]
        threshold, selected=select(values,Q(1,4));fdp=false_fraction(selected,{3,4,5});fdr+=fdp/8
        table.append({'signs':list(signs),'threshold':None if threshold is None else str(threshold),
                      'selected_1based':[i+1 for i in selected],'FDP':str(fdp)})
    check(fdr==Q(7,40),'conditional_eight_sign_FDR')
    return {'W':list(map(str,w)),'conditional_FDR':str(fdr),'sign_table':table}

def sign_budget_checks():
    # Conditional hypergeometric deletion identity, including the s=0 deficit.
    for j in range(1,31):
        for s in range(j+1):
            expected=Q(j-s,j)*Q(j,1+s)
            if s:
                expected+=Q(s,j)*Q(j,s)
            current=Q(j+1,1+s)
            check(expected==current-(1 if s==0 else 0),'reverse_one_step_budget')
        initial=sum((Q(comb(j,s),2**j)*Q(j+1,s+1) for s in range(j+1)),Q(0))
        check(initial==2-Q(1,2**j),'binomial_initial_budget')
    maximum=Q(0);cases=0
    # A null/non-null label is part of each configuration; each non-null sign
    # is fixed, while all nonzero null signs are averaged. Ties and zeros vary.
    for p in range(1,7):
        magnitude_families=[list(range(1,p+1)),[i//2+1 for i in range(p)],
                            [0]+list(range(1,p))]
        for magnitudes in magnitude_families:
            for roles in product([-1,0,1],repeat=p):
                nulls={i for i,role in enumerate(roles) if role==0}
                random_indices=[i for i in nulls if magnitudes[i]]
                for q in [Q(1,4),Q(1,2),Q(3,4)]:
                    risk=Q(0)
                    for signs in product([-1,1],repeat=len(random_indices)):
                        w=[Q(m*(r if r else 1)) for m,r in zip(magnitudes,roles)]
                        for i,sign in zip(random_indices,signs):w[i]=Q(sign*magnitudes[i])
                        t,selected=select(w,q)
                        check((t,selected)==direct_select(w,q),'group_scan_vs_direct_definition')
                        if selected:
                            check(q*len(selected)>=1,'minimum_discovery_count')
                        risk+=false_fraction(selected,nulls)/2**len(random_indices)
                    check(risk<=q,'finite_conditional_FDR_with_ties_zeros');cases+=1
                    maximum=max(maximum,risk/q)
    # Marginal symmetry without joint fairness, and the omitted-offset case.
    bad=sum((Q(1,2)*false_fraction(select([sign]*3,Q(1,3))[1],{0,1,2}) for sign in [-1,1]),Q(0))
    check(bad==Q(1,2)>Q(1,3),'common_sign_invalid_counterexample')
    no_offset_risk=Q(0)
    for sign in [-1,1]:
        accepted=Q(int(sign<0),max(1,int(sign>0)))<=Q(1,4)
        no_offset_risk+=Q(int(accepted and sign>0),2)
    check(no_offset_risk==Q(1,2)>Q(1,4),'offset_zero_one_null_counterexample')
    return {'conditional_configurations':cases,'largest_risk_over_q':str(maximum),
            'common_sign_FDR':str(bad),'no_offset_one_null_FDR':'1/2'}

def nonlinear_model_checks():
    sigma=matrix([[1,Q(1,2),Q(1,4)],[Q(1,2),1,Q(1,2)],[Q(1,4),Q(1,2),1]])
    d=matrix([[Q(1,2),0,0],[0,Q(1,3),0],[0,0,Q(1,2)]])
    precision=inverse(sigma);b=add(identity(3),multiply(d,precision),-1)
    v=add(scale(d,2),multiply(multiply(d,precision),d),-1)
    expected_b=matrix([[Q(1,3),Q(1,3),0],[Q(2,9),Q(4,9),Q(2,9)],[0,Q(1,3),Q(1,3)]])
    expected_v=matrix([[Q(2,3),Q(1,9),0],[Q(1,9),Q(13,27),Q(1,9)],[0,Q(1,9),Q(2,3)]])
    check(b==expected_b and v==expected_v,'three_variable_conditional_generator')
    check(multiply(sigma,transpose(b))==add(sigma,d,-1),'random_cross_covariance')
    check(add(multiply(multiply(b,sigma),transpose(b)),v)==sigma,'random_copy_covariance')
    minors=[]
    for a in [sigma,add(scale(sigma,2),d,-1),v]:
        row=[determinant([r[:k] for r in a[:k]]) for k in range(1,4)]
        check(all(z>0 for z in row),'positive_definite_leading_minors');minors.append(list(map(str,row)))
    check(minors==[['1','3/4','9/16'],['3/2','3/2','4/3'],['2/3','25/81','16/81']],'displayed_minors')
    mean=multiply(b,[[Q(1)],[Q(2)],[Q(3)]])
    check(mean==[[Q(1)],[Q(16,9)],[Q(5,3)]],'conditional_mean_record')
    check(sigma[0][2]-sigma[0][1]*sigma[1][2]/sigma[1][1]==0,'conditional_interaction_covariance')
    check(1/precision[0][0]==Q(3,4) and 1/precision[2][2]==Q(3,4),'nondegenerate_nonnull_conditional_variances')
    def interaction(x,z,y):
        return [(x[j]-z[j])*y*sum(x[k]+z[k] for k in range(3) if k!=j) for j in range(3)]
    x=list(map(Q,[1,2,3]));z=list(map(Q,[0,1,4]));w=interaction(x,z,Q(2))
    check(w==[20,16,-8] and select(w,Q(1,2))==(Q(16),(0,1)),'interaction_score_and_selection')
    for mask in product([False,True],repeat=3):
        xx=[z[i] if mask[i] else x[i] for i in range(3)]
        zz=[x[i] if mask[i] else z[i] for i in range(3)]
        check(interaction(xx,zz,Q(2))==[-a if mask[i] else a for i,a in enumerate(w)],'interaction_full_antisymmetry')
    return {'B':[[str(a) for a in row] for row in b],'V':[[str(a) for a in row] for row in v],
            'principal_minors':minors,'conditional_mean':[str(a[0]) for a in mean],
            'true_nulls_1based':[2],'interaction_W':list(map(str,w))}

def conditional_nulls(law):
    """law[(x_tuple,y)] gives exact probability; test all conditional cells."""
    p=len(next(iter(law))[0]);result=[]
    for j in range(p):
        total=defaultdict(Q);joint=defaultdict(Q);left=defaultdict(Q);right=defaultdict(Q)
        xvalues=set();yvalues=set()
        for (x,y),weight in law.items():
            z=x[:j]+x[j+1:];total[z]+=weight;joint[(z,x[j],y)]+=weight
            left[(z,x[j])]+=weight;right[(z,y)]+=weight;xvalues.add(x[j]);yvalues.add(y)
        valid=True
        for z in total:
            for x,y in product(xvalues,yvalues):
                valid &= joint[(z,x,y)]*total[z]==left[(z,x)]*right[(z,y)]
        if valid:result.append(j)
    return result

def discrete_copy_checks():
    # A correlated, singular law: a fresh complete vector is not pairwise valid.
    law={((u,u,u),u):Q(1,2) for u in [0,1]}
    check(conditional_nulls(law)==[0,1,2],'singular_model_all_conditional_nulls')
    badrisk=Q(0)
    for u,v in product([0,1],repeat=2):
        w=[Q(int(u==u)-int(v==u))]*3
        badrisk+=false_fraction(select(w,Q(1,3))[1],{0,1,2})/4
    check(badrisk==Q(1,2),'independent_whole_copy_FDR_failure')
    # A valid non-Gaussian construction from independent symmetric pair tables.
    theta=[Q(1,3),Q(1,2),Q(2,3)];off=[Q(1,6),Q(1,4),Q(1,6)]
    tables=[{(0,0):1-a-r,(0,1):r,(1,0):r,(1,1):a-r} for a,r in zip(theta,off)]
    pairs={}
    for states in product([(0,0),(0,1),(1,0),(1,1)],repeat=3):
        weight=Q(1)
        for j,state in enumerate(states):weight*=tables[j][state]
        pairs[(tuple(a for a,b in states),tuple(b for a,b in states))]=weight
    check(sum(pairs.values())==1,'discrete_pair_law_normalized')
    for (x,z),weight in pairs.items():
        for mask in product([False,True],repeat=3):
            xx=tuple(z[i] if mask[i] else x[i] for i in range(3))
            zz=tuple(x[i] if mask[i] else z[i] for i in range(3))
            check(pairs[(xx,zz)]==weight,'discrete_all_subset_pair_exchangeability')
    risks={}
    for response in ['independent','xor']:
        joint=defaultdict(Q);scored=[]
        for (x,z),weight in pairs.items():
            outcomes=[(0,Q(1,2)),(1,Q(1,2))] if response=='independent' else [(x[0]^x[1],Q(1))]
            for y,mass in outcomes:
                prob=weight*mass;joint[(x,y)]+=prob
                w=tuple(Q(int(a==y)-int(b==y)) for a,b in zip(x,z));scored.append((w,prob))
        nulls=set(conditional_nulls(joint))
        check(nulls==({0,1,2} if response=='independent' else {2}),'discrete_response_conditional_targets')
        groups=defaultdict(lambda:defaultdict(Q))
        for w,prob in scored:
            active=tuple(j for j in sorted(nulls) if w[j])
            key=(tuple(map(abs,w)),tuple((j,w[j]) for j in range(3) if j not in nulls))
            groups[(key,active)][tuple(1 if w[j]>0 else -1 for j in active)]+=prob
        for (key,active),weights in groups.items():
            total=sum(weights.values())
            for signs in product([-1,1],repeat=len(active)):
                check(weights[signs]==total/2**len(active),'conditional_joint_sign_uniformity_from_full_law')
        for q in [Q(1,3),Q(1,2),Q(2,3)]:
            risk=sum((mass*false_fraction(select(w,q)[1],nulls) for w,mass in scored),Q(0))
            check(risk<=q,'discrete_modelx_actual_FDR');risks[response+' q='+str(q)]=str(risk)
    return {'invalid_independent_copy_FDR':str(badrisk),'valid_finite_model_risks':risks}

def full_binary_law(probabilities):
    law={}
    for bits in product([0,1],repeat=len(probabilities)):
        mass=Q(1)
        for bit,p in zip(bits,probabilities):mass*=p if bit else 1-p
        law[bits]=mass
    return law

def resampling_checks():
    background=(0,0,1,1)
    law=full_binary_law([Q(1,5),Q(1,5),Q(3,5),Q(3,5)])
    observed=background
    score=lambda x:sum(a==b for a,b in zip(x,background))
    tail=sum((mass for x,mass in law.items() if score(x)>=score(observed)),Q(0))
    selected={x:mass for x,mass in law.items() if sum(x)==2};normalizer=sum(selected.values())
    conditioned={x:mass/normalizer for x,mass in selected.items()}
    check(tail==Q(144,625) and normalizer==Q(244,625),'nonuniform_background_full_probability')
    check(conditioned[observed]==Q(36,61),'extra_count_conditional_tail')
    check(sorted(conditioned.values())==[Q(1,61)]+[Q(6,61)]*4+[Q(36,61)],'six_point_conditional_kernel')
    for x in selected:
        check(conditioned[x]*normalizer==law[x],'conditional_kernel_reconstruction')
    # Original page's stronger uniform-permutation error, conditional and global.
    wrong_conditional=Q(9,10)**4;wrong_global=Q(3,8)*wrong_conditional
    check(wrong_conditional==Q(6561,10000) and wrong_global==Q(19683,80000)>Q(1,5),'incorrect_permutation_unconditional_error_lower')
    simulated=[4,4,3,4,2,3,4,4,3]
    p=Q(1+sum(a>=4 for a in simulated),len(simulated)+1)
    check(p==Q(3,5),'specified_monte_carlo_record')
    laws=[{1:Q(1)},{0:Q(1,5),1:Q(3,5),2:Q(1,5)}]
    for distribution in [law,conditioned,full_binary_law([Q(1,10),Q(1,10),Q(9,10),Q(9,10)])]:
        scores=defaultdict(Q)
        for x,mass in distribution.items():scores[score(x)]+=mass
        laws.append(dict(scores))
    # Integrate over T0 and the exact binomial count of >= T0 among M draws.
    for scores in laws:
        for m in [1,3,7,19]:
            for k in range(m+2):
                reject=Q(0)
                for value,mass in scores.items():
                    exceed=sum(p for v,p in scores.items() if v>=value)
                    conditional=sum((Q(comb(m,b))*exceed**b*(1-exceed)**(m-b)
                                     for b in range(min(m,k-1)+1)),Q(0))
                    reject+=mass*conditional
                check(reject<=Q(k,m+1),'exact_monte_carlo_rank_level_for_discrete_laws')
    # Pure exchangeable rank counting, every three-valued array of length <= 7.
    for n in range(2,8):
        for values in product(range(3),repeat=n):
            ranks=[sum(v>=a for v in values) for a in values]
            histogram=Counter(ranks);count=0
            for k in range(1,n+1):
                count+=histogram[k]
                check(count<=k,'all_tied_rank_counting_arrays')
    check(Q(1,20)==Q(1,19+1),'minimum_monte_carlo_resolution')
    # A frozen memory score equals 1 only at the observed continuous point;
    # the almost-sure zero match probability is proved in the page, not simulated.
    return {'unconditional_background_tail':str(tail),'count_probability':str(normalizer),
            'count_conditioned_tail':str(conditioned[observed]),
            'count_expected_rejection_sampling_attempts':str(1/normalizer),
            'six_point_kernel':[{ 'ones_1based':[i+1 for i,a in enumerate(x) if a], 'probability':str(mass)} for x,mass in conditioned.items()],
            'incorrect_permutation_error_lower':str(wrong_global),'monte_carlo_record':str(p)}

def repeated_run_checks():
    one=sum((Q(1,16)*false_fraction(select(signs,Q(1,4))[1],set(range(4)))
             for signs in product([-1,1],repeat=4)),Q(0))
    check(one==Q(1,16),'one_allnull_run_FDR')
    selected_risk=Q(0)
    for successes in product([False,True],repeat=5):
        weight=Q(1)
        for yes in successes:weight*=one if yes else 1-one
        if any(successes):selected_risk+=weight
    check(selected_risk==1-Q(15,16)**5==Q(289201,1048576)>Q(1,4),'five_valid_runs_postselection_FDR')
    return {'one_run_FDR':str(one),'five_run_report_FDR':str(selected_risk)}

def main():
    CHECKS.clear()
    parser=ArgumentParser(description=__doc__);parser.add_argument('--output',type=Path);args=parser.parse_args()
    results={'fixed_design':fixed_design_checks(),'symbolic_sign_budget':sign_budget_checks(),
             'nonlinear_gaussian':nonlinear_model_checks(),'discrete_modelx':discrete_copy_checks(),
             'conditional_resampling':resampling_checks(),'repeated_runs':repeated_run_checks()}
    output={'status':'PASS','checks':sum(CHECKS.values()),'groups':dict(CHECKS),'results':results,
            'scope':'Exact finite matrix/distribution/rank certificates. No simulation, estimated design law, or unproved numerical tolerance. General theorems are proved in the companion pages.'}
    text=json.dumps(output,ensure_ascii=False,indent=2)+'\n'
    if args.output:
        args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(text)
    print(text,end='')

if __name__=='__main__':
    main()
