"""Export actual Scanpy outputs as literal Lean certificates. No Python sum oracle."""
from __future__ import annotations

import argparse
import copy
import hashlib
import importlib.metadata as metadata
import json
import math
import os
from pathlib import Path
import subprocess
import sys

ROOT = Path(__file__).resolve().parents[1]
os.environ.setdefault("MPLCONFIGDIR", str(ROOT / ".cache/matplotlib"))
os.environ.setdefault("NUMBA_CACHE_DIR", str(ROOT / ".cache/numba"))
os.environ.setdefault("NUMBA_NUM_THREADS", "2")
os.environ.setdefault("NUMBA_THREADING_LAYER", "workqueue")
os.environ.setdefault("NUMBA_DISABLE_JIT", "0")
os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
import anndata as ad
import numpy as np
import pandas as pd
import scanpy as sc
from scipy import sparse

BOUND = 2**53
COMMIT = "ff2133115f639fccce9633aa7f5309c36a0972e7"


def dump(path, value):
    Path(path).write_text(json.dumps(value, indent=2, sort_keys=True, ensure_ascii=True) + "\n")


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


def label(value):
    if not isinstance(value, str) or not value or any(not 32 <= ord(c) <= 126 for c in value):
        raise ValueError("labels must be nonempty printable ASCII strings")
    return value


def natural(value):
    # Float results are interpreted exactly as binary64 values, never rounded to integers.
    if isinstance(value, (np.floating, float)):
        if not math.isfinite(value):
            raise ValueError("nonfinite output")
        numerator, denominator = float(value).as_integer_ratio()
        if denominator != 1 or numerator < 0:
            raise ValueError("output is not an exact nonnegative integer")
        return numerator
    if isinstance(value, (np.integer, int)) and not isinstance(value, (bool, np.bool_)):
        if value < 0:
            raise ValueError("negative integer")
        return int(value)
    raise ValueError("unsupported numeric type")


def make_adata(spec, backend):
    # Keep the raw matrix's mathematical integers before any narrowing conversion.
    features = [label(x) for x in spec["features"]]
    categories = [label(x) for x in spec["categories"]]
    rows = spec["rows"]
    if not features or not rows:
        raise ValueError("the executable contract requires nonempty observations/features")
    if len(set(features)) != len(features) or len(set(categories)) != len(categories):
        raise ValueError("duplicate features/categories")
    names = [label(r["name"]) for r in rows]
    if len(set(names)) != len(names):
        raise ValueError("duplicate observation names")
    for row in rows:
        if type(row["keep"]) is not bool:
            raise ValueError("mask must be Boolean")
        if row["group"] is not None and label(row["group"]) not in categories:
            raise ValueError("group outside declared categories")
        if len(row["counts"]) != len(features):
            raise ValueError("ragged matrix")
        if any(type(x) is not int or x < 0 or x > BOUND for x in row["counts"]):
            raise ValueError("counts must be integers in [0, 2^53]")
    if all(r["group"] is None for r in rows):
        raise ValueError("at least one assigned group is required by this adapter")
    x = np.array([r["counts"] for r in rows], dtype=np.int64)
    if backend == "csr":
        # Store explicit zeros deliberately; use canonical indices with no duplicates.
        n, p = x.shape
        x = sparse.csr_matrix((x.ravel(), np.tile(np.arange(p), n), np.arange(n+1)*p), shape=x.shape)
    elif backend == "csc":
        n, p = x.shape
        x = sparse.csc_matrix((x.T.ravel(), np.tile(np.arange(n), p), np.arange(p+1)*n), shape=x.shape)
    elif backend != "dense":
        raise ValueError("backend must be dense, csr or csc")
    return ad.AnnData(x, obs=pd.DataFrame({"group": pd.Categorical(
        [r["group"] for r in rows], categories=categories, ordered=True),
        "keep": [r["keep"] for r in rows]}, index=names),
        var=pd.DataFrame(index=features))


def snapshot_input(adata):
    x = adata.X
    if x.dtype != np.dtype("int64"):
        raise ValueError("input dtype must be int64")
    if sparse.issparse(x):
        if not isinstance(x, (sparse.csr_matrix, sparse.csc_matrix)) or not x.has_canonical_format:
            raise ValueError("sparse input must be canonical CSR/CSC matrix")
        x.check_format(full_check=True)
        x = x.toarray()
    elif type(x) is not np.ndarray:
        raise ValueError("unsupported matrix representation")
    return {"features": [label(v) for v in adata.var_names],
            "categories": [label(v) for v in adata.obs["group"].cat.categories],
            "rows": [{"id": {"index": i, "name": label(str(name))},
                      "group": None if pd.isna(group) else label(group),
                      "keep": bool(keep), "counts": [natural(v) for v in x[i]]}
                     for i, (name, group, keep) in enumerate(zip(adata.obs_names,
                         adata.obs["group"], adata.obs["keep"], strict=True))]}


