"""Deterministic comparison: run from any working directory after build.sh."""
from __future__ import annotations

import hashlib
import itertools
import json
import subprocess
import sys
import time
from fractions import Fraction
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "src"))
import numpy as np
from resampling import canonical_bytes, core_ranks, digest, expected, fixture, produce, producer_assignments, verify
from bridge import run_r_bytes, strict_json

OUT = ROOT / "results"
OUT.mkdir(exist_ok=True)


def save(name, data):
    (OUT / name).write_text(json.dumps(data, indent=2, sort_keys=True) + "\n")


def fraction_string(value):
    return f"{value.numerator}/{value.denominator}"


def reject_fraction(rows, alpha=Fraction(1, 20)):
    return Fraction(sum(Fraction(r["p_num"], r["p_den"]) <= alpha for r in rows), len(rows))


def run_r(data, mode, name):
    evidence, input_bytes, output_bytes = run_r_bytes(canonical_bytes(data), mode)
    # Archive the same immutable snapshots that were executed and checked.
    (OUT / f"{name}-input.json").write_bytes(input_bytes)
    (OUT / f"{name}-output.json").write_bytes(output_bytes)
    output = strict_json(output_bytes)
    evidence["reject_fraction"] = fraction_string(reject_fraction(output["rows"]))
    return evidence


def comparison():
    # Fixed ordinary fixtures; golden p numerators are specified independently.
    ordinary = [
        (fixture([[0], [0], [1], [1]]), [2, 6, 6, 6, 6, 2], [4, 0, 0, 0, 0, 4]),
        (fixture([[3], [3], [3], [3]]), [6] * 6, [0] * 6),
    ]
    rng = np.random.default_rng(20261002)
    strong = [x for x, _, _ in ordinary]
    for i in range(40):
        n = [4, 6, 8][i % 3]
        p = [1, 3, 10, 35][i % 4]
        matrix = rng.integers(0, 20, size=(n, p)).tolist()
        if i % 2:
            blocks = [f"b{j//2}" for j in range(n)]
            strong.append(fixture(matrix, blocks, {b: 1 for b in blocks}))
        else:
            strong.append(fixture(matrix))
    targets = [expected(data) for data in strong]
    modes = ["correct", "strict_tail", "drop_assignment", "freeze_selection", "wrong_denominator",
             "ignore_blocks", "duplicate_assignment", "wrong_score", "wrong_identity"]
    results = []
    for mode in modes:
        ordinary_failed = []
        for i, (data, nums, scores) in enumerate(ordinary):
            out = produce(data, mode)
            if len(out["rows"]) != 6 or [r["p_num"] for r in out["rows"]] != nums or [r["score"] for r in out["rows"]] != scores or any(r["p_den"] != 6 for r in out["rows"]):
                ordinary_failed.append(i)
        candidates = [produce(data, mode) for data in strong]
        strong_failed = [i for i, (a, b) in enumerate(zip(candidates, targets)) if a != b]
        spec_failed = []
        for i, (data, output) in enumerate(zip(strong, candidates)):
            try:
                verify(data, output)
            except ValueError as error:
                spec_failed.append({"fixture": i, "reason": str(error)})
        results.append({"producer": mode, "ordinary_detected": bool(ordinary_failed),
                        "strong_detected": bool(strong_failed), "spec_detected": bool(spec_failed),
                        "ordinary_failed_fixtures": ordinary_failed,
                        "strong_failed_count": len(strong_failed), "spec_failed_count": len(spec_failed),
                        "first_spec_failure": spec_failed[0] if spec_failed else None})
    assert not any(results[0][key] for key in ("ordinary_detected", "strong_detected", "spec_detected"))
    assert all(r["strong_detected"] and r["spec_detected"] for r in results[1:])
    save("assurance-comparison.json", {"ordinary_fixtures": 2, "strong_fixtures": len(strong), "spec_fixtures_per_mode": len(strong), "mutants": results})
    return strong


