"""Independent correctness checks for genealogical ancestry propagation.

These tests validate the ancestry bookkeeping used in the Rohde benchmark
replication. They are deliberately small enough that the same pedigree can be
checked by a slow, explicit graph-traversal implementation.
"""
from __future__ import annotations
import random


def bitset_descendants(parent_layers, n):
    """Return descendant bitsets for every generation in a fixed pedigree.

    parent_layers[g][child] = (p1,p2) for transition from generation g to g+1
    going backward in time. Generation 0 is the present.
    """
    current = [1 << i for i in range(n)]
    out = [current]
    for parents in parent_layers:
        prev = [0] * n
        for child, (p1,p2) in enumerate(parents):
            prev[p1] |= current[child]
            prev[p2] |= current[child]
        current = prev
        out.append(current)
    return out


def traversal_descendants(parent_layers, n):
    """Slow independent check by traversing every present person's parent paths."""
    # desc[g][a] is a Python set of present-day descendants of ancestor a at g.
    desc = [[set() for _ in range(n)] for _ in range(len(parent_layers)+1)]
    for present in range(n):
        reachable = {present}
        desc[0][present].add(present)
        for g, parents in enumerate(parent_layers, start=1):
            new = set()
            for child in reachable:
                p1,p2 = parents[child]
                new.add(p1); new.add(p2)
            reachable = new
            for anc in reachable:
                desc[g][anc].add(present)
    return desc


def test_bitsets_against_traversal(trials=250):
    for trial in range(trials):
        rng = random.Random(910000 + trial)
        n = rng.randint(2, 12)
        generations = rng.randint(1, 10)
        layers=[]
        for _ in range(generations):
            layers.append([(rng.randrange(n),rng.randrange(n)) for _ in range(n)])
        bits = bitset_descendants(layers,n)
        sets = traversal_descendants(layers,n)
        for g in range(generations+1):
            for a in range(n):
                expected = sum(1 << x for x in sets[g][a])
                assert bits[g][a] == expected, (trial,n,g,a,bits[g][a],expected)
    return trials


def test_disconnected_components(n_per_component=40, generations=80, trials=25):
    """With zero migration between two components, no global CA can appear."""
    n=2*n_per_component
    everyone=(1<<n)-1
    for trial in range(trials):
        rng=random.Random(920000+trial)
        current=[1<<i for i in range(n)]
        for _ in range(generations):
            prev=[0]*n
            for child in range(n):
                lo=0 if child<n_per_component else n_per_component
                hi=n_per_component if child<n_per_component else n
                p1=rng.randrange(lo,hi); p2=rng.randrange(lo,hi)
                prev[p1] |= current[child]; prev[p2] |= current[child]
            assert not any(x==everyone for x in prev)
            current=prev
    return trials


def one_run_panmixia(n, seed, max_generations=80):
    rng=random.Random(seed)
    current=[1<<i for i in range(n)]
    everyone=(1<<n)-1
    T=None
    for g in range(1,max_generations+1):
        prev=[0]*n
        for child,bits in enumerate(current):
            p1=rng.randrange(n); p2=rng.randrange(n)
            prev[p1] |= bits; prev[p2] |= bits
        if T is None and any(x==everyone for x in prev):
            T=g
        if T is not None and all(x==0 or x==everyone for x in prev):
            return T,g
        current=prev
    raise RuntimeError('max_generations too small')


def test_seed_reproducibility():
    a=one_run_panmixia(250,123456)
    b=one_run_panmixia(250,123456)
    assert a==b
    return a


def test_mrca_persistence(n=100, trials=25):
    """Once a CA exists, at least one CA must exist in each earlier generation."""
    for t in range(trials):
        rng=random.Random(930000+t)
        current=[1<<i for i in range(n)]
        everyone=(1<<n)-1
        seen=False
        for _ in range(60):
            prev=[0]*n
            for child,bits in enumerate(current):
                p1=rng.randrange(n); p2=rng.randrange(n)
                prev[p1] |= bits; prev[p2] |= bits
            has=any(x==everyone for x in prev)
            if seen:
                assert has
            seen = seen or has
            current=prev
    return trials


if __name__=='__main__':
    print('bitset_vs_explicit_trials=', test_bitsets_against_traversal())
    print('disconnected_zero_migration_trials=', test_disconnected_components())
    print('seed_reproducibility_result=', test_seed_reproducibility())
    print('mrca_persistence_trials=', test_mrca_persistence())
    print('ALL TESTS PASSED')
