"""Deterministic analysis of retained dice outcomes; does not draw or filter samples."""
import argparse
import csv
import hashlib
import json
import platform
import re
from pathlib import Path
from xml.sax.saxutils import escape

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


def sha256(path):
    return hashlib.sha256(path.read_bytes()).hexdigest()


def holm(pvalues):
    """Holm step-down adjusted p-values, retaining monotonicity and original order."""
    order = np.argsort(pvalues, kind="stable")
    adjusted = np.empty(len(order))
    running = 0.0
    for rank, index in enumerate(order):
        running = max(running, (len(order) - rank) * pvalues[index])
        adjusted[index] = min(1.0, running)
    return adjusted


def clopper_pearson(count, n, confidence=0.95):
    tail = (1 - confidence) / 2
    lower = 0.0 if count == 0 else stats.beta.ppf(tail, count, n - count + 1)
    upper = 1.0 if count == n else stats.beta.ppf(1 - tail, count + 1, n - count)
    return [float(lower), float(upper)]


def analyze(entry, data, protocol, total_faces):
    path = data / entry["outcomesFile"]
    assert sha256(path) == entry["outcomesSha256"], entry["id"]
    assert sha256(data / entry["seedsFile"]) == entry["seedsSha256"]
    values = np.fromfile(path, dtype=np.uint8)
    n, sides = len(values), entry["sides"]
    assert n == entry["outcomes"] and np.all((values >= 1) & (values <= sides))
    counts = np.bincount(values, minlength=sides + 1)[1:]
    chi = stats.chisquare(counts, np.full(sides, n / sides))
    buckets = min(sides, 20)
    categorized = (values.astype(np.int64) - 1) * buckets // sides
    pair_count = n // 2
    left, right = categorized[:2 * pair_count:2], categorized[1:2 * pair_count:2]
    contingency = np.bincount(left * buckets + right, minlength=buckets ** 2).reshape(buckets, buckets)
    serial = stats.chi2_contingency(contingency, correction=False)
    assert np.min(serial.expected_freq) >= 5, "Expected cell counts too small"
    centered = values.astype(np.float64) - np.mean(values)
    denominator = np.dot(centered, centered)
    correlations = [float(np.dot(centered[:-lag], centered[lag:]) / denominator)
                    for lag in range(1, protocol["autocorrelationLags"] + 1)]
    tail = protocol["alpha"] / (2 * total_faces)
    envelope = [int(stats.binom.ppf(tail, n, 1 / sides)), int(stats.binom.isf(tail, n, 1 / sides))]
    record = {"id": entry["id"], "environment": entry["environment"], "mode": entry["mode"],
              "algorithm": entry["algorithm"], "sides": sides, "n": n, "counts": counts.tolist(),
              "expectedCount": n / sides, "chiSquare": float(chi.statistic),
              "degreesOfFreedom": sides - 1, "pUniformity": float(chi.pvalue),
              "serialChiSquare": float(serial.statistic), "serialDegreesOfFreedom": int(serial.dof),
              "pSerial": float(serial.pvalue), "serialBuckets": buckets, "serialPairs": pair_count,
              "minimumExpectedPairCount": float(np.min(serial.expected_freq)),
              "mean": float(np.mean(values)), "theoreticalMean": (sides + 1) / 2,
              "sampleVariance": float(np.var(values, ddof=1)), "theoreticalVariance": (sides ** 2 - 1) / 12,
              "totalVariationDistance": float(np.sum(np.abs(counts / n - 1 / sides)) / 2),
              "maxAbsoluteDeviation": float(np.max(np.abs(counts / n - 1 / sides))),
              "faceNullEnvelope": envelope, "faceEnvelopeOutside": int(np.sum((counts < envelope[0]) | (counts > envelope[1]))),
              "individualFaceCI95": [clopper_pearson(int(c), n) for c in counts],
              "autocorrelation": correlations}
    return record, values