def snapshot_output(result, before, backend):
    x = result.layers["sum"]
    wanted_dtype = np.dtype("float64" if backend == "dense" else "int64")
    if type(x) is not np.ndarray or x.dtype != wanted_dtype:
        raise ValueError(f"unexpected output array/dtype: {type(x)}, {x.dtype}")
    groups = []
    for k, (name, annotation, nobs) in enumerate(zip(result.obs_names,
            result.obs["group"], result.obs["n_obs_aggregated"], strict=True)):
        name = label(name)
        # This is an adapter-supplied witness. Scanpy does not return source indices.
        ids = [dict(r["id"]) for r in before["rows"] if r["keep"] and r["group"] == name]
        groups.append({"label": name, "annotation": label(annotation), "sources": ids,
                       "nObs": natural(nobs), "counts": [natural(v) for v in x[k]]})
    return {"features": [label(x) for x in result.var_names], "groups": groups}


def ls(values, encode=str):
    return "[" + ", ".join(encode(v) for v in values) + "]"


def st(value):
    return json.dumps(label(value), ensure_ascii=True)


def identity(value):
    return "⟨" + str(value["index"]) + ", " + st(value["name"]) + "⟩"


def source_row(row):
    group = "none" if row["group"] is None else "some " + st(row["group"])
    return "⟨" + ", ".join([identity(row["id"]), group, str(row["keep"]).lower(), ls(row["counts"])]) + "⟩"


def lean_input(value):
    return "⟨" + ", ".join([ls(value["features"], st), ls(value["categories"], st), ls(value["rows"], source_row)]) + "⟩"


def group_row(row):
    return "⟨" + ", ".join([st(row["label"]), st(row["annotation"]), ls(row["sources"], identity),
                             str(row["nObs"]), ls(row["counts"])]) + "⟩"


def lean_output(value):
    return "⟨" + ", ".join([ls(value["features"], st), ls(value["groups"], group_row)]) + "⟩"


def certificate_text(cases):
    lines = ["import ScanpyAggregate", "open ScanpyAggregate",
             "set_option maxRecDepth 100000", "set_option maxHeartbeats 4000000"]
    for k, case in enumerate(cases):
        lines += [f"-- {case['name']}", f"def input{k} : Input := {lean_input(case['input'])}",
                  f"def output{k} : Output := {lean_output(case['output'])}",
                  f"theorem contract{k} : valid input{k} = {str(case.get('input_valid', True)).lower()} := by decide",
                  f"theorem checked{k} : check input{k} output{k} = {str(case['accept']).lower()} := by decide"]
        if case["accept"]:
            lines += [f"theorem sound{k} : valid input{k} = true ∧ output{k} = expected input{k} :=",
                      f"  check_sound input{k} output{k} checked{k}"]
        lines += [f"#print axioms checked{k}"]
    return "\n".join(lines) + "\n"


def run_kernel(path):
    build = subprocess.run(["lake", "build"], cwd=ROOT / "formal", text=True,
                           stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
    if build.returncode or "sorryAx" in build.stdout or "ofReduceBool" in build.stdout:
        raise RuntimeError("Lean library build failed:\n" + build.stdout)
    result = subprocess.run(["lake", "env", "lean", str(Path(path).resolve())], cwd=ROOT / "formal",
                            text=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
    if result.returncode or "sorryAx" in result.stdout or "ofReduceBool" in result.stdout:
        raise RuntimeError("Lean kernel certification failed:\n" + result.stdout)
    return result.stdout


def collect(spec, backend, name):
    adata = make_adata(spec, backend)
    before = snapshot_input(adata)
    result = sc.get.aggregate(adata, by="group", func="sum", axis="obs", mask="keep")
    after = snapshot_input(adata)
    if before != after:
        raise AssertionError("Scanpy mutated input")
    out = snapshot_output(result, before, backend)
    return {"name": name, "backend": backend, "input": before, "output": out,
            "accept": True, "actual_dtype": str(result.layers["sum"].dtype),
            "input_storage": "dense" if backend == "dense" else "canonical_explicit_zeros",
            "actual_binary64_hex": [[float(v).hex() for v in row] for row in result.layers["sum"]]
                if backend == "dense" else None}


def verify(spec, backend, directory):
    from pins import verify_pins
    verify_pins()
    directory = Path(directory)
    directory.mkdir(parents=True, exist_ok=True)
    (directory / "receipt.json").unlink(missing_ok=True)
    case = collect(spec, backend, "user_run")
    dump(directory / "transcript.json", case)
    proof = directory / "Certificate.lean"
    proof.write_text(certificate_text([case]))
    (directory / "kernel.log").write_text(run_kernel(proof))
    dump(directory / "receipt.json", {"status": "kernel-checked observed output",
        "transcript_sha256": sha(directory / "transcript.json"), "certificate_sha256": sha(proof),
        "checker_sha256": sha(ROOT / "formal/ScanpyAggregate.lean"), "scanpy_commit": COMMIT})
    return case


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path)
    parser.add_argument("--backend", choices=["dense", "csr", "csc"], default="csr")
    parser.add_argument("--evidence", type=Path, required=True)
    args = parser.parse_args()
    verify(json.loads(args.input.read_text()), args.backend, args.evidence)
    print("PASS: Lean kernel checked actual Scanpy output, labels and source witnesses.")
