#!/usr/bin/env python3
"""phylo-evidence.py -- exact Bayesian evidence for any four-taxon alignment
(Jukes-Cantor 1969, exponential branch-length prior, alpha = 1).

Give it four aligned nucleotide sequences (FASTA or NEXUS) or the 15
site-pattern-class counts directly; it returns, for each of the three
unrooted quartet topologies T in {12|34, 13|24, 14|23},

    Z(T; u)  as an EXACT rational p/q (full integer strings + digit counts),
    log Z(T) to the requested --dps,
    and the three exact-rational Bayes factors with their logs.

DETERMINISM (headline feature).  The answer path is exact integer
arithmetic end to end: no sampling, no seeds, no floating point.  The
evidence integral Z = int_{[0,1]^5} prod_k p_k(x)^{u_k} dx is a rational
number (PhyloNote, Lemma 1): its denominator divides
256^N * prod_{e=1}^{5} prod_{j=0}^{N} (j+alpha), so for alpha = 1 the
integer D = 256^N * lcm(1..N+1)^5 clears it.  Z is computed modulo a
prime stack whose SIZE IS FIXED A PRIORI by that Lemma-1 bound, and the
exact fraction is reassembled by the Chinese remainder theorem.  Identical
input gives bit-identical output on any machine (results go to stdout;
progress/timing go to stderr and never affect the result block).

SELF-CERTIFICATION (every run).  The prime stack is split into two
DISJOINT halves; each half must reconstruct the identical rational
independently (disjoint-CRT gate, printed per topology; any mismatch is a
hard failure, exit code 3).  `--verify` adds an independent float64
Gauss-Legendre quadrature cross-check (Felsenstein pruning code path,
never touching the exact expansion) at float-attainable digits.
`--selftest` recomputes the stored references from vendored count vectors
(phylo-evidence-selftest.json, shipped alongside) and byte-compares the
full integer strings: the 20-site synthetic benchmark (49/87-digit
fraction, log BF 3.18823969...), the Rfam RF00005 mt tRNA-Gln quartet
(188/302-digit fraction, log BFs 14.5368... / 15.5563...), and two
recorded real-data battery rows (hominid mt tRNA-Lys N=70, tRNA-Gln
anchor-control N=68) recomputed through the consa engine.

ENGINES (validated math only; the tool prints which engine ran and why).
  dense     small N (wired for N <= 30): dense expansion of the
            5-variable coefficient tensor mod p ((N+1)^5 lattice),
            multi-prime CRT.  Measured limits: memory ~3*(N+1)^5*8 B
            per prime (refuses near N~65 under a 31 GiB address cap;
            N=120 would need ~209 GB), and serial full-exact wall ~2 h
            already at N~40 -- hence the conservative N <= 30 ceiling.
  parseval  mid N (N <= 120 hard ceiling): streamed DFT-dual (Parseval)
            engine, O(L) vectors + O(L^3) grids per prime, L = N+1;
            kills the memory wall (measured <= ~0.3 GB per prime),
            compute-bound at ~L^5 per prime.
  consa     Construction-A backward-cone tensor split (engine from the
            original real-data analysis, 2026-07-08).  Splits the counts u into
            u_B on the two collapsible classes {xxxx, xxyy} (3-D latent
            coefficient tensor V, side n+1 = sum(u_B)+1) and u_A = the
            13 remaining classes (dense 5-D core, side nA+1 =
            sum(u_A)+1), contracted through b-shifted modified-I
            integral ladders and three mod-p GEMMs (Z = 4^-(4*nA+6*n) *
            sum_beta QA(beta) sum_ABC V(A,B,C) * prod_e Itil(.;b_e)).
            Cost is driven by nA, NOT by N: real quartet alignments
            concentrate counts on xxxx/xxyy, so gene-length data with a
            small non-collapsible remainder becomes exact-affordable.
            AUTO-SELECTED when the engine's own split gate passes on
            all three topology-permuted count vectors (the gate is the
            engine's own validated limits: (nA+1)^2 <= 2^16 GEMM inner
            dimension; (n+1)^3 < 2^32 final-sum lattice; QA-core memory
            3*(nA+1)^5*8 B per worker inside the process address
            budget); otherwise it REFUSES BY NAME on stderr and the
            tool falls back to the direct engines above, which stay
            wired UNCHANGED as the reference path.  Validation record
            (original real-data analysis): byte-identical num/den reproduction
            of all 8 recorded exact battery rows (N=63..75, both disjoint
            CRT sets, all 3 topologies; 8.6x aggregate faster than the
            direct engines, 1.7x-1270x per row) and the production row
            Rfam RF00001 (N=116, gate_all=true).  Wired ceiling
            N <= 390, the original analysis's production regime (lysozyme-
            class genes); beyond it the tool refuses by name.  numba is
            OPTIONAL: if importable it accelerates the 5-D core
            convolution (fused sparse gather); otherwise a pure-numpy
            path (lazy hi/lo-split modular convolution + exact
            float64-BLAS 16-bit-split GEMM, partial sums < 2^48 < 2^53)
            computes the identical residues (gated byte-identical; set
            PHYLO_EVIDENCE_NO_NUMBA=1 to force the numpy path).  The
            engine line names the backend, e.g. "consa[numba]".
Measured cost (CPU-seconds per prime PER TOPOLOGY; power-law rate fits
from a dedicated scaling-ladder measurement of these engines
(2026-07-07), timed points N in {20,30,40,60,68,120}, three points
re-measured independently from scratch with bit-identical residues;
multiply by 2*n_half(N) = 2*((bits(D)+8)//30+2) ~= 1.02*N+4 primes per
topology and x3 topologies for a full --selftest or --verify-grade run;
jobs are embarrassingly parallel over primes):
    t_prime(N) ~= 2.8e-7 * (N+1)^5 CPU-s  (parseval, idle-box fit)
    t_prime(N) ~= 3.6e-7 * N^5.23 CPU-s   (same engine re-fit on a
    heavily contended box -- the pessimistic envelope)
    N=20: ~2.3 CPU-s/prime x 24 primes/topology (seconds total);
    N=40: ~87 CPU-s x 44/topology;
    N=68: ~1.4e3 CPU-s x 72/topology (~28 CPU-h/topology, ~84 CPU-h for
    all 3 topologies of --selftest's tRNA-Gln leg) -- an EARLIER estimate
    in this file (460 CPU-s/prime, "minutes on a many-core box") was
    stale/too low by ~3x; do not trust it;
    N=120 (the refusal ceiling): ~7.3e3 CPU-s x 126/topology, ~7.7e2
    CPU-h for all 3 topologies even by the idle-box fit.  The ceiling
    is a wall, not a comfort zone: the practical "few hours on a few
    dozen cores" reach ends near N~90 (~1.4e2 CPU-h all 3 topologies).
On a machine with P truly-idle cores divide CPU-h by P for wall time
(e.g. the tRNA-Gln leg is ~1.75 h at P=48); on a shared/contended
machine measure your own per-process CPU efficiency first --
contention can cut it 5-10x.
The N~890 full-alignment regime is OUT OF REACH for the direct engines
(~10^17 ops, i.e. ~10^11-10^12 CPU-s by the same fits); the consa
engine reaches gene-length alignments whenever the non-collapsible
remainder nA is small (its cost driver; measured 296 CPU-s/prime at
N=116/nA=38, 1799 CPU-s/prime at N=390/nA=39-46, original-analysis
rate probes on a contended 96-core box), but data whose counts spread
over many non-collapsible classes (e.g. the brown-1982 and
hayasaka-1988 primate benchmarks, nA ~ 200) exceed its own gate and
are STILL refused honestly.  The order-14 contiguity-recurrence
transport recorded in the original analysis passes that wall in minutes but
applies to 1-parameter count families only (validated to n = 5000),
and its large-N seeding for general 15-class data is an open gap, so
it is NOT wired in here.  Beyond the validated ceilings the tool
refuses by name rather than silently grinding.

SCOPE.  JC69 substitution model + exponential branch-length prior with
shape alpha = 1 only (that is what these engines are validated on:
exact equality vs an independent dense engine at N=20/40, two disjoint
prime stacks at N=68).  `--alpha` is exposed and validated for input, but
any value other than 1 is refused by name (AlphaScopeError); for small-N
rational alpha see the companion phylo-evaluate.py --point.  GTR and
+Gamma rate heterogeneity are exact-capable in principle (PhyloNote
Lemma 2 resolvent identity and the Euler-Gompertz closure, Sec. 7.1) but
are OUT OF SCOPE in v1 -- this stub note is the only trace of them here.

INPUT.  FASTA or NEXUS (--nexus), exactly four taxa (a named error lists
the taxa found otherwise); columns containing anything outside A/C/G/T
(after U->T) are dropped by default and the count of dropped columns is
reported (--keep-policy is reserved for future policies).  Or
--pattern-counts 'xxxx:35,xxyy:8,...' on the 15 canonical first-appearance
labels.  The output block echoes the counts and a SHA-256 of the
canonical count vector (the exact input to the mathematics); the raw
input-file SHA-256 is reported on stderr as provenance.

EXIT CODES.  0 ok; 3 disjoint-CRT gate mismatch; 4 selftest mismatch;
5 --verify quadrature disagreement; 10 taxa-count error; 11 input-format
error; 12 alpha out of validated scope; 13 beyond engine ceiling;
14 missing dependency.  Dependencies: Python 3.9+, mpmath, numpy (all
pip-installable; nothing else).

USAGE
  phylo-evidence.py alignment.fasta [--dps 50] [--verify] [--procs P]
  phylo-evidence.py alignment.nex --nexus
  phylo-evidence.py --pattern-counts 'xxxx:10,xxyy:4,xyxy:2,xyyx:2,xxxy:1,xyzw:1'
  phylo-evidence.py --selftest [--quick]     # stored-reference reproduction
  phylo-evidence.py --selftest --anchor      # tRNA-Gln anchor quartet only,
                                             # consa engine leg only
"""
import argparse
import hashlib
import itertools
import json
import math
import os
import re
import sys
from fractions import Fraction