def finite_checks():
    vectors = [list(xs) for n in range(1, 8) for xs in itertools.product(range(3), repeat=n)]
    ranks = core_ranks(vectors)
    for xs, ns in zip(vectors, ranks):
        assert ns == [sum(y >= x for y in xs) for x in xs]
        assert all(sum(r <= k for r in ns) <= k for k in range(len(xs)+1))
    checked, worst_correct, worst_frozen = 0, Fraction(0), Fraction(0)
    failures = 0
    for bits in itertools.product((0, 1), repeat=12):
        data = fixture([list(bits[2*i:2*i+2]) for i in range(6)])
        correct, frozen = produce(data), produce(data, "freeze_selection")
        assert correct == expected(data)
        good = reject_fraction(correct["rows"], Fraction(1, 10))
        bad = reject_fraction(frozen["rows"], Fraction(1, 10))
        assert good <= Fraction(1, 10)
        worst_correct, worst_frozen = max(worst_correct, good), max(worst_frozen, bad)
        failures += bad > Fraction(1, 10)
        checked += 1
    save("finite-checks.json", {"lean_python_rank_vectors": len(vectors), "score_alphabet": [0, 1, 2], "lengths": [1, 7],
                                "binary_matrices": checked, "matrix_shape": [6, 2], "assignments_each": 20,
                                "alpha": "1/10", "worst_correct_rejection_fraction": fraction_string(worst_correct),
                                "worst_frozen_rejection_fraction": fraction_string(worst_frozen),
                                "matrices_with_frozen_inflation": failures})