def figures(records, values_by_id, protocol, out, language="pt"):
    plt.rcParams.update({"figure.facecolor": "#141320", "axes.facecolor": "#1d1b2b",
                         "savefig.facecolor": "#141320", "text.color": "#f0eff5",
                         "axes.labelcolor": "#dedce7", "xtick.color": "#dedce7",
                         "ytick.color": "#dedce7", "axes.edgecolor": "#6d6a7d",
                         "grid.color": "#6d6a7d", "grid.alpha": 0.25, "font.size": 11,
                         "font.family": "DejaVu Sans", "svg.hashsalt": protocol["studyId"]})
    out.mkdir(parents=True, exist_ok=True)
    def save(fig, name):
        if language == "en":
            for label in fig.findobj(matplotlib.text.Text):
                label.set_text(english_label(label.get_text()))
        fig.savefig(out / f"{name}.svg", bbox_inches="tight", metadata={"Date": None})
        svg = out / f"{name}.svg"
        title = fig._suptitle.get_text() if fig._suptitle is not None else fig.axes[0].get_title()
        svg.write_text(re.sub(r"(<svg\b[^>]*>)", lambda match: match[0] + "\n <title>" + escape(title) + "</title>",
                              svg.read_text(encoding="utf-8"), count=1), encoding="utf-8")
        fig.savefig(out / f"{name}.png", dpi=170, bbox_inches="tight", metadata={"Software": "Matplotlib"})
        plt.close(fig)
    def find(mode, sides, algorithm="xoshiro128ss", environment="Node.js"):
        return next(r for r in records if (r["mode"], r["sides"], r["algorithm"], r["environment"])
                    == (mode, sides, algorithm, environment))

    fig, axes = plt.subplots(2, 1, figsize=(11.6, 7.8), constrained_layout=True)
    for ax, record, label in zip(axes, [find("batch", 20), find("fresh", 20, environment="Browser")],
                                ["Node · lotes de 500 dados", "Chromium · semente nova a cada dado"]):
        n = record["n"]
        ax.bar(np.arange(1, 21), np.array(record["counts"]) / n * 100, color="#76bfff", width=0.75)
        lo, hi = np.array(record["faceNullEnvelope"]) / n * 100
        ax.axhspan(lo, hi, color="#d9bd80", alpha=0.16, label="Faixa binomial simultânea sob H₀ (99%, 500 faces)")
        ax.axhline(5, color="#f6d492", linestyle="--", linewidth=1.5, label="P(face) = 5%")
        ax.set(xlabel="Face do d20", ylabel="Frequência (%)", xticks=np.arange(1, 21), ylim=(0, 6.3),
               title=f"{label} · N = {n:,}".replace(",", "."))
        ax.grid(axis="y")
        ax.legend(loc="lower right", facecolor="#1d1b2b", labelcolor="#f0eff5", fontsize=9)
    fig.suptitle("d20: observações da API pública, sem forçar equilíbrio", fontsize=17, weight="bold")
    save(fig, "d20-frequency")

    fig, axes = plt.subplots(2, 4, figsize=(13, 7.8), constrained_layout=True)
    for ax, sides in zip(axes.flat, protocol["sides"]):
        record = find("batch", sides)
        x = np.arange(1, sides + 1)
        residual = (np.array(record["counts"]) - record["expectedCount"]) / np.sqrt(record["expectedCount"] * (1 - 1 / sides))
        ax.bar(x, residual, color="#76bfff", width=0.8)
        ax.axhline(0, color="#dedce7", linewidth=0.8)
        ax.set(title=f"d{sides} · N = 1.000.000", xlabel="Face", ylabel="Desvio / σ", ylim=(-5, 5))
        ax.grid(axis="y")
        if sides <= 12: ax.set_xticks(x)
        elif sides == 20: ax.set_xticks([1, 5, 10, 15, 20])
    axes.flat[-1].axis("off")
    axes.flat[-1].text(0.05, 0.65, "xoshiro128**\n\nσ = √[Np(1 − p)]\np = 1 / lados\n\nDesvios pequenos são\nesperados mesmo sob H₀.", fontsize=13)
    fig.suptitle("Sete dados: resíduos de frequência, com escala comparável", fontsize=17, weight="bold")
    save(fig, "all-dice-residuals")

    fig, axes = plt.subplots(1, 2, figsize=(11.6, 4.9), constrained_layout=True)
    faces = np.arange(1, 21)
    naive = np.bincount(np.arange(256) % 20 + 1, minlength=21)[1:] / 256 * 100
    rejected = np.bincount(np.arange(240) % 20 + 1, minlength=21)[1:] / 240 * 100
    for ax, bars, title in zip(axes, [naive, rejected], ["Módulo direto: 256 valores", "Rejeição: aceita somente 0…239"]):
        ax.bar(faces, bars, color="#f4a1a8" if ax is axes[0] else "#89d7b1")
        ax.axhline(5, color="#f6d492", linestyle="--")
        ax.set(title=title, xlabel="Face do d20", ylabel="Probabilidade exata (%)", xticks=[1, 5, 10, 15, 20], ylim=(4.4, 5.3))
        ax.grid(axis="y")
    fig.suptitle("Exemplo didático de 8 bits: viés de módulo e sua eliminação", fontsize=16, weight="bold")
    save(fig, "rejection-sampling")

    fig, ax = plt.subplots(figsize=(11.6, 4.8), constrained_layout=True)
    labels = [("batch", "Node.js", "#76bfff"), ("fresh", "Node.js", "#89d7b1"), ("fresh", "Browser", "#f4a1a8")]
    lags = np.arange(1, protocol["autocorrelationLags"] + 1)
    z = stats.norm.isf(protocol["alpha"] / (2 * len(records) * len(lags)))
    for mode, environment, color in labels:
        record = find(mode, 20, environment=environment)
        ax.plot(lags, record["autocorrelation"], marker="o", markersize=4, color=color,
                label=f"{environment} · {mode} · N={record['n']:,}".replace(",", "."))
        if mode == "fresh" and environment == "Browser":
            ax.axhspan(-z / np.sqrt(record["n"]), z / np.sqrt(record["n"]), color="#d9bd80", alpha=0.1,
                       label="Referência assintótica de 99%, Bonferroni (440 lags); N=50.000")
    ax.axhline(0, color="#dedce7", linewidth=0.8)
    ax.set(xlabel="Defasagem (lag)", ylabel="Autocorrelação estimada", xticks=lags,
           title="d20: dependência linear nas 20 primeiras defasagens")
    ax.grid()
    ax.legend(facecolor="#1d1b2b", labelcolor="#f0eff5", fontsize=9, loc="upper right")
    save(fig, "d20-autocorrelation")

    fig, axes = plt.subplots(1, 3, figsize=(13, 4.6), constrained_layout=True)
    d20 = values_by_id[find("batch", 20)["id"]].astype(np.int64).reshape(-1, 2)
    d6 = values_by_id[find("batch", 6)["id"]].astype(np.int64).reshape(-1, 2)
    derived = [(np.max(d20, axis=1), faces, (2 * faces - 1) / 400, "Máximo de 2d20 (vantagem)"),
               (np.min(d20, axis=1), faces, (41 - 2 * faces) / 400, "Mínimo de 2d20 (desvantagem)"),
               (np.sum(d6, axis=1), np.arange(2, 13), (6 - np.abs(7 - np.arange(2, 13))) / 36, "Soma de 2d6")]
    for ax, (sample, xs, theory, title) in zip(axes, derived):
        counts = np.bincount(sample, minlength=int(xs[-1]) + 1)[xs]
        ax.bar(xs, counts / len(sample) * 100, color="#76bfff", label="Derivados dos dados brutos")
        ax.plot(xs, theory * 100, color="#f6d492", marker="o", markersize=3, label="Modelo analítico")
        ax.set(title=title, xlabel="Resultado selecionado / soma", ylabel="Probabilidade (%)")
        ax.grid(axis="y")
    axes[0].legend(facecolor="#1d1b2b", labelcolor="#f0eff5", fontsize=8)
    fig.suptitle("Dados justos não tornam todo resultado agregado uniforme", fontsize=16, weight="bold")
    save(fig, "derived-distributions")


