#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
xoverlap_memory_test.py — do the x1..x20 overlap columns have any memory?
=========================================================================
The x-columns record, for each draw, how many numbers it shares with the draw N
places before it: xN(t) = |draw_t ∩ draw_{t-N}|, for N = 1..20. This script asks
whether xN = 0 ("the next draw dodged those 5 numbers") is *predictable* — via the
column's own history, gap/recurrence analysis, Markov structure, or the "diagonal"
(per-number recurrence of a fixed earlier draw).

It runs the same battery for Joker (5/45) and Eurojackpot (5/50):

  1. FOUNDATION — per-number recurrence hazard P(appear | delay since last seen).
  2. Per-column base rate P(xN=0) with Wilson CIs + a chi-square uniformity test.
  3. Gap analysis — gaps between 0-events vs a Geometric fit (memoryless ⇔ geometric).
  4. Markov — order-1 and order-2 transition tables for (xN=0).
  5. Walk-forward — pick 6 columns each draw (gap / base-rate / per-number "diagonal"
     predictors) and measure the out-of-sample 0-rate vs the base rate (+ Wilson LB).
  6. FAIR-NULL — shuffle the draw order many times and recompute the walk-forward lift,
     to see whether the real result is anything a coin-flip sequence wouldn't produce.

Needs each game's `hist_df.csv` (with st1..st5 and x1..x20). Writes 4 PNGs to
static/img/blog/x-overlap-memory/ and prints a full numeric summary.
"""
from __future__ import annotations
from math import comb
from pathlib import Path

import numpy as np
import pandas as pd
from scipy import stats
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

GAMES = {"Joker (5/45)": ("games/joker/data/hist_df.csv", 45),
         "Eurojackpot (5/50)": ("games/eurojackpot/data/hist_df.csv", 50)}
ST = ["st1", "st2", "st3", "st4", "st5"]
XIS = [f"x{i}" for i in range(1, 21)]
IMG = Path("static/img/blog/x-overlap-memory"); IMG.mkdir(parents=True, exist_ok=True)
K, WARM, B_NULL = 6, 250, 800
COL = {"Joker (5/45)": "#e0a417", "Eurojackpot (5/50)": "#3ba7d6"}
plt.rcParams.update({"figure.dpi": 130, "font.size": 11, "axes.grid": True, "grid.alpha": .25})


def wilson(k, n, z=1.96):
    p = k / n; d = 1 + z * z / n
    c = (p + z * z / (2 * n)) / d
    hw = z * np.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d
    return c - hw, c + hw


def load(path):
    h = pd.read_csv(path); h.columns = [str(c).strip() for c in h.columns]
    draws = h[ST].astype(int).to_numpy()
    X0 = (h[XIS].astype(float).to_numpy() == 0).astype(int)   # 1 where xN == 0
    return draws, X0


def per_number_hazard(draws, npool, maxd=15):
    last = {m: None for m in range(1, npool + 1)}; hz = {}
    for t in range(len(draws)):
        s = set(int(x) for x in draws[t])
        for m in range(1, npool + 1):
            if last[m] is not None:
                hz.setdefault(t - last[m], []).append(1 if m in s else 0)
        for m in draws[t]:
            last[int(m)] = t
    delays = np.array(sum(([d] * len(v) for d, v in hz.items()), []))
    appeared = np.array(sum((list(v) for v in hz.values()), []))
    corr = float(np.corrcoef(delays, appeared)[0, 1])
    curve = {d: (float(np.mean(hz[d])), len(hz[d])) for d in range(1, maxd + 1) if d in hz}
    return curve, corr


def gap_predictor(X0, warm=WARM, k=K):
    """Pick the k columns with the longest current gap since their last 0; report the
    out-of-sample 0-rate of those picks. Pure x-column predictor (used for the fair-null)."""
    n = X0.shape[0]; last0 = np.full(20, -1); hits = []
    for t in range(n - 1):
        gap = t - last0
        if t >= warm:
            pick = np.argsort(-gap)[:k]
            hits.append(X0[t + 1, pick].mean())
        last0[X0[t] == 1] = t
    return np.array(hits)


def diagonal_predictor(draws, X0, npool, warm=WARM, k=K):
    """The 'diagonal': score each lag N by P(next draw dodges reference-draw_{t+1-N})
    = prod(1 - per-number appearance rate). Ranks lags, picks k, reports 0-rate."""
    n = len(draws); pnum = 5 / npool
    last = np.full(npool + 1, -10**9); hits = []
    for t in range(n - 1):
        for m in draws[t]:
            last[int(m)] = t
        if t >= warm:
            score = np.zeros(20)
            for j, N in enumerate(range(1, 21)):
                ri = (t + 1) - N
                if ri < 0:
                    score[j] = -1; continue
                p0 = 1.0
                for m in draws[ri]:
                    p0 *= (1 - pnum)   # flat hazard ⇒ per-number rate is the base rate
                score[j] = p0
            pick = np.argsort(-score)[:k]
            hits.append(X0[t + 1, pick].mean())
    return np.array(hits)


def fair_null_lift(draws, npool, base, b=B_NULL, warm=WARM, k=K, seed=7):
    """Shuffle the draw ORDER b times; each time recompute the x-columns from the overlap
    matrix and rerun the gap predictor. Returns the null distribution of its lift."""
    rng = np.random.default_rng(seed)
    n = len(draws)
    A = np.zeros((n, npool), dtype=np.int8)
    for t in range(n):
        A[t, draws[t] - 1] = 1
    O = A @ A.T                     # overlap matrix O[i,j] = |draw_i ∩ draw_j|
    lifts = []
    for _ in range(b):
        perm = rng.permutation(n)
        Xs = np.zeros((n, 20), dtype=np.int8)
        for N in range(1, 21):
            Xs[N:, N - 1] = (O[perm[N:], perm[:-N]] == 0).astype(np.int8)
        h = gap_predictor(Xs, warm, k)
        lifts.append(h.mean() / base)
    return np.array(lifts)


# --------------------------------------------------------------------------- run
summary = {}
haz_fig, haz_ax = plt.subplots(figsize=(7.2, 4.2))
p0_fig, p0_ax = plt.subplots(1, 2, figsize=(11, 4.2))
gap_fig, gap_ax = plt.subplots(1, 2, figsize=(11, 4.2))
null_fig, null_ax = plt.subplots(1, 2, figsize=(11, 4.2))

for gi, (game, (path, npool)) in enumerate(GAMES.items()):
    draws, X0 = load(path)
    n = len(draws); base = comb(npool - 5, 5) / comb(npool, 5); pnum = 5 / npool
    c = COL[game]
    S = {"n": n, "base": base, "npool": npool}

    # 1. foundation hazard
    curve, corr = per_number_hazard(draws, npool)
    S["hazard_corr"] = corr
    ds = sorted(curve); haz_ax.plot(ds, [curve[d][0] for d in ds], "o-", color=c, label=f"{game}  (corr={corr:+.3f})")
    haz_ax.axhline(pnum, color=c, ls="--", alpha=.5)

    # 2. per-column base rate + CI + uniformity chi-square
    kc = X0.sum(axis=0); p0 = kc / n
    ci = np.array([wilson(int(x), n) for x in kc])
    pbar = p0.mean()
    chi = float(((kc - n * pbar) ** 2 / (n * pbar * (1 - pbar))).sum())
    chi_p = float(stats.chi2.sf(chi, 19))
    S.update(p0_min=float(p0.min()), p0_max=float(p0.max()), chi=chi, chi_p=chi_p)
    ax = p0_ax[gi]
    ax.errorbar(range(1, 21), p0, yerr=[p0 - ci[:, 0], ci[:, 1] - p0], fmt="o", color=c, capsize=2, ms=4)
    ax.axhline(base, color="k", ls="--", lw=1, label=f"base = {base:.3f}")
    ax.set_title(f"{game}\nP(xN=0) per column  (uniformity χ² p={chi_p:.2f})", fontsize=10)
    ax.set_xlabel("column N (lag)"); ax.set_ylabel("P(xN = 0)"); ax.legend(fontsize=9)

    # 3. gap analysis + geometric fit
    gaps = np.concatenate([np.diff(np.where(X0[:, cc] == 1)[0]) for cc in range(20)])
    pgeo = 1.0 / gaps.mean()
    obs = np.array([(gaps == g).sum() for g in range(1, 6)] + [(gaps >= 6).sum()], float)
    egp = np.array([(1 - pgeo) ** (g - 1) * pgeo for g in range(1, 6)] + [(1 - pgeo) ** 5], float) * len(gaps)
    gchi = float(((obs - egp) ** 2 / egp).sum()); gchi_p = float(stats.chi2.sf(gchi, len(obs) - 2))
    S.update(gap_mean=float(gaps.mean()), pgeo=float(pgeo), geo_chi=gchi, geo_p=gchi_p)
    ax = gap_ax[gi]
    xs = np.arange(1, 7); w = .38
    ax.bar(xs - w / 2, obs / len(gaps), w, color=c, label="observed", alpha=.85)
    ax.bar(xs + w / 2, egp / len(gaps), w, color="#888", label="Geometric fit", alpha=.7)
    ax.set_xticks(xs); ax.set_xticklabels(["1", "2", "3", "4", "5", "6+"])
    ax.set_title(f"{game}\ngaps between 0-events vs Geometric  (χ² p={gchi_p:.2f})", fontsize=10)
    ax.set_xlabel("gap (draws)"); ax.set_ylabel("fraction"); ax.legend(fontsize=9)

    # 4. Markov order-1 / order-2 (pooled over columns)
    b = X0
    prev = b[:-1].ravel(); cur = b[1:].ravel()
    p_00, p_0n = float(cur[prev == 1].mean()), float(cur[prev == 0].mean())
    S.update(mk_p00=p_00, mk_p0n=p_0n)
    o2 = {}
    for cc in range(20):
        s = b[:, cc]
        for t in range(2, n):
            o2.setdefault((int(s[t - 1]), int(s[t - 2])), []).append(int(s[t]))
    S["mk2"] = {kk: float(np.mean(v)) for kk, v in sorted(o2.items())}

    # 5. walk-forward predictors
    def wf(h):
        m = h.mean(); lo, _ = wilson(int(round(h.sum())), len(h))
        return float(m), float(m / base), float(lo / base)
    hb = np.array([X0[t + 1, :K].mean() for t in range(WARM, n - 1)])   # fixed cols 1..6
    S["wf"] = {"fixed_1-6": wf(hb),
               "gap_longest": wf(gap_predictor(X0)),
               "diagonal": wf(diagonal_predictor(draws, X0, npool))}

    # 6. fair-null
    nulls = fair_null_lift(draws, npool, base)
    real = gap_predictor(X0).mean() / base
    pval = float((nulls >= real).mean())
    S.update(null_mean=float(nulls.mean()), null_lo=float(np.percentile(nulls, 2.5)),
             null_hi=float(np.percentile(nulls, 97.5)), real_lift=float(real), null_p=pval)
    ax = null_ax[gi]
    ax.hist(nulls, bins=30, color="#999", alpha=.75, label="shuffled-order null")
    ax.axvline(real, color=c, lw=2.5, label=f"real = {real:.3f}")
    ax.axvline(1.0, color="k", ls="--", lw=1)
    ax.set_title(f"{game}\nwalk-forward lift vs null  (p={pval:.2f})", fontsize=10)
    ax.set_xlabel("lift (hit-rate / base)"); ax.set_ylabel("null trials"); ax.legend(fontsize=9)

    summary[game] = S

haz_ax.set_title("Foundation: per-number recurrence hazard P(appear | delay)")
haz_ax.set_xlabel("delay since the number was last drawn"); haz_ax.set_ylabel("P(number appears)")
haz_ax.legend(fontsize=9)
for fig, name in [(haz_fig, "img1"), (p0_fig, "img2"), (gap_fig, "img3"), (null_fig, "img4")]:
    fig.tight_layout(); fig.savefig(IMG / f"{name}.png", bbox_inches="tight"); plt.close(fig)

# --------------------------------------------------------------------------- print
print(f"\ncharts -> {IMG}\n" + "=" * 70)
for game, S in summary.items():
    print(f"\n### {game} | {S['n']} draws | base P(xN=0)={S['base']:.4f} "
          f"| per-number base={5/S['npool']:.4f}")
    print(f"  foundation:  corr(delay, appears) = {S['hazard_corr']:+.4f}")
    print(f"  per-column:  P(0) range [{S['p0_min']:.3f}, {S['p0_max']:.3f}]  "
          f"uniformity chi2={S['chi']:.1f} (df 19)  p={S['chi_p']:.3f}")
    print(f"  gaps:        mean={S['gap_mean']:.2f}  geometric p_hat={S['pgeo']:.3f}  "
          f"GOF chi2 p={S['geo_p']:.3f}")
    print(f"  markov o1:   P(0|prev=0)={S['mk_p00']:.4f}  P(0|prev>0)={S['mk_p0n']:.4f}")
    print(f"  markov o2:   " + "  ".join(f"{k}:{v:.3f}" for k, v in S["mk2"].items()))
    for name, (hit, lift, lb) in S["wf"].items():
        print(f"  wf {name:<11} hit={hit:.4f}  lift={lift:.3f}  WilsonLB(lift)={lb:.3f}")
    print(f"  FAIR-NULL:   real lift={S['real_lift']:.3f}  null mean={S['null_mean']:.3f}  "
          f"95% null [{S['null_lo']:.3f}, {S['null_hi']:.3f}]  p={S['null_p']:.3f}")