def statistical_examples(strong):
    base = fixture([[0] for _ in range(8)])
    assignments = producer_assignments(base)
    # One representative of each complementary pair: 35 fixed binary features.
    panel = [z for z in assignments if z[0] == 0]
    adversarial = fixture([[z[i] for z in panel] for i in range(8)])
    save("adversarial-data.json", adversarial)
    good, bad = produce(adversarial), produce(adversarial, "freeze_selection")
    verify(adversarial, good)
    assert reject_fraction(good["rows"]) == 0
    assert reject_fraction(bad["rows"]) == 1
    # External preselection can violate history while passing every numerical check.
    external_accepted, external_rejected, truthful_refused = 0, 0, 0
    for observed in assignments:
        m = sum(observed)
        effects = [abs(8 * sum(adversarial["matrix"][i][g] for i in range(8) if observed[i])
                       - m * sum(row[g] for row in adversarial["matrix"])) for g in range(35)]
        winner = max(range(35), key=effects.__getitem__)
        selected = fixture([[row[winner]] for row in adversarial["matrix"]])
        selected["panel_history"]["basis"] = "Deliberately false fixed-panel declaration for the boundary experiment"
        candidate = produce(selected)
        external_accepted += verify(selected, candidate)
        observed_key = "".join(map(str, observed))
        actual = next(row for row in candidate["rows"] if row["assignment"] == observed_key)
        external_rejected += Fraction(actual["p_num"], actual["p_den"]) <= Fraction(1,20)
        selected["panel_history"] = {"status": "selected_using_observed_assignment", "basis": "Observed winning feature retained"}
        try:
            verify(selected, candidate)
        except ValueError:
            truthful_refused += 1
    assert (external_accepted, external_rejected, truthful_refused) == (70,70,70)
    save("external-selection-boundary.json", {"observed_assignments": 70, "false_fixed_panel_declarations_accepted": external_accepted,
                                             "original_global_null_rejections": external_rejected, "observed_pvalue": "1/35",
                                             "truthful_adaptive_history_declarations_refused": truthful_refused,
                                             "conclusion": "The checker certifies the supplied-panel table, not actual panel-selection history"})
    r_evidence = {"adversarial_correct": run_r(adversarial, "correct", "r-adversarial-correct"),
                  "adversarial_frozen": run_r(adversarial, "freeze_selection", "r-adversarial-frozen"),
                  "blocked_correct": run_r(strong[3], "correct", "r-blocked-correct"),
                  "random_correct": run_r(strong[4], "correct", "r-random-correct")}
    assert r_evidence["adversarial_correct"]["accepted"]
    assert not r_evidence["adversarial_frozen"]["accepted"]
    assert r_evidence["blocked_correct"]["accepted"] and r_evidence["random_correct"]["accepted"]
    save("r-correspondence.json", r_evidence)
    results = {"adversarial": {"units": 8, "features": 35, "assignments": 70, "alpha": "1/20",
                              "correct_pvalue": "1/1", "frozen_pvalue": "1/35",
                              "correct_rejection_fraction": fraction_string(reject_fraction(good["rows"])),
                              "frozen_rejection_fraction": fraction_string(reject_fraction(bad["rows"]))}}
    rng = np.random.default_rng(104729)
    rows = []
    for i in range(200):
        mean = rng.gamma(2.0, 5.0, size=(1, 100))
        if i >= 100:
            mean = mean * rng.gamma(2.0, 0.5, size=(8, 100))
        data = fixture(rng.poisson(mean, size=(8, 100)).tolist())
        a = reject_fraction(produce(data)["rows"])
        b = reject_fraction(produce(data, "freeze_selection")["rows"])
        assert a <= Fraction(1, 20)
        rows.append({"index": i, "generator": "poisson" if i < 100 else "gamma_poisson",
                     "correct": fraction_string(a), "frozen": fraction_string(b)})
    save("count-matrix-rejection-fractions.json", rows)
    for generator in ("poisson", "gamma_poisson"):
        subset = [r for r in rows if r["generator"] == generator]
        results[generator] = {"matrices": len(subset), "units": 8, "features": 100, "assignments_each": 70,
                              "correct_mean": fraction_string(sum(Fraction(r["correct"]) for r in subset) / len(subset)),
                              "frozen_mean": fraction_string(sum(Fraction(r["frozen"]) for r in subset) / len(subset)),
                              "frozen_inflated_matrices": sum(Fraction(r["frozen"]) > Fraction(1, 20) for r in subset)}
    # Same correct computation, different actual assignment law.
    data = fixture([[0], [0], [0], [0], [1], [1], [1], [1]])
    out = produce(data)
    verify(data, out)
    rejected = [r for r in out["rows"] if Fraction(r["p_num"], r["p_den"]) <= Fraction(1, 20)]
    assert len(rejected) == 2
    results["assumption_failure"] = {"checker_accepted": True, "uniform_rejection_probability": "1/35",
                                     "actual_law": "Each of the two rejecting assignments has mass 9/20; the other 68 share 1/10",
                                     "actual_rejection_probability": "9/10",
                                     "meaning": "A declaration of uniform assignment does not prove the real assignment mechanism"}
    save("statistical-results.json", results)


def main():
    start = time.perf_counter()
    strong = comparison()
    finite_checks()
    statistical_examples(strong)
    files = sorted(str(p.relative_to(ROOT)) for folder in ("formal", "src", "tests", "experiments") for p in (ROOT/folder).iterdir() if p.suffix in {".lean", ".py", ".R", ".sh"}) + ["formal/rank-core", "formal/build.json", "docs/contract.md"]
    save("run-manifest.json", {"python": sys.version, "numpy": np.__version__,
                               "lean": subprocess.check_output(["lean", "--version"], text=True).strip(),
                               "R": subprocess.check_output(["Rscript", "--version"], text=True).strip(),
                               "jsonlite": subprocess.check_output(["Rscript", "-e", "cat(as.character(packageVersion('jsonlite')))"], text=True).strip(),
                               "sha256": {name: hashlib.sha256((ROOT/name).read_bytes()).hexdigest() for name in files}})
    print(json.dumps({"status": "all assertions passed", "elapsed_seconds": round(time.perf_counter()-start, 3), "results": str(OUT)}, indent=2))


if __name__ == "__main__":
    main()