# keep worker processes single-threaded (they are parallel over primes)
for _v in ('OPENBLAS_NUM_THREADS', 'MKL_NUM_THREADS', 'OMP_NUM_THREADS',
           'NUMEXPR_NUM_THREADS'):
    os.environ.setdefault(_v, '1')

try:
    import mpmath as mp
except ImportError:  # pragma: no cover
    sys.stderr.write("ERROR[DependencyError]: mpmath is required "
                     "(pip install mpmath)\n")
    sys.exit(14)
try:
    import numpy as np
except ImportError:  # pragma: no cover
    sys.stderr.write("ERROR[DependencyError]: numpy is required for the "
                     "exact engines (pip install numpy)\n")
    sys.exit(14)

VERSION = "1.1 (2026-07-08)"
DENSE_NMAX = 30        # dense-tensor engine: (N+1)^5 int64 lattice per prime
PARSEVAL_NMAX = 120    # hard validated/practical ceiling, see --help text
CONSA_NMAX = 390       # consa engine: original-analysis production ceiling
CONSA_COLLAPSIBLE = ('xxxx', 'xxyy')   # u_B classes (3-D collapsible atoms)
# measured interpreter+numpy+numba-JIT address-space footprint on the
# reference box (single-threaded BLAS): < 0.4 GiB; budgeted at 1 GiB.
CONSA_SLACK_BYTES = 1 << 30
U64 = np.uint64

# exit codes
RC_GATE, RC_SELFTEST, RC_VERIFY = 3, 4, 5
RC_TAXA, RC_FORMAT, RC_ALPHA, RC_CEILING, RC_DEP = 10, 11, 12, 13, 14


class CLIError(Exception):
    name, rc = "CLIError", 1


class TaxaError(CLIError):
    name, rc = "TaxaError", RC_TAXA


class InputFormatError(CLIError):
    name, rc = "InputFormatError", RC_FORMAT


class AlphaScopeError(CLIError):
    name, rc = "AlphaScopeError", RC_ALPHA


class EngineCeilingError(CLIError):
    name, rc = "EngineCeilingError", RC_CEILING


def log(msg):
    """Progress/provenance -> stderr (never part of the result block)."""
    sys.stderr.write(msg + "\n")
    sys.stderr.flush()


def out(msg=""):
    sys.stdout.write(msg + "\n")
    sys.stdout.flush()


# ===========================================================================
# Pattern classes (15 = Bell(4) orbits of nucleotide relabelling)
# ===========================================================================
def canonical_label(col):
    seen, o = {}, []
    for c in col:
        if c not in seen:
            seen[c] = 'xyzw'[len(seen)]
        o.append(seen[c])
    return ''.join(o)


ALL_LABELS = sorted({canonical_label(s)
                     for s in itertools.product('ACGT', repeat=4)})
assert len(ALL_LABELS) == 15

TOPOLOGIES = ('12|34', '13|24', '14|23')
# tau[i] = which taxon (1-based) is read at leaf i of the reference 12|34 tree
TOPO_TAU = {'12|34': (1, 2, 3, 4), '13|24': (1, 3, 2, 4), '14|23': (1, 4, 3, 2)}


def permute_counts(u, tau):
    """Counts u' with Z(tau-topology; u) = Z(12|34; u') (validated identity)."""
    o = {}
    for lab, cnt in u.items():
        rep = tuple('xyzw'.index(ch) for ch in lab)
        s2 = tuple(rep[tau[i] - 1] for i in range(4))
        lab2 = canonical_label(s2)
        o[lab2] = o.get(lab2, 0) + cnt
    return o


# ---- W-side count atoms g_k (Parseval engine; b3_5d_smart.g_kernel) ------
def g_support(lab):
    """g_k(y) = sum_{a,b in {0..3}^2} y^e, e = ([a=s1],[a=s2],[b=s3],[b=s4],
    [a=b]) -- as a list of (e-bits, multiplicity)."""
    s = tuple('xyzw'.index(ch) for ch in lab)
    K = {}
    for a in range(4):
        for b in range(4):
            e = (int(a == s[0]), int(a == s[1]),
                 int(b == s[2]), int(b == s[3]), int(a == b))
            K[e] = K.get(e, 0) + 1
    return sorted(K.items())


# ---- multilinear 256*p_k kernels (dense engine; Felsenstein pruning) -----
def build_kernels():
    """Independent derivation of the pattern polynomials by pruning:
    P_same = (1+3x)/4, P_diff = (1-x)/4, uniform root, sum over the two
    internal states.  Returns {label: {bits in {0,1}^5: int}} for 256*p_k."""
    def P(i, j):
        return (Fraction(1, 4), Fraction(3, 4)) if i == j else \
               (Fraction(1, 4), Fraction(-1, 4))
    kernels = {}
    for s in itertools.product(range(4), repeat=4):
        lab = canonical_label(s)
        if lab in kernels:
            continue
        poly = {}
        for a in range(4):
            for b in range(4):
                fac = [P(a, s[0]), P(a, s[1]), P(b, s[2]), P(b, s[3]), P(a, b)]
                for bits in itertools.product((0, 1), repeat=5):
                    c = Fraction(1, 4)
                    for e in range(5):
                        c *= fac[e][bits[e]]
                    poly[bits] = poly.get(bits, Fraction(0)) + c
        k = {}
        for bits, c in poly.items():
            v = c * 256
            assert v.denominator == 1
            if v:
                k[bits] = int(v)
        kernels[lab] = k
    assert len(kernels) == 15
    return kernels


KERNELS = build_kernels()


# ===========================================================================
# Number theory (stdlib only; deterministic)
# ===========================================================================
def is_prime(n):
    """Deterministic Miller-Rabin for n < 3.3e24."""
    if n < 2:
        return False
    for p in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37):
        if n % p == 0:
            return n == p
    d, s = n - 1, 0
    while d % 2 == 0:
        d //= 2
        s += 1
    for a in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37):
        x = pow(a, d, n)
        if x in (1, n - 1):
            continue
        for _ in range(s - 1):
            x = x * x % n
            if x == n - 1:
                break
        else:
            return False
    return True


