"""Generate Lean kernel-checkable claims from exact decoded R binary64 outputs.

This decoder/generator is trusted for byte-to-value correspondence, not arithmetic
proof. Lean checks every generated rational bound with kernel-checked `norm_num` proofs.
"""
from __future__ import annotations
import csv
import hashlib
import json
import math
import struct
from fractions import Fraction
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "evidence/outputs"
EPS = Fraction(1, 10**12)

def read(case: str, kind: str, size: int) -> list[Fraction]:
    raw = (OUT / f"{case}-{kind}.bin").read_bytes()
    assert len(raw) == size * 8
    values = struct.unpack("<" + "d" * size, raw)
    assert all(math.isfinite(v) and v > 0 for v in values)
    return [Fraction.from_float(v) for v in values]

def lean_q(value: Fraction) -> str:
    return f"({value.numerator} / {value.denominator} : ℚ)"

checks: list[dict] = []
lines = ["import Normalization", "open DESeq2Verification", "namespace ProductionCertificates",
         "set_option maxRecDepth 10000", "set_option maxHeartbeats 0"]

def certify(name: str, actual: Fraction, expected: Fraction, kind: str) -> None:
    assert expected > 0
    relative = abs(actual - expected) / expected
    assert relative <= EPS, (name, relative)
    a, e, eps = map(lean_q, (actual, expected, EPS))
    lines.extend([
        f"def actual_{name} : ℚ := {a}",
        f"def expected_{name} : ℚ := {e}",
        f"theorem accepted_{name} : checkRelative actual_{name} expected_{name} {eps} = true := by",
        f"  norm_num [checkRelative, WithinRelative, actual_{name}, expected_{name}]",
        f"theorem bound_{name} : |actual_{name} - expected_{name}| / expected_{name} ≤ {eps} :=",
        f"  checkRelative_sound _ _ _ accepted_{name}"])
    checks.append(dict(name=name, kind=kind, actual=str(actual), expected=str(expected),
                       relative_error=str(relative), relative_error_float=float(relative)))

with (ROOT / "evidence/cases.tsv").open() as stream:
    cases = list(csv.DictReader(stream, delimiter="\t"))
for case in cases:
    name, n = case["case"], int(case["samples"])
    scale = [Fraction(v) for v in case["scales"].split(",")]
    base, scaled, globally, filtered = [read(name, k, n) for k in ("base", "scaled", "global", "filtered")]
    for j in range(n):
        certify(f"{name}_global_{j}", globally[j], base[j], "global_scaling")
        if int(case["retained"]) != int(case["rows"]):
            certify(f"{name}_filtered_{j}", filtered[j], base[j], "zero_row_filtering")
        # Exact cross-products express equivariance of relative sample size factors.
        if j != 0:
            certify(f"{name}_sample_{j}", scaled[j] * base[0] * scale[0],
                    scaled[0] * base[j] * scale[j], "sample_scaling_cross_product")

# The checker must reject altered values and a zero reference denominator.
lines.extend([
    "example : checkRelative (actual_odd_global_0 * (101/100)) expected_odd_global_0 (1/1000000000000) = false := by",
    "  norm_num [checkRelative, WithinRelative, actual_odd_global_0, expected_odd_global_0]",
    "example : checkRelative (101/100) 1 (1/1000000000000) = false := by norm_num [checkRelative, WithinRelative]",
    "example : checkRelative 0 0 (1/1000000000000) = false := by norm_num [checkRelative, WithinRelative]",
    "#eval (" + " && ".join(f"checkRelative actual_{c['name']} expected_{c['name']} (1/1000000000000)" for c in checks) + ")",
    "#print axioms bound_odd_sample_1", "end ProductionCertificates", ""])
(ROOT / "ProductionCertificates.lean").write_text("\n".join(lines))
report = dict(checks=len(checks), matrices=len(cases), epsilon=str(EPS),
              max_relative_error=max(c["relative_error_float"] for c in checks), details=checks,
              nonzero_residuals=sum(c["actual"] != c["expected"] for c in checks),
              binary_sha256={p.name: hashlib.sha256(p.read_bytes()).hexdigest()
                             for p in sorted(OUT.glob("*.bin"))})
(ROOT / "evidence/certificates.json").write_text(json.dumps(report, indent=2) + "\n")
print(json.dumps({k: report[k] for k in ("checks", "matrices", "epsilon", "max_relative_error")}))
