# Independent exact check of the pseudoreplication benchmark: beta-binomial donors, arm totals by
# convolution, two-sided Fisher (probability ordering, ties included) at 1/20, and the known-null oracle.
from fractions import Fraction as F
from math import comb
import itertools, sys

def betabin(k, a):
    # P(X=x) = C(k,x) B(x+a, k-x+a) / B(a,a), as exact fractions via products.
    def B(p, q):  # Beta(p,q) for positive integers, as a Fraction
        from math import factorial
        return F(factorial(p - 1) * factorial(q - 1), factorial(p + q - 1))
    return [comb(k, x) * B(x + a, k - x + a) / B(a, a) for x in range(k + 1)]

def arm_total(m, per_donor):
    dist = [F(1)]
    for _ in range(m):
        new = [F(0)] * (len(dist) + len(per_donor) - 1)
        for i, p in enumerate(dist):
            for j, q in enumerate(per_donor):
                new[i + j] += p * q
        dist = new
    return dist

def fisher_rejects(n, A, B):
    t = A + B
    w = lambda x: comb(n, x) * comb(n, t - x)
    obs = w(A); tot = comb(2 * n, t)
    p = F(sum(w(x) for x in range(max(0, t - n), min(n, t) + 1) if w(x) <= obs), tot)
    return p <= F(1, 20)

def type1(m, k, per_donor):
    n = m * k; d = arm_total(m, per_donor)
    naive = sum(d[A] * d[B] for A in range(n + 1) for B in range(n + 1) if fisher_rejects(n, A, B))
    # oracle: tail P(|A'-B'| >= |A-B|) under the clustered null
    diff = {}
    for A in range(n + 1):
        for B in range(n + 1):
            diff[abs(A - B)] = diff.get(abs(A - B), 0) + d[A] * d[B]
    tail = {}; run = F(0)
    for v in sorted(diff, reverse=True):
        run += diff[v]; tail[v] = run
    oracle = sum(diff[v] for v in diff if tail[v] <= F(1, 20))
    return naive, oracle

sel = betabin(20, 5)
naive, oracle = type1(5, 20, sel)
print("selected m=5 k=20 a=5: Fisher", float(naive), "oracle", float(oracle))
print("one observation per donor (m=5, k=1, a=5): Fisher", float(type1(5, 1, betabin(1, 5))[0]))
print("weaker dependence (m=5, k=20, a=50): Fisher", float(type1(5, 20, betabin(20, 50))[0]))
# Full grid
above = 0; oracle_ok = 0; indep_ok = 0; total = 0; mx = F(0)
for m in (3, 5, 10):
    for k in (1, 2, 5, 10, 20):
        designs = [("beta", a, betabin(k, a)) for a in (50, 5, 1)]
        designs.append(("independent", None, [F(comb(k, x), 2 ** k) for x in range(k + 1)]))
        designs.append(("perfect_copy", None, [F(1, 2)] + [F(0)] * (k - 1) + [F(1, 2)]))
        for kind, a, pd in designs:
            nv, orc = type1(m, k, pd); total += 1
            above += nv > F(1, 20); oracle_ok += orc <= F(1, 20); mx = max(mx, nv)
            if kind == "independent": indep_ok += nv <= F(1, 20)
print("designs", total, "Fisher above nominal", above, "oracle at or below nominal", oracle_ok, "independent controls at or below", indep_ok, "max Fisher", float(mx))