def factorize(n):
    f, d = set(), 2
    while d * d <= n:
        while n % d == 0:
            f.add(d)
            n //= d
        d += 1
    if n > 1:
        f.add(n)
    return sorted(f)


def find_parseval_primes(L, count, pmax=2**31 - 1):
    """31-bit primes p = c*L+1 (descending c) with a verified primitive
    L-th root of unity mod p."""
    o = []
    fac = factorize(L)
    c = (pmax - 1) // L
    while len(o) < count and c > 0:
        p = c * L + 1
        if is_prime(p):
            for h in range(2, 200):
                w = pow(h, (p - 1) // L, p)
                if w != 1 and pow(w, L, p) == 1 and \
                   all(pow(w, L // q, p) != 1 for q in fac):
                    o.append((p, w))
                    break
        c -= 1
    if len(o) < count:
        raise InputFormatError(f"only {len(o)} Parseval primes for L={L}")
    return o


def find_plain_primes(count, start=2**31 - 1):
    o, c = [], start
    while len(o) < count:
        if is_prime(c):
            o.append(c)
        c -= 2
    return o


def crt_list(rems, mods):
    R, M = rems[0], mods[0]
    for r, m in zip(rems[1:], mods[1:]):
        g, s, t = _xgcd(M, m)
        assert g == 1
        R, M = (R * t * m + r * s * M) % (M * m), M * m
    return R, M


def _xgcd(a, b):
    x0, x1, y0, y1 = 1, 0, 0, 1
    while b:
        q, a, b = a // b, b, a % b
        x0, x1 = x1, x0 - q * x1
        y0, y1 = y1, y0 - q * y1
    return a, x0, y0


def lemma1_denominator(N):
    """D with D*Z(T;u,alpha=1) in Z and |D*Z| < D.

    A PRIORI prime count: PhyloNote Lemma 1 (label lem:rational) proves
    Z's denominator divides 256^N * prod_{e=1}^5 prod_{j=0}^N (j+alpha);
    for alpha = 1 each edge integral is a sum of fractions 1/(a_e+1),
    a_e <= N, so the per-edge denominator divides lcm(1..N+1) and
    D = 256^N * lcm(1..N+1)^5 clears Z (recorded refinement, identical to
    the original engines).  The number of 30-bit CRT primes needed is
    fixed in advance by bits(D) -- nothing about the prime stack depends
    on the data values themselves."""
    lc = 1
    for k in range(2, N + 2):
        lc = lc * k // math.gcd(lc, k)
    return (256 ** N) * (lc ** 5)


def nprimes_half(D):
    """Primes per CRT half-stack, a priori from the Lemma-1 bound."""
    return (D.bit_length() + 8) // 30 + 2


# ===========================================================================
# Engine 1: dense 5D tensor mod p (small N; port of the original b4 engine)
# ===========================================================================
def conv5d_mod(Q, K, p):
    s = tuple(d + 1 for d in Q.shape)
    o = np.zeros(s, dtype=np.int64)
    for eps, c in K.items():
        c = int(c) % p
        if c == 0:
            continue
        sl = tuple(slice(e, e + Q.shape[i]) for i, e in enumerate(eps))
        o[sl] = (o[sl] + c * Q) % p
    return o


def integrate_mod(Q, p):
    s = Q.shape
    acc = Q
    for d in range(5):
        inv = np.array([pow(j + 1, p - 2, p) for j in range(s[d])],
                       dtype=np.int64)
        shape = [1] * 5
        shape[d] = s[d]
        acc = (acc * inv.reshape(shape)) % p
    return int(acc.sum() % p)


def _dense_worker(args):
    u, N, p = args
    Q = np.ones((1,) * 5, dtype=np.int64)
    for lab, e in sorted(u.items()):
        K = KERNELS[lab]
        for _ in range(e):
            Q = conv5d_mod(Q, K, p)
    I = integrate_mod(Q, p)
    Zp = I * pow(pow(256, N, p), p - 2, p) % p
    return p, Zp


# ===========================================================================
# Engine 2: streamed Parseval (DFT-dual) mod p (mid N; original parseval port)
# ===========================================================================
def J_raw_mod(N, p):
    """J(m) = int_0^1 (1+3x)^m (1-x)^{N-m} dx mod p;
    J(m) = (3m J(m-1)+1)/(N-m+1), J(0) = 1/(N+1)."""
    inv = [0] + [pow(j, p - 2, p) for j in range(1, N + 2)]
    J = np.zeros(N + 1, dtype=np.int64)
    J[0] = inv[N + 1]
    for m in range(1, N + 1):
        J[m] = (3 * m % p) * int(J[m - 1]) % p
        J[m] = (int(J[m]) + 1) * inv[N - m + 1] % p
    return J


def _omega_powers(w, L, p):
    pw = np.zeros(L, dtype=U64)
    pw[0] = 1
    for j in range(1, L):
        pw[j] = int(pw[j - 1]) * w % p
    return pw


def _Jhat_mod(J, pw, L, p):
    """Jhat[j] = sum_m J[m] w^{-jm} (inverse root: Parseval pairing)."""
    N = len(J) - 1
    Jh = np.zeros(L, dtype=U64)
    Ju = J.astype(U64)
    pp = U64(p)
    for j in range(L):
        idx = (-j * np.arange(N + 1)) % L
        Jh[j] = int((Ju * pw[idx] % pp).sum() % pp)
    return Jh


def _powmod_grid(base, expo, p):
    pp = U64(p)
    r = None
    b = base.copy()
    while expo:
        if expo & 1:
            r = b.copy() if r is None else (r * b % pp)
        expo >>= 1
        if expo:
            b = b * b % pp
    return r if r is not None else np.ones_like(base)


def _parseval_worker(args):
    u, N, L, p, w, jrows = args
    pw = _omega_powers(w, L, p)
    J = J_raw_mod(N, p)
    Jh = _Jhat_mod(J, pw, L, p)
    pp = U64(p)
    a3 = pw.reshape(-1, 1, 1)
    a4 = pw.reshape(1, -1, 1)
    a5 = pw.reshape(1, 1, -1)
    T = [None] * 8   # tail grids indexed by t = e3*4 + e4*2 + e5
    T[0] = np.ones((L, L, L), dtype=U64)
    T[1] = np.broadcast_to(a5, (L, L, L)).copy()
    T[2] = np.broadcast_to(a4, (L, L, L)).copy()
    T[3] = a4 * a5 % pp
    T[4] = np.broadcast_to(a3, (L, L, L)).copy()
    T[5] = a3 * a5 % pp
    T[6] = a3 * a4 % pp
    T[7] = T[6] * a5 % pp
    W3 = Jh.reshape(-1, 1, 1) * Jh.reshape(1, -1, 1) % pp
    W3 = W3 * Jh.reshape(1, 1, -1) % pp
    pats = [(g_support(lab), e) for lab, e in sorted(u.items()) if e > 0]
    total = 0
    rows = range(L) if jrows is None else jrows
    for j1 in rows:
        z1 = int(pw[j1])
        for j2 in range(L):
            z2 = int(pw[j2])
            prod = None
            for sup, ek in pats:
                D8 = [0] * 8
                for e, c in sup:
                    v = c
                    if e[0]:
                        v = v * z1 % p
                    if e[1]:
                        v = v * z2 % p
                    D8[e[2] * 4 + e[3] * 2 + e[4]] += v
                g = None
                for t in range(8):
                    d = D8[t] % p
                    if d == 0:
                        continue
                    term = U64(d) * T[t] % pp
                    g = term if g is None else (g + term) % pp
                if g is None:
                    g = np.zeros((L, L, L), dtype=U64)
                gp = _powmod_grid(g, ek, p)
                prod = gp if prod is None else (prod * gp % pp)
            w12 = int(Jh[j1]) * int(Jh[j2]) % p
            s3 = int((prod * W3 % pp).sum() % pp)
            total = (total + w12 * s3) % p
    invL5 = pow(pow(L, 5, p), p - 2, p)
    S = total * invL5 % p
    Zp = S * pow(pow(4096, N, p), p - 2, p) % p
    return p, Zp


# ===========================================================================
# Engine 3: Construction-A backward-cone tensor split ("consa")
# (port of the original real-data analysis engine, 2026-07-08;
#  validated 8/8 recorded exact rows +
#  RF00001 production row byte-identical -- see the ENGINES note above.
#  Self-contained: atoms/kernels rebuilt from this file's own g_support /
#  KERNELS derivations, byte-equality gated against the original engine.)
# ===========================================================================
def _consa_atom3d(lab):
    """Collapsed 2x2x2 atom of g_lab on (Y=y1y2, Z=y3y4, y5).
    Only defined for the CONSA_COLLAPSIBLE classes (e1==e2, e3==e4 on
    every monomial); total mass is 16 (4x4 latent states per site)."""
    A = np.zeros((2, 2, 2), dtype=np.int64)
    for e, c in g_support(lab):
        if not (e[0] == e[1] and e[2] == e[3]):
            raise EngineCeilingError(
                f"consa: class {lab} is not collapsible at monomial {e}")
        A[e[0], e[2], e[4]] += c
    assert int(A.sum()) == 16
    return A


def _consa_V_array(c_xxxx, c_xxyy, p):
    """3-D coeff array of g~_xxxx^c1 * g~_xxyy^c2 mod p, shape (n+1,)^3.
    int64-safe: atom coeffs <= 8, V < p < 2^31 -> terms < 2^37."""
    V = np.ones((1, 1, 1), dtype=np.int64)
    for lab, c in (('xxxx', c_xxxx), ('xxyy', c_xxyy)):
        A = _consa_atom3d(lab)
        for _ in range(c):
            ns = tuple(d + 1 for d in V.shape)
            o = np.zeros(ns, dtype=np.int64)
            for e in itertools.product((0, 1), repeat=3):
                c8 = int(A[e])
                if c8 == 0:
                    continue
                sl = tuple(slice(ei, ei + V.shape[i])
                           for i, ei in enumerate(e))
                o[sl] += c8 * V
            V = o % p
    return V


def _consa_Itable(n, bmax, p):
    """Itil_n(m;b) = int_0^1 x^b (1+3x)^m (1-x)^{n-m} dx mod p, shape
    (n+1, bmax+1), by the b-shifted IBP ladder (original construction,
    validated there against direct Fraction integration):
      Itil(0;0) = 1/(n+1);  Itil(0;b) = b/(n+b+1) * Itil(0;b-1)
      Itil(m;0) = (3m*Itil(m-1;0) + 1)/(n-m+1)
      Itil(m;b) = (b*Itil(m;b-1) + 3m*Itil(m-1;b))/(n-m+1+b)."""
    inv = np.zeros(n + bmax + 2, dtype=np.int64)
    for j in range(1, n + bmax + 2):
        inv[j] = pow(j, p - 2, p)
    T = np.zeros((n + 1, bmax + 1), dtype=np.int64)
    T[0, 0] = inv[n + 1]
    for b in range(1, bmax + 1):
        T[0, b] = b * T[0, b - 1] % p * inv[n + b + 1] % p
    for m in range(1, n + 1):
        T[m, 0] = (3 * m * T[m - 1, 0] + 1) % p * inv[n - m + 1] % p
        for b in range(1, bmax + 1):
            T[m, b] = (b * T[m, b - 1] + 3 * m * T[m - 1, b]) % p \
                * inv[n - m + 1 + b] % p
    return T


def _consa_mm_mod(A, B, p, blk=16_000_000):
    """Exact (A @ B) mod p via 16-bit 4-way split + float64 BLAS dgemm.
    Inputs reduced mod p < 2^31; halves < 2^16 -> partial sums over an
    inner dim <= 2^16 stay < 2^48 << 2^53 (exact in float64).  B's
    columns are chunked to cap the float64 copies."""
    assert A.shape[-1] == B.shape[0] and A.shape[-1] <= (1 << 16)
    a1 = (A >> 16).astype(np.float64)
    a0 = (A & 0xFFFF).astype(np.float64)
    o = np.empty((A.shape[0], B.shape[1]), dtype=np.int64)
    step = max(1, blk // max(1, B.shape[0]))
    for j in range(0, B.shape[1], step):
        Bj = B[:, j:j + step]
        b1 = (Bj >> 16).astype(np.float64)
        b0 = (Bj & 0xFFFF).astype(np.float64)
        h11 = np.dot(a1, b1).astype(np.int64) % p
        h10 = (np.dot(a1, b0) + np.dot(a0, b1)).astype(np.int64) % p
        h00 = np.dot(a0, b0).astype(np.int64) % p
        o[:, j:j + step] = \
            (h11 * ((1 << 32) % p) + h10 * ((1 << 16) % p) + h00) % p
    return o


def _consa_conv5d_lazy(Q, K, p):
    """One convolution step Q * (256*p_k) mod p over the dict kernel K,
    lazy hi/lo accumulation (<= 32 taps, c0 < 2^16, Q < p < 2^31 ->
    32 partial terms < 2^52, int64-safe; 2 modulos per cell)."""
    s = tuple(d + 1 for d in Q.shape)
    lo = np.zeros(s, dtype=np.int64)
    hi = np.zeros(s, dtype=np.int64)
    buf = np.empty(Q.shape, dtype=np.int64)
    for eps, c in K.items():
        c = int(c) % p
        if c == 0:
            continue
        sl = tuple(slice(e, e + Q.shape[i]) for i, e in enumerate(eps))
        c1, c0 = c >> 16, c & 0xFFFF
        if c0:
            np.multiply(Q, c0, out=buf)
            lo[sl] += buf
        if c1:
            np.multiply(Q, c1, out=buf)
            hi[sl] += buf
    lo %= p
    hi %= p
    return (lo + hi * ((1 << 16) % p)) % p


def _consa_QA_numpy(uA, p):
    """Dense 5-D coeff array of prod_{k in A} (256 p_k)^{u_k} mod p
    (pure-numpy backend)."""
    Q = np.ones((1,) * 5, dtype=np.int64)
    for lab, e in sorted(uA.items()):
        K = KERNELS[lab]
        for _ in range(e):
            Q = _consa_conv5d_lazy(Q, K, p)
    return Q


_CONSA_NUMBA = None    # None = undecided; False = unavailable/disabled


def _consa_numba_backend():
    """Optional numba fused sparse-gather for the 5-D core convolution.
    Returns the compiled gather kernel, or False (pure-numpy fallback,
    identical residues -- gated).  PHYLO_EVIDENCE_NO_NUMBA=1 disables."""
    global _CONSA_NUMBA
    if _CONSA_NUMBA is not None:
        return _CONSA_NUMBA
    if os.environ.get('PHYLO_EVIDENCE_NO_NUMBA'):
        _CONSA_NUMBA = False
        return False
    try:
        from numba import njit

        @njit(cache=False)
        def _gather5(Qp, offs, cs, o, p):
            # o[i] = (sum_t cs[t] * Qp[i - offs[t] + 1]) mod p; Qp is
            # zero-padded by 1 per axis.  Single int64 accumulator:
            # <= 16 taps, |c| <= 12, q < 2^31 -> |acc| < 2^39; numba
            # int64 % follows Python semantics -> non-negative.
            n0, n1, n2, n3, n4 = o.shape
            T = cs.shape[0]
            acc = np.empty(n4, dtype=np.int64)
            for i0 in range(n0):
                for i1 in range(n1):
                    for i2 in range(n2):
                        for i3 in range(n3):
                            acc[:] = 0
                            for t in range(T):
                                row = Qp[i0 - offs[t, 0] + 1,
                                         i1 - offs[t, 1] + 1,
                                         i2 - offs[t, 2] + 1,
                                         i3 - offs[t, 3] + 1]
                                c = cs[t]
                                off4 = 1 - offs[t, 4]
                                for i4 in range(n4):
                                    acc[i4] += c * np.int64(row[i4 + off4])
                            for i4 in range(n4):
                                o[i0, i1, i2, i3, i4] = np.int32(acc[i4] % p)

        # warm the JIT in the parent so forked pool workers inherit it
        _q = np.zeros((3,) * 5, dtype=np.int32)
        _q[(1,) * 5] = 1
        _o = np.empty((1,) * 5, dtype=np.int32)
        _gather5(_q, np.zeros((1, 5), dtype=np.int64),
                 np.ones(1, dtype=np.int64), _o, 2_147_483_629)
        _CONSA_NUMBA = _gather5
    except Exception as e:   # numba absent/broken -> graceful degradation
        log(f"# consa: numba unavailable ({type(e).__name__}); "
            f"using the pure-numpy backend (identical residues, slower)")
        _CONSA_NUMBA = False
    return _CONSA_NUMBA


def _consa_sparse(K):
    """Dict kernel -> (offsets, coeffs) arrays for the numba gather.
    The asserts are the source engine's own bounds (all 15 kernels)."""
    offs = np.array(sorted(K.keys()), dtype=np.int64)
    cs = np.array([int(K[tuple(o)]) for o in offs.tolist()], dtype=np.int64)
    assert np.abs(cs).max() <= 12 and cs.shape[0] <= 16
    return np.ascontiguousarray(offs), np.ascontiguousarray(cs)


def _consa_QA_numba(uA, p, gather):
    """numba fused sparse-gather QA build (int32, values < p < 2^31)."""
    Q = np.ones((1,) * 5, dtype=np.int32)
    for lab, e in sorted(uA.items()):
        offs, cs = _consa_sparse(KERNELS[lab])
        for _ in range(e):
            Qp = np.zeros(tuple(d + 2 for d in Q.shape), dtype=np.int32)
            Qp[(slice(1, -1),) * 5] = Q
            o = np.empty(tuple(d + 1 for d in Q.shape), dtype=np.int32)
            gather(Qp, offs, cs, o, p)
            Q = o
    return Q


def _consa_split(u):
    uB = {k: u.get(k, 0) for k in CONSA_COLLAPSIBLE if u.get(k, 0) > 0}
    uA = {k: v for k, v in u.items()
          if k not in CONSA_COLLAPSIBLE and v > 0}
    return uA, uB


def _address_budget(procs):
    """Per-worker address-space budget: RLIMIT_AS if set, else physical
    memory divided by the worker count.  Used only by the consa gate."""
    try:
        import resource
        soft, _ = resource.getrlimit(resource.RLIMIT_AS)
        if soft not in (resource.RLIM_INFINITY, -1):
            return int(soft)
    except Exception:
        pass
    try:
        phys = os.sysconf('SC_PAGE_SIZE') * os.sysconf('SC_PHYS_PAGES')
        return int(phys) // max(1, procs)
    except (ValueError, OSError, AttributeError):
        return 1 << 62   # undetectable; the engine's asserts still apply


def consa_gate(u, procs=1):
    """Construction-A applicability gate for ONE (topology-permuted)
    count vector -- the ENGINE'S OWN validated limits, nothing else:
      (i)   (nA+1)^2 <= 2^16   (GEMM inner-dimension bound, exactness
                                of the 16-bit-split float64 dgemm)
      (ii)  (n+1)^3  <  2^32   (final-sum lattice, int64 accumulator)
      (iii) 3*(nA+1)^5*8 B     (measured per-worker peak of the dense QA
            core, from the original analysis) + slack fits the per-worker
            address budget.
    Returns (ok, reason-string)."""
    uA, uB = _consa_split(u)
    nA, n = sum(uA.values()), sum(uB.values())
    d, L = nA + 1, n + 1
    if d * d > (1 << 16):
        return False, (f"GEMM inner dimension (nA+1)^2 = {d * d} exceeds "
                       f"2^16 (nA = {nA})")
    if L ** 3 >= (1 << 32):
        return False, (f"final-sum lattice (n+1)^3 = {L ** 3} >= 2^32 "
                       f"(n = {n})")
    need = 3 * d ** 5 * 8
    budget = _address_budget(procs)
    if need + CONSA_SLACK_BYTES > budget:
        return False, (f"QA core needs ~{need / 1e9:.2f} GB/worker "
                       f"(3*(nA+1)^5*8 B at nA = {nA}) + "
                       f"{CONSA_SLACK_BYTES / 1e9:.1f} GB slack, exceeding "
                       f"the per-worker address budget "
                       f"{budget / 1e9:.2f} GB")
    return True, f"nA={nA} n={n} QA-core {need / 1e6:.1f} MB/worker"


def _consa_Z_mod(u, p):
    """Z(12|34-frame; u) mod p by the tensor split (the original engine
    core, fast backend).  Preconditions enforced by consa_gate()."""
    uA, uB = _consa_split(u)
    nA, n = sum(uA.values()), sum(uB.values())
    V = _consa_V_array(uB.get('xxxx', 0), uB.get('xxyy', 0), p)  # (n+1)^3
    It = _consa_Itable(n, nA, p)                                 # (n+1,nA+1)
    gather = _consa_numba_backend()
    if gather:
        QA = _consa_QA_numba(uA, p, gather)
    else:
        QA = _consa_QA_numpy(uA, p)
    d, L = nA + 1, n + 1
    # K2(A; b1 b2) = Itil(A;b1) Itil(A;b2)              (L, d^2)
    K2 = (It[:, :, None] * It[:, None, :] % p).reshape(L, d * d)
    # J12(A; b3 b4 b5) = sum_{b1 b2} K2 @ QA            (L, d^3)
    J12 = _consa_mm_mod(K2, QA.reshape(d * d, d ** 3), p)
    # J34(B; A b5) = sum_{b3 b4} K2 @ J12^T-blocks      (L, L*d)
    T = np.ascontiguousarray(
        J12.reshape(L, d * d, d).transpose(1, 0, 2)).reshape(d * d, L * d)
    J34 = _consa_mm_mod(K2, T, p).reshape(L, L, d)               # [B, A, b5]
    # J5(B A; C) = sum_{b5} J34 @ Itil(C;b5)            (L*L, L)
    J5 = _consa_mm_mod(J34.reshape(L * L, d),
                       np.ascontiguousarray(It.T), p).reshape(L, L, L)
    # S = sum_{A,B,C} V(A,B,C) * J5(B,A,C)
    W = V.transpose(1, 0, 2) * J5 % p       # elementwise < 2^62, int64-safe
    assert L ** 3 < (1 << 32)               # gate precondition backstop
    S = int(W.sum() % p)
    inv4 = pow(4, p - 2, p)
    return S * pow(inv4, 4 * nA + 6 * n, p) % p


def _consa_worker(args):
    u, N, p = args
    return p, _consa_Z_mod(u, p)


# ===========================================================================
# Exact evidence: engine dispatch + disjoint-CRT reconstruction
# ===========================================================================
def _consa_gate_all_topologies(counts, procs):
    """consa gate on all three topology-permuted count vectors.
    Returns (ok, per-topology notes list, first-refusal string or None)."""
    notes, refusal = [], None
    for T in TOPOLOGIES:
        up = {k: v for k, v in
              permute_counts(counts, TOPO_TAU[T]).items() if v > 0}
        ok, why = consa_gate(up, procs)
        notes.append(f"{T}: {why}")
        if not ok and refusal is None:
            refusal = f"{T}: {why}"
    return refusal is None, notes, refusal


def choose_engine(N, forced, counts=None, procs=1):
    """Engine selection.  Returns (engine, note); the note is printed in
    the output block (the tool names its engine and why, honestly)."""
    direct_note = ''
    if forced == 'consa':
        if N > CONSA_NMAX:
            raise EngineCeilingError(
                f"N={N} exceeds the consa engine's wired ceiling "
                f"N={CONSA_NMAX} (original-analysis production regime); "
                f"refusing rather than running unvalidated territory")
        if counts is None:
            raise EngineCeilingError(
                "consa engine requires the count vector for its split "
                "gate (internal error: counts not passed)")
        ok, notes, refusal = _consa_gate_all_topologies(counts, procs)
        if not ok:
            raise EngineCeilingError(
                f"consa split gate REFUSES this dataset ({refusal}); "
                f"use --engine auto/dense/parseval for the direct path")
        return 'consa', 'forced; split gate PASS (' + '; '.join(notes) + ')'
    if forced != 'auto':
        eng = forced
        note = 'forced'
    else:
        if counts is not None and N <= CONSA_NMAX:
            ok, notes, refusal = _consa_gate_all_topologies(counts, procs)
            if ok:
                return 'consa', ('auto: Construction-A split gate PASS ('
                                 + '; '.join(notes) + ')')
            direct_note = (f"auto: consa REFUSED ({refusal}) -- "
                           f"falling back to the direct engine")
            log(f"# engine auto-select: {direct_note}")
        elif N > CONSA_NMAX:
            direct_note = (f"auto: N={N} exceeds the consa ceiling "
                           f"N={CONSA_NMAX} -- direct engine")
        eng = 'dense' if N <= DENSE_NMAX else 'parseval'
        note = direct_note if direct_note else 'auto: by N'
    if eng == 'dense' and N > DENSE_NMAX:
        raise EngineCeilingError(
            f"dense engine is validated for N <= {DENSE_NMAX} (got N={N}); "
            f"use the parseval engine")
    if N > PARSEVAL_NMAX:
        raise EngineCeilingError(
            f"N={N} exceeds the validated/practical ceiling "
            f"N={PARSEVAL_NMAX} of the streamed Parseval engine "
            f"(measured cost ~2.8e-7*(N+1)^5 CPU-s per prime, idle-box "
            f"fit, x ~{2*nprimes_half(lemma1_denominator(min(N,1000)))} primes per topology; "
            f"the N~890 full-alignment regime is ~10^17 ops and OUT OF "
            f"REACH for the direct engines"
            + (f"; additionally {direct_note}" if direct_note else "")
            + "; recurrence transport for general counts is an open "
            f"gap -- refusing rather than silently grinding)")
    return eng, note


def Z_exact(u, N, engine, procs, tag=""):
    """Exact Fraction Z(12|34-frame; u) with the disjoint-CRT gate.

    Returns (Fraction, gate_dict).  Two disjoint half-stacks of primes,
    each sized a priori from the Lemma-1 denominator bound, must
    reconstruct the identical rational."""
    import multiprocessing as mproc
    D = lemma1_denominator(N)
    nh = nprimes_half(D)
    if engine == 'parseval':
        L = N + 1
        primes = find_parseval_primes(L, 2 * nh)
        args = [(u, N, L, p, w, None) for p, w in primes]
        worker = _parseval_worker
    elif engine == 'consa':
        ok, why = consa_gate(u, procs)      # fail-closed at the point of use
        if not ok:
            raise EngineCeilingError(f"consa split gate REFUSES: {why}")
        _consa_numba_backend()              # decide/warm JIT before the fork
        primes = [(p, None) for p in find_plain_primes(2 * nh)]
        args = [(u, N, p) for p, _ in primes]
        worker = _consa_worker
    else:
        primes = [(p, None) for p in find_plain_primes(2 * nh)]
        args = [(u, N, p) for p, _ in primes]
        worker = _dense_worker
    log(f"[{tag}] engine={engine} N={N} primes=2x{nh} "
        f"(Lemma-1 bound: bits(D)={D.bit_length()}) procs={procs}")
    if procs > 1:
        with mproc.Pool(procs) as pool:
            res = pool.map(worker, args)
    else:
        res = [worker(a) for a in args]
    halves = []
    for half in (res[:nh], res[nh:]):
        rems = [Zp * (D % p) % p for p, Zp in half]
        mods = [p for p, _ in half]
        R, M = crt_list(rems, mods)
        if R > M // 2:
            R -= M
        halves.append(Fraction(R, D))
    gate = {
        'nprimes_per_half': nh,
        'half_A_range': [int(primes[0][0]), int(primes[nh - 1][0])],
        'half_B_range': [int(primes[nh][0]), int(primes[2 * nh - 1][0])],
        'match': halves[0] == halves[1],
    }
    ZA = halves[0]
    if not (0 < ZA < 1):
        gate['match'] = False
    return ZA, gate


# ===========================================================================
# Independent quadrature cross-check (--verify; float64, pruning code path)
# ===========================================================================
def _jc_vec(s_obs, x):
    same = (1.0 + 3.0 * x) / 4.0
    diff = (1.0 - x) / 4.0
    o = np.empty((4, len(x)))
    for a in range(4):
        o[a] = same if a == s_obs else diff
    return o


def _quad_slice(args):
    (i5, t5, nodes, logw, reps, uvec) = args
    K = len(nodes)
    same5 = (1.0 + 3.0 * t5) / 4.0
    diff5 = (1.0 - t5) / 4.0
    P5 = np.full((4, 4), diff5)
    np.fill_diagonal(P5, same5)
    V = [[_jc_vec(s, nodes) for s in rep] for rep in reps]
    lw2 = logw[:, None, None]
    lw34 = logw[None, :, None] + logw[None, None, :]
    o = np.empty(K)
    ncls = len(reps)
    Smat = np.empty((ncls, K, K, K))
    for i1 in range(K):
        for k, rep in enumerate(reps):
            w_a = V[k][0][:, i1]
            A = V[k][1]
            M = 0.25 * (w_a[:, None] * P5).T @ A
            pgrid = np.einsum('bi,bj,bk->ijk', M, V[k][2], V[k][3],
                              optimize=True)
            Smat[k] = np.log(pgrid)
        S = np.tensordot(uvec, Smat, axes=(0, 0))
        S += lw2
        S += lw34
        m = S.max()
        o[i1] = m + np.log(np.exp(S - m).sum())
    return logw[i5] + (o + logw)


def logZ_quadrature(u, K, procs):
    """log Z(12|34; u) via K^5 Gauss-Legendre grid, float64 log-sum-exp.
    Independent route: pattern probabilities from pruning transition
    matrices directly, never from the expanded polynomials."""
    import multiprocessing as mproc
    xs, ws = np.polynomial.legendre.leggauss(K)
    nodes = (xs + 1.0) / 2.0
    logw = np.log(ws / 2.0)
    labs = [l for l, c in sorted(u.items()) if c > 0]
    reps = [tuple('xyzw'.index(ch) for ch in l) for l in labs]
    uvec = np.array([float(u[l]) for l in labs])
    args = [(i5, nodes[i5], nodes, logw, reps, uvec) for i5 in range(K)]
    if procs > 1:
        with mproc.Pool(min(procs, K)) as pool:
            parts = pool.map(_quad_slice, args)
    else:
        parts = [_quad_slice(a) for a in args]
    allv = np.concatenate(parts)
    m = allv.max()
    return float(m + np.log(np.exp(allv - m).sum()))


# ===========================================================================
# Input handling
# ===========================================================================
def parse_fasta(text):
    seqs, name = {}, None
    order = []
    for line in text.splitlines():
        line = line.strip()
        if not line:
            continue
        if line.startswith('>'):
            name = line[1:].split()[0] if line[1:].split() else line[1:]
            if name in seqs:
                raise InputFormatError(f"duplicate FASTA record '{name}'")
            seqs[name] = []
            order.append(name)
        else:
            if name is None:
                raise InputFormatError("FASTA: sequence before first '>'")
            seqs[name].append(line)
    if not seqs:
        raise InputFormatError("no FASTA records found")
    return {n: ''.join(seqs[n]).upper() for n in order}, order


def parse_nexus(text):
    m = re.search(r'matrix(.*?);', text, re.S | re.I)
    if not m:
        raise InputFormatError("NEXUS: no MATRIX block found")
    body = m.group(1)
    seqs, order = {}, []
    for line in body.splitlines():
        ls = line.strip()
        if not ls or ls.startswith('['):
            continue
        parts = ls.split()
        if len(parts) >= 2:
            if parts[0] not in seqs:
                seqs[parts[0]] = ''
                order.append(parts[0])
            seqs[parts[0]] += ''.join(parts[1:]).upper()
    if not seqs:
        raise InputFormatError("NEXUS: empty MATRIX block")
    return seqs, order


ACGT = set('ACGT')


def collapse_alignment(seqs, order):
    if len(order) != 4:
        raise TaxaError(
            f"exactly 4 taxa required, found {len(order)}: "
            f"{', '.join(order) if order else '(none)'}")
    lens = {len(seqs[t]) for t in order}
    if len(lens) != 1:
        raise InputFormatError(
            "sequences are not aligned (unequal lengths: "
            + ', '.join(f"{t}:{len(seqs[t])}" for t in order) + ")")
    cols = list(zip(*[seqs[t].replace('U', 'T') for t in order]))
    counts = {lab: 0 for lab in ALL_LABELS}
    dropped = 0
    for col in cols:
        if all(c in ACGT for c in col):
            counts[canonical_label(col)] += 1
        else:
            dropped += 1
    return counts, len(cols), dropped


def parse_pattern_counts(spec):
    counts = {lab: 0 for lab in ALL_LABELS}
    for tok in spec.split(','):
        tok = tok.strip()
        if not tok:
            continue
        try:
            lab, v = tok.split(':')
            v = int(v)
        except ValueError:
            raise InputFormatError(f"bad --pattern-counts token '{tok}' "
                                   f"(expect label:count)")
        if lab not in counts:
            raise InputFormatError(
                f"unknown pattern label '{lab}'; the 15 canonical "
                f"first-appearance labels are: {', '.join(ALL_LABELS)}")
        if v < 0:
            raise InputFormatError(f"negative count for '{lab}'")
        counts[lab] += v
    return counts


def counts_canonical(counts):
    return ','.join(f"{lab}:{counts.get(lab, 0)}" for lab in ALL_LABELS)


# ===========================================================================
# Result block
# ===========================================================================
def digits_agree(a, b):
    a, b = mp.mpf(a), mp.mpf(b)
    if a == b:
        return float('inf')
    return float(-mp.log10(abs((a - b) / b)))


def run_point(counts, args, taxa=None):
    """Full evidence run on one count vector.  Returns (rc, results dict)."""
    N = sum(counts.values())
    if N == 0:
        raise InputFormatError("no usable columns (N = 0)")
    try:
        alpha = Fraction(args.alpha)
    except (ValueError, ZeroDivisionError):
        print(f"AlphaParseError: --alpha must be a rational number "
              f"(e.g. 1, 2/3), got {args.alpha!r}", file=sys.stderr)
        sys.exit(12)
    if alpha != 1:
        raise AlphaScopeError(
            f"alpha = {args.alpha} is outside the validated scope of the "
            f"validated engines (alpha = 1 only: uniform prior on "
            f"x_e = exp(-4/3 mu t_e), i.e. Exp(1) on the rescaled branch "
            f"length).  Rational alpha at small N is available in the "
            f"companion phylo-evaluate.py via --point.  Refusing -- no "
            f"new mathematics in this tool.")
    engine, engine_note = choose_engine(N, args.engine, counts, args.procs)
    engine_label = engine
    if engine == 'consa':
        engine_label += '[numba]' if _consa_numba_backend() else '[numpy]'
    dps = args.dps
    out(f"# phylo-evidence v{VERSION}")
    out(f"# model: JC69, exponential branch-length prior, alpha = 1")
    out(f"# N = {N} site patterns; engine = {engine_label} (mod-p CRT, "
        f"disjoint half-stacks)")
    out(f"# engine selection: {engine_note}")
    if taxa:
        out(f"# taxa (leaf order 1..4): {', '.join(taxa)}")
        out(f"# topology frame: 12|34 = (1,2)|(3,4) in the leaf order above")
    out(f"# counts: {counts_canonical(counts)}")
    out(f"# counts sha256: "
        f"{hashlib.sha256(counts_canonical(counts).encode()).hexdigest()}")
    out(f"# alpha = 1, dps = {dps}")
    results, gates = {}, {}
    rc = 0
    for T in TOPOLOGIES:
        up = permute_counts(counts, TOPO_TAU[T])
        Z, gate = Z_exact(up, N, engine, args.procs, tag=T)
        results[T] = Z
        gates[T] = gate
    out()
    with mp.workdps(dps + 15):
        for T in TOPOLOGIES:
            Z = results[T]
            g = gates[T]
            out(f"T = {T}")
            out(f"  Z = {Z.numerator}")
            out(f"      / {Z.denominator}")
            out(f"  digits: {len(str(Z.numerator))} / "
                f"{len(str(Z.denominator))}")
            lz = mp.log(mp.mpf(Z.numerator)) - mp.log(mp.mpf(Z.denominator))
            out(f"  log Z = {mp.nstr(lz, dps)}")
            out(f"  disjoint-CRT gate: half-A ({g['nprimes_per_half']} primes,"
                f" p {g['half_A_range'][0]}..{g['half_A_range'][1]}) vs "
                f"half-B ({g['nprimes_per_half']} primes, p "
                f"{g['half_B_range'][0]}..{g['half_B_range'][1]}): "
                f"{'MATCH' if g['match'] else 'MISMATCH'}")
            out()
        pairs = [('12|34', '13|24'), ('12|34', '14|23'), ('13|24', '14|23')]
        out("Exact Bayes factors:")
        for a, b in pairs:
            bf = results[a] / results[b]
            lbf = (mp.log(mp.mpf(bf.numerator))
                   - mp.log(mp.mpf(bf.denominator)))
            out(f"  BF({a} : {b}) = {bf.numerator}")
            out(f"                    / {bf.denominator}")
            out(f"  log BF({a} : {b}) = {mp.nstr(lbf, dps)}")
        out()
    if not all(g['match'] for g in gates.values()):
        out("DISJOINT-CRT GATE: MISMATCH -- exact reconstruction NOT "
            "certified (exit 3)")
        rc = RC_GATE
    else:
        out("DISJOINT-CRT GATE: all topologies MATCH -- exact rational "
            "certified by two disjoint prime stacks")
    if args.verify and rc == 0:
        out()
        out("Independent quadrature cross-check (float64 Gauss-Legendre, "
            "pruning route):")
        K1 = max(24, N // 2 + 2)
        K2 = K1 + 16
        worst = float('inf')
        with mp.workdps(40):
            for T in TOPOLOGIES:
                up = permute_counts(counts, TOPO_TAU[T])
                lq1 = logZ_quadrature(up, K1, args.procs)
                lq2 = logZ_quadrature(up, K2, args.procs)
                Z = results[T]
                lz = mp.log(mp.mpf(Z.numerator)) - mp.log(mp.mpf(Z.denominator))
                d = digits_agree(mp.mpf(lq2), lz)
                stab = digits_agree(mp.mpf(lq1), mp.mpf(lq2))
                worst = min(worst, d)
                out(f"  {T}: quad logZ = {lq2!r}  agree(exact) = {d:.1f}d  "
                    f"K-sweep({K1},{K2}) stable to {stab:.1f}d")
        if worst > 10.0:
            out(f"  VERIFY: PASS (all topologies agree > 10 digits; "
                f"float64-attainable)")
        else:
            out(f"  VERIFY: FAIL (worst agreement {worst:.1f}d <= 10d) "
                f"(exit 5)")
            rc = RC_VERIFY
    return rc, results


# ===========================================================================
# Selftest (vendored references; byte-compare full integer strings)
# ===========================================================================
def run_selftest(args):
    path = args.selftest_data or os.path.join(
        os.path.dirname(os.path.abspath(__file__)),
        'phylo-evidence-selftest.json')
    with open(path) as f:
        data = json.load(f)
    failures = []
    names = (['b4_n20'] if args.quick
             else ['trnaq_rf00005'] if args.anchor
             else list(data['datasets']))
    for name in names:
        ds = data['datasets'][name]
        counts = {lab: int(v) for lab, v in ds['counts'].items()}
        N = sum(counts.values())
        out(f"== selftest: {name} (N={N}) ==")
        # (a) pattern-collapse check when sequences are vendored
        if 'sequences' in ds:
            seqs = {k: v for k, v in ds['sequences'].items()}
            order = list(ds['sequences'].keys())
            c2, ncols, dropped = collapse_alignment(seqs, order)
            okc = (c2 == counts)
            out(f"  collapse check (vendored sequences -> counts, "
                f"{ncols} cols, {dropped} dropped): "
                f"{'PASS' if okc else 'FAIL'}")
            if not okc:
                failures.append(f"{name}: sequence collapse mismatch")
        # (b) engines: per-dataset list if vendored (new consa rows),
        # else the original default (parseval + dense when in range) --
        # the direct-engine legs are UNCHANGED from v1.0
        engines = ds.get('engines',
                         ['parseval'] + (['dense'] if N <= DENSE_NMAX
                                         else []))
        if args.anchor:
            engines = [e for e in engines if e == 'consa']
        for engine in engines:
            allok = True
            zs = {}
            for T in TOPOLOGIES:
                up = permute_counts(counts, TOPO_TAU[T])
                Z, gate = Z_exact(up, N, engine, args.procs,
                                  tag=f"{name}/{engine}/{T}")
                zs[T] = Z
                ref = ds['references'][T]
                ok = (str(Z.numerator) == ref['num']
                      and str(Z.denominator) == ref['den']
                      and gate['match'])
                out(f"  [{engine}] {T}: num {len(str(Z.numerator))}d / "
                    f"den {len(str(Z.denominator))}d  CRT-gate "
                    f"{'MATCH' if gate['match'] else 'MISMATCH'}  "
                    f"byte-compare vs vendored: {'PASS' if ok else 'FAIL'}")
                if not ok:
                    allok = False
                    failures.append(f"{name}/{engine}/{T}: integer-string "
                                    f"or gate mismatch")
            # (c) log Bayes factors vs vendored decimals
            if allok and 'log_bfs' in ds:
                with mp.workdps(60):
                    for key, refstr in sorted(ds['log_bfs'].items()):
                        a, b = key.split('_vs_')
                        bf = zs[a] / zs[b]
                        lbf = (mp.log(mp.mpf(bf.numerator))
                               - mp.log(mp.mpf(bf.denominator)))
                        ref = mp.mpf(refstr)
                        d = digits_agree(lbf, ref)
                        ok = d > (len(refstr.split('.')[-1]) - 2)
                        out(f"  [{engine}] log BF({a}:{b}) = "
                            f"{mp.nstr(lbf, 20)} vs vendored {refstr[:22]}..."
                            f" agree {d if d != float('inf') else 99:.1f}d: "
                            f"{'PASS' if ok else 'FAIL'}")
                        if not ok:
                            failures.append(f"{name}/{engine}/BF {key}")
    out()
    if failures:
        out(f"SELFTEST: FAIL ({len(failures)} mismatches) (exit 4)")
        for f_ in failures:
            out(f"  - {f_}")
        return RC_SELFTEST
    if args.quick:
        out("SELFTEST: PASS (--quick: 20-site reference only)")
    elif args.anchor:
        out("SELFTEST: PASS (--anchor: tRNA-Gln anchor quartet, "
            "consa engine leg only)")
    else:
        out(f"SELFTEST: PASS -- stored references ({', '.join(names)}) "
            f"reproduced byte-identically")
    return 0


# ===========================================================================
def main():
    ap = argparse.ArgumentParser(
        prog='phylo-evidence.py',
        formatter_class=argparse.RawDescriptionHelpFormatter,
        description=__doc__)
    ap.add_argument('alignment', nargs='?', help='FASTA (or NEXUS with '
                    '--nexus) alignment of exactly four taxa')
    ap.add_argument('--nexus', action='store_true',
                    help='parse the input file as NEXUS')
    ap.add_argument('--pattern-counts', metavar='SPEC',
                    help="direct input: 'xxxx:35,xxyy:8,...' on the 15 "
                    "canonical labels")
    ap.add_argument('--alpha', default='1',
                    help='prior shape (validated: 1 only; anything else '
                    'is refused by name)')
    ap.add_argument('--dps', type=int, default=50,
                    help='decimal digits for log Z / log BF output '
                    '(exact rationals are dps-independent)')
    ap.add_argument('--engine',
                    choices=['auto', 'dense', 'parseval', 'consa'],
                    default='auto',
                    help='engine override (default auto: consa when its '
                    'split gate passes, else direct by N)')
    ap.add_argument('--procs', type=int, default=min(os.cpu_count() or 1, 32),
                    help='worker processes for the prime stack')
    ap.add_argument('--verify', action='store_true',
                    help='add the independent float64 quadrature '
                    'cross-check (float-attainable digits)')
    ap.add_argument('--selftest', action='store_true',
                    help='reproduce the vendored reference values and '
                    'byte-compare full integer strings (the direct-engine '
                    'tRNA-Gln leg costs ~84 CPU-hours -- measured, see the '
                    'ENGINES note above -- spread over --procs workers; '
                    'the consa legs (20-site, tRNA-Lys N=70, tRNA-Gln '
                    'anchor N=68) cost ~1.5 CPU-h combined; on a shared/'
                    'contended box the full run can be many hours of wall '
                    'time; use --quick for the fast 20-site leg only)')
    ap.add_argument('--quick', action='store_true',
                    help='with --selftest: 20-site reference only')
    ap.add_argument('--anchor', action='store_true',
                    help='with --selftest: the Rfam RF00005 mt tRNA-Gln '
                    'anchor quartet (mouse/rat/chicken/frog, N=68) only, '
                    'consa engine leg only -- reproduces the 188/302-digit '
                    'fraction and the 14.5368/15.5563 log BFs without the '
                    '~84 CPU-h direct-engine leg')
    ap.add_argument('--selftest-data', metavar='PATH', default=None,
                    help='override the vendored selftest JSON path')
    ap.add_argument('--keep-policy', default='drop-nonACGT',
                    help='column policy (reserved; only drop-nonACGT '
                    'is implemented)')
    args = ap.parse_args()
    try:
        if args.keep_policy != 'drop-nonACGT':
            raise InputFormatError(
                f"--keep-policy '{args.keep_policy}' is reserved for "
                f"future use; only 'drop-nonACGT' is implemented")
        if args.selftest:
            sys.exit(run_selftest(args))
        if args.pattern_counts and args.alignment:
            raise InputFormatError(
                "give either an alignment file or --pattern-counts, not both")
        if args.pattern_counts:
            counts = parse_pattern_counts(args.pattern_counts)
            rc, _ = run_point(counts, args)
            sys.exit(rc)
        if not args.alignment:
            ap.print_usage(sys.stderr)
            raise InputFormatError("no input (alignment file, "
                                   "--pattern-counts, or --selftest)")
        with open(args.alignment, 'rb') as f:
            raw = f.read()
        log(f"# input file: {os.path.basename(args.alignment)} "
            f"({len(raw)} bytes) sha256 "
            f"{hashlib.sha256(raw).hexdigest()}")
        text = raw.decode('utf-8', errors='replace')
        seqs, order = (parse_nexus if args.nexus else parse_fasta)(text)
        counts, ncols, dropped = collapse_alignment(seqs, order)
        out(f"# alignment: {len(order)} taxa x {ncols} columns; "
            f"{dropped} non-ACGT columns dropped (policy: drop-nonACGT), "
            f"{ncols - dropped} used")
        rc, _ = run_point(counts, args, taxa=order)
        sys.exit(rc)
    except CLIError as e:
        log(f"ERROR[{type(e).name}]: {e}")
        sys.exit(type(e).rc)


if __name__ == '__main__':
    main()
