"""Checks for lava_leap.py. Run: .venv/bin/python test_lava_leap.py  (or pytest)."""
import numpy as np

import lava_leap as ll


def channels(lava):
    """List of (start, width) lava channels on a floor."""
    out, t = [], 0
    while t < ll.FLOOR:
        if lava[t]:
            w = 1
            while t + w < ll.FLOOR and lava[t + w]:
                w += 1
            out.append((t, w))
            t += w
        else:
            t += 1
    return out


def test_floor_rules():
    rng = np.random.default_rng(0)
    counts = set()
    for _ in range(5000):
        lava = ll.new_floor(rng)
        ch = channels(lava)
        counts.add(len(ch))
        assert 3 <= len(ch) <= 7
        assert all(w in (1, 2) for _, w in ch)
        assert ch[0][0] >= 5 and ch[-1][0] + ch[-1][1] - 1 <= 36
        for (s1, w1), (s2, _) in zip(ch, ch[1:]):
            assert s2 - (s1 + w1) >= 2, "need >= 2 safe tiles between channels"
    assert counts == {3, 4, 5, 6, 7}


def test_sensors_and_play():
    lava = np.zeros(ll.FLOOR, dtype=bool)
    lava[[10, 20, 21]] = True
    assert list(ll.sensors(lava, 9)) == [1, 0, 0, 1]
    assert list(ll.sensors(lava, 18)) == [0, 1, 1, 1]
    assert list(ll.sensors(lava, 38)) == [0, 0, 0, 1]   # past the exit reads safe
    rng = np.random.default_rng(1)
    for _ in range(200):
        run = ll.play(np.zeros((4, 2)), lava, rng)
        assert len(run.states) == len(run.actions) == len(run.path) - 1
        assert 0 < run.reward <= 1
        moves = np.diff(run.path)
        for a, m, end in zip(run.actions, moves, run.path[1:]):
            assert m == 1 if a == ll.STEP else (m in (2, 3) or end == ll.EXIT)
        assert run.reward == 1.0 or lava[run.path[-1]]


def test_exact_matches_simulation():
    rng = np.random.default_rng(2)
    floors = [ll.new_floor(rng) for _ in range(60)]
    W = rng.normal(size=(4, 2))
    table = ll.leap_table(W)
    v, p = ll.exact_outcomes(table, floors)
    sim_v, sim_p = ll.simulate(table, floors, plays=400, rng=np.random.default_rng(3))
    assert abs(v.mean() - sim_v) < 0.01 and abs(p.mean() - sim_p) < 0.01


def test_reinforce_is_the_log_prob_gradient():
    rng = np.random.default_rng(4)
    lava = ll.new_floor(rng)
    W = rng.normal(scale=0.5, size=(4, 2))
    runs = [ll.play(W, lava, rng) for _ in range(8)]
    S, acts, R = ll.batch(runs, [r.reward for r in runs])

    def J(Wx):  # (1/N) sum_t r_t log pi(a_t|s_t)
        return np.mean(R * np.log(ll.policy(Wx, S)[np.arange(len(acts)), acts]))

    lr = 0.5
    analytic = (ll.reinforce_update(W, runs, lr) - W) / lr
    numeric = np.zeros_like(W)
    for i in np.ndindex(W.shape):
        d = np.zeros_like(W)
        d[i] = 1e-6
        numeric[i] = (J(W + d) - J(W - d)) / 2e-6
    assert np.allclose(analytic, numeric, atol=1e-6)


def test_clipped_gradient_matches_objective():
    rng = np.random.default_rng(5)
    lava = ll.new_floor(rng)
    W0 = np.zeros((4, 2))
    runs = [ll.play(W0, lava, rng) for _ in range(8)]
    r = np.array([run.reward for run in runs])
    S, acts, A = ll.batch(runs, (r - r.mean()) / (r.std() + 1e-8))
    rows = np.arange(len(acts))
    pi_old = ll.policy(W0, S)[rows, acts]
    trace = []
    ll.clipped_update(W0, S, acts, A, lr=3.0, passes=4, trace=trace)  # big lr: force clipping
    last = trace[-1]
    assert last["clipped"].any(), "test needs some clipped steps"

    def J(Wx):  # (1/N) sum_t min(rho A, clip(rho) A)
        rho = ll.policy(Wx, S)[rows, acts] / pi_old
        return np.mean(np.minimum(rho * A, np.clip(rho, 0.8, 1.2) * A))

    numeric = np.zeros_like(W0)
    for i in np.ndindex(W0.shape):
        d = np.zeros_like(W0)
        d[i] = 1e-7
        numeric[i] = (J(last["W"] + d) - J(last["W"] - d)) / 2e-7
    assert np.allclose(last["grad"], numeric, atol=1e-6)
    assert np.allclose(trace[0]["rho"], 1.0)            # pass 1: every ratio is exactly 1


def test_pass_one_is_reinforce_with_advantages():
    rng = np.random.default_rng(6)
    lava = ll.new_floor(rng)
    W = rng.normal(scale=0.3, size=(4, 2))
    runs = [ll.play(W, lava, rng) for _ in range(8)]
    r = np.array([run.reward for run in runs])
    adv = (r - r.mean()) / (r.std() + 1e-8)
    assert abs(adv.mean()) < 1e-9 and abs(adv.std() - 1) < 1e-6
    fake = [run._replace(reward=a) for run, a in zip(runs, adv)]
    one_pass = ll.grpo_update(W, runs, lr=0.5, passes=1)
    assert np.allclose(one_pass, ll.reinforce_update(W, fake, lr=0.5))


def test_optimal_rule_is_in_the_top_tied_set():
    rng = np.random.default_rng(7)
    floors = [ll.new_floor(rng) for _ in range(100)]
    scores, _ = ll.all_deterministic_policies(floors)
    leap_iff_lava_next = sum(1 << b for b in range(8) if b >> 2 & 1)  # patterns 1xx
    assert np.isclose(scores[leap_iff_lava_next], scores.max())


if __name__ == "__main__":
    tests = [f for name, f in sorted(globals().items()) if name.startswith("test_")]
    for f in tests:
        f()
        print("ok ", f.__name__)
    print(f"{len(tests)} tests passed")