def english_label(label):
    translations = {
        "Node · lotes de 500 dados": "Node · batches of 500 dice",
        "Chromium · semente nova a cada dado": "Chromium · fresh seed for every die",
        "Faixa binomial simultânea sob H₀ (99%, 500 faces)": "Simultaneous binomial null envelope (99%, 500 face counts)",
        "Face do d20": "d20 face", "Frequência (%)": "Frequency (%)",
        "d20: observações da API pública, sem forçar equilíbrio": "d20: public API observations, without enforced balance",
        "Desvio / σ": "Deviation / σ", "Face": "Face",
        "Sete dados: resíduos de frequência, com escala comparável": "Seven dice: frequency residuals on a shared scale",
        "xoshiro128**\n\nσ = √[Np(1 − p)]\np = 1 / lados\n\nDesvios pequenos são\nesperados mesmo sob H₀.": "xoshiro128**\n\nσ = √[Np(1 − p)]\np = 1 / sides\n\nSmall deviations are\nexpected even under H₀.",
        "Módulo direto: 256 valores": "Direct modulo: 256 values",
        "Rejeição: aceita somente 0…239": "Rejection: accepts only 0…239",
        "Probabilidade exata (%)": "Exact probability (%)",
        "Exemplo didático de 8 bits: viés de módulo e sua eliminação": "8-bit teaching example: modulo bias and its removal",
        "Defasagem (lag)": "Lag", "Autocorrelação estimada": "Estimated autocorrelation",
        "d20: dependência linear nas 20 primeiras defasagens": "d20: linear dependence at the first 20 lags",
        "Referência assintótica de 99%, Bonferroni (440 lags); N=50.000": "99% asymptotic Bonferroni guide (440 lags); N=50,000",
        "Máximo de 2d20 (vantagem)": "Maximum of 2d20 (advantage)",
        "Mínimo de 2d20 (desvantagem)": "Minimum of 2d20 (disadvantage)",
        "Soma de 2d6": "Sum of 2d6", "Derivados dos dados brutos": "Derived from retained raw faces",
        "Modelo analítico": "Analytical model", "Resultado selecionado / soma": "Selected result / sum",
        "Probabilidade (%)": "Probability (%)",
        "Dados justos não tornam todo resultado agregado uniforme": "Fair dice do not make every aggregate outcome uniform",
    }
    for original, translated in translations.items():
        label = label.replace(original, translated)
    label = label.replace(" · batch · ", " · batches · ").replace(" · fresh · ", " · fresh seeds · ")
    return re.sub(r"(?<=\d)\.(?=\d{3}(?:\D|$))", ",", label)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--data", type=Path, default=Path("audit-data"))
    parser.add_argument("--out", type=Path, default=Path("analysis"))
    args = parser.parse_args()
    manifest = json.loads((args.data / "manifest.json").read_text(encoding="utf-8"))
    protocol = json.loads((args.data / "protocol.json").read_text(encoding="utf-8"))
    assert not manifest["smokeTest"], "Smoke samples are not publication data"
    assert sha256(args.data / "protocol.json") == manifest["protocolSha256"]
    assert len(manifest["datasets"]) == 22
    expected = {}
    for algorithm in protocol["algorithms"]:
        for sides in protocol["sides"]:
            expected[f"node-batch-{algorithm}-d{sides}"] = ("Node.js", "batch", algorithm, sides, protocol["batch"]["outcomesPerDataset"])
    for sides in protocol["sides"]:
        expected[f"node-fresh-xoshiro128ss-d{sides}"] = ("Node.js", "fresh", "xoshiro128ss", sides, protocol["fresh"]["outcomesPerDataset"])
    expected["browser-fresh-xoshiro128ss-d20"] = ("Browser", "fresh", "xoshiro128ss", 20, protocol["browser"]["outcomes"])
    observed = {e["id"]: tuple(e[k] for k in ["environment", "mode", "algorithm", "sides", "outcomes"]) for e in manifest["datasets"]}
    assert observed == expected, "Dataset profiles must match the frozen protocol"
    assert manifest["version"] == protocol["version"]
    assert manifest["packageIntegrity"] == protocol["packageIntegrity"]
    assert sum(entry["outcomes"] for entry in manifest["datasets"]) == protocol["totalOutcomes"]
    total_faces = sum(entry["sides"] for entry in manifest["datasets"])
    records, values_by_id = [], {}
    for entry in manifest["datasets"]:
        record, values = analyze(entry, args.data, protocol, total_faces)
        records.append(record)
        values_by_id[entry["id"]] = values
    pvalues = [p for record in records for p in [record["pUniformity"], record["pSerial"]]]
    adjusted = holm(pvalues)
    for index, record in enumerate(records):
        record["pUniformityHolm"], record["pSerialHolm"] = map(float, adjusted[index * 2:index * 2 + 2])
    result = {"studyId": protocol["studyId"], "protocolSha256": manifest["protocolSha256"],
              "alpha": protocol["alpha"], "primaryTestCount": len(pvalues),
              "rejectedPrimaryTests": int(np.sum(adjusted < protocol["alpha"])),
              "totalOutcomes": protocol["totalOutcomes"], "totalFaceEnvelopes": total_faces,
              "analysisRuntime": {"python": platform.python_version(), "numpy": np.__version__,
                                  "scipy": scipy.__version__, "matplotlib": matplotlib.__version__},
              "datasets": records}
    args.out.mkdir(parents=True, exist_ok=True)
    (args.out / "statistics.json").write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    csv_keys = [key for key in records[0] if key not in ["counts", "individualFaceCI95", "autocorrelation", "faceNullEnvelope"]]
    with (args.out / "results.csv").open("w", newline="", encoding="utf-8") as file:
        writer = csv.DictWriter(file, fieldnames=csv_keys, extrasaction="ignore")
        writer.writeheader()
        writer.writerows(records)
    figures(records, values_by_id, protocol, args.out / "figures")
    figures(records, values_by_id, protocol, args.out / "figures/en", language="en")
    print(f"Analyzed {result['totalOutcomes']:,} faces; {len(pvalues)} primary tests; "
          f"{result['rejectedPrimaryTests']} rejected at Holm family alpha={protocol['alpha']}")


if __name__ == "__main__":
    main()
