#!/usr/bin/env python3
"""Reorder the saved tutorial heatmap without rerunning analyses or changing values.

Requires Python 3.10+, numpy, pandas, scipy, matplotlib and a Chinese font.
python replot_integrated_heatmap.py --execution-zip iobrx-eight-execution-records.zip --output-dir redraw
"""
from __future__ import annotations

import argparse
import hashlib
import io
import json
from pathlib import Path
import platform
import zipfile

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib import font_manager
from matplotlib.colors import ListedColormap
from matplotlib.patches import Patch
import numpy as np
import pandas as pd
import scipy
from scipy.cluster.hierarchy import leaves_list, linkage


GROUPS = ["TME1", "TME2", "TME3"]
KINDS = ["Cells", "Function", "LR"]
GROUP_COLORS = ["#268F9C", "#C45F76", "#C4942F"]
KIND_COLORS = ["#516FAD", "#399077", "#9A6F9F"]


def digest(data):
    return hashlib.sha256(data).hexdigest()


def leaf_order(frame, method, metric):
    """Stable input order and orientation; preserve the hierarchy while seriating."""
    frame = frame.sort_index()
    if len(frame) < 2:
        return frame.index.tolist()
    tree = linkage(frame.to_numpy(), method=method, metric=metric, optimal_ordering=True)
    names = frame.index[leaves_list(tree)].tolist()
    return min(names, names[::-1])


def adjacent_distance(frame):
    return float(np.linalg.norm(np.diff(frame.to_numpy(), axis=0), axis=1).mean())


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--execution-zip", type=Path, default=(
        Path(__file__).resolve().parents[1] / "iobrx-eight-execution-records.zip"))
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--font", help="Installed Chinese font family")
    args = parser.parse_args()
    out = args.output_dir.resolve()
    out.mkdir(parents=True, exist_ok=True)
    sources = {}
    with zipfile.ZipFile(args.execution_zip) as archive:
        def read(suffix):
            names = [n for n in archive.namelist() if n.endswith("/results/05_interactions/" + suffix)]
            if len(names) != 1:
                raise ValueError(f"Expected one archived table: {suffix}")
            data = archive.read(names[0])
            sources[names[0]] = digest(data)
            return pd.read_csv(io.BytesIO(data), index_col=0, float_precision="round_trip")
        original = read("figures/fig3_integrated_heatmap.csv")
        primary = read("runs/01_tme_cluster_primary/result.csv").set_index("ID")

    assert original.shape == (61, 412)
    assert original.index.is_unique and original.columns.is_unique and primary.index.is_unique
    assert set(original.columns) == set(primary.index)
    assert np.isfinite(original.to_numpy()).all(), "Missing values require an explicit policy"
    assert (original.std(axis=1) > 0).all(), "Constant rows cannot use correlation distance"
    # This archived figure table already contains whole-cohort row z-scores (ddof=1).
    # Do not z-score within clusters or transform the matrix a second time.
    assert np.allclose(original.mean(axis=1), 0, atol=1e-12)
    assert np.allclose(original.std(axis=1, ddof=1), 1, atol=1e-12)
    clusters = primary.loc[original.columns, "cluster"]
    counts = clusters.value_counts().reindex(GROUPS)
    assert counts.tolist() == [114, 101, 197]
    cell_rows = original.index[:22].tolist()
    assert set(cell_rows) == set(primary.columns) - {"cluster"}
    kinds = pd.Series(["Cells"] * 22 + ["Function"] * 15 + ["LR"] * 24, index=original.index)
    means = original.T.groupby(clusters, observed=True).mean().reindex(GROUPS).T
    # A descriptive row grouping, not a new patient classifier or a significance test.
    # idxmax resolves exact ties by the declared TME1, TME2, TME3 order.
    peak = means.idxmax(axis=1)
    rows = []
    for group in GROUPS:
        for kind in KINDS:
            selected = original.loc[(peak == group) & (kinds == kind)]
            rows.extend(leaf_order(selected, "average", "correlation"))
    columns = []
    adjacency = {}
    for group in GROUPS:
        old_ids = clusters.index[clusters == group].tolist()
        cells = original.loc[cell_rows, old_ids].T
        ids = leaf_order(cells, "ward", "euclidean")
        columns.extend(ids)
        adjacency[group] = {"before": adjacent_distance(cells),
                            "after": adjacent_distance(cells.loc[ids])}
    plotted = original.loc[rows, columns]
    assert plotted.reindex(index=original.index, columns=original.columns).equals(original)
    assert clusters.loc[columns].tolist() == sum(([g] * int(counts[g]) for g in GROUPS), [])
    assert peak.loc[rows].tolist() == sorted(peak.loc[rows], key=GROUPS.index)
    assert kinds.loc[rows].value_counts().to_dict() == {"LR": 24, "Cells": 22, "Function": 15}
    # Recompute the summary from the plotted cells, checking the same label alignment.
    pd.testing.assert_frame_equal(
        plotted.T.groupby(clusters.loc[columns], observed=True).mean().reindex(GROUPS).T,
        means.loc[rows], check_exact=False, atol=1e-12, rtol=1e-12)

    installed = {f.name for f in font_manager.fontManager.ttflist}
    candidates = [args.font] if args.font else ["Microsoft YaHei", "Noto Sans CJK SC", "SimHei", "WenQuanYi Zen Hei"]
    font = next((name for name in candidates if name in installed), None)
    if font is None:
        raise RuntimeError("Install a Chinese font, or pass its family name with --font")
    plt.rcParams.update({"font.family": font, "font.size": 10, "axes.unicode_minus": False,
                         "svg.fonttype": "none", "svg.hashsalt": "iobrx-heatmap-order-v1"})
    fig = plt.figure(figsize=(17, 13), facecolor="white")
    bottom, height = .12, .735
    left, width = .315, .51
    ax = fig.add_axes([left, bottom, width, height])
    annotation = fig.add_axes([left + width + .006, bottom, .009, height])
    summary = fig.add_axes([.85, bottom, .09, height])
    for axis in [ax, annotation, summary]:
        axis.set_xticks([])
        axis.set_yticks([])
        for edge in axis.spines.values():
            edge.set_visible(False)
    heat = ax.imshow(plotted.values, vmin=-3, vmax=3, cmap="RdBu_r",
                     aspect="auto", interpolation="nearest")
    ax.set_yticks(np.arange(61), labels=rows, fontsize=9)
    ax.tick_params(axis="y", length=0, pad=7)
    annotation.imshow([[KINDS.index(k)] for k in kinds.loc[rows]],
                      cmap=ListedColormap(KIND_COLORS), vmin=-.5, vmax=2.5,
                      aspect="auto", interpolation="nearest")
    summary.imshow(means.loc[rows].values, vmin=-3, vmax=3, cmap="RdBu_r",
                   aspect="auto", interpolation="nearest")
    summary.set_xticks(range(3), labels=GROUPS, fontsize=9)
    summary.tick_params(axis="x", length=0, pad=9)
    summary.set_xticks([.5, 1.5], minor=True)
    summary.grid(which="minor", axis="x", color="white", linewidth=1)
    summary.tick_params(which="minor", bottom=False)

    top = fig.add_axes([left, bottom + height + .014, width, .024])
    top.set_xlim(-.5, 411.5)
    top.set_ylim(0, 1)
    top.axis("off")
    start = 0
    for group, color in zip(GROUPS, GROUP_COLORS):
        n = int(counts[group])
        top.barh(.5, n, left=start - .5, height=1, color=color, linewidth=0)
        top.text(start + (n - 1) / 2, .5, f"{group}  ·  {n} 人", ha="center", va="center", color="white", fontsize=11)
        if start:
            ax.axvline(start - .5, color="white", linewidth=2)
            top.axvline(start - .5, color="white", linewidth=2)
        start += n
    row_start = 0
    for group, color in zip(GROUPS, GROUP_COLORS):
        n = int(peak.eq(group).sum())
        y = bottom + height * (1 - (row_start + n / 2) / 61)
        fig.text(.04, y, f"{group}\n相对较高\n{n} 项", color=color, fontsize=12,
                 weight="bold", va="center", ha="left", linespacing=1.7)
        if row_start:
            for axis in [ax, annotation, summary]:
                axis.axhline(row_start - .5, color="white", linewidth=2.5)
        row_start += n
    summary.set_title("组内均值", fontsize=11, pad=26)
    cax = fig.add_axes([.961, .34, .010, .23])
    colorbar = fig.colorbar(heat, cax=cax, ticks=[-3, 0, 3], extend="both")
    colorbar.outline.set_visible(False)
    colorbar.ax.tick_params(length=0, labelsize=9)
    colorbar.ax.set_title("行 z 分数", fontsize=9, pad=12)
    ax.set_xlabel("每列一位患者；组内按 22 类细胞特征的相似性排列", fontsize=11, labelpad=14)
    fig.text(.04, .958, "STAD 微环境：三组患者的共同模式", fontsize=23, weight="bold", color="#243A2E")
    fig.text(.04, .920, "61 个特征  ·  412 位患者  ·  行按在哪一组相对更高排列；原有分群保持不变", fontsize=12, color="#52645A")
    handles = [Patch(facecolor=color, label=label) for color, label in
               zip(KIND_COLORS, ["细胞（22）", "功能（15）", "LR（24）"])]
    fig.legend(handles=handles, loc="lower left", bbox_to_anchor=(.034, .052),
               ncol=3, frameon=False, fontsize=10, title="主图右侧细条：特征类型", title_fontsize=10)
    fig.text(.40, .073, "红 / 蓝：该特征在队列中的相对高 / 低值；主图与组内均值共用色标。", fontsize=10, color="#52645A")
    fig.text(.40, .047, "行按三组平均 z 分数的最高值分组；阶梯形排列用于读图，不是独立验证。", fontsize=10, color="#52645A")
    for ext in ["png", "svg"]:
        fig.savefig(out / f"fig3_integrated_heatmap.{ext}", dpi=180,
                    metadata={"Date": None} if ext == "svg" else {})
    plt.close(fig)

    plotted.index.name = "feature"
    plotted.to_csv(out / "plotted-zscore-matrix.csv", lineterminator="\n")
    means.loc[rows].rename_axis("feature").to_csv(out / "group-mean-zscores.csv", lineterminator="\n")
    pd.DataFrame({"position": np.arange(1, 413), "sample_id": columns,
                  "cluster": clusters.loc[columns].values}).to_csv(out / "sample-order.csv", index=False, lineterminator="\n")
    pd.DataFrame({"position": np.arange(1, 62), "feature": rows,
                  "kind": kinds.loc[rows].values, "highest_mean_group": peak.loc[rows].values}).to_csv(
                      out / "feature-order.csv", index=False, lineterminator="\n")
    report = {
        "revision": "continuous heatmap reorder; no statistical reanalysis",
        "source_table_sha256": sources, "script_sha256": digest(Path(__file__).read_bytes()),
        "environment": {"python": platform.python_version(), "numpy": np.__version__,
                        "pandas": pd.__version__, "scipy": scipy.__version__,
                        "matplotlib": matplotlib.__version__, "font": font},
        "shape": list(plotted.shape), "cluster_counts": counts.to_dict(),
        "row_order": "highest group mean z-score: TME1/TME2/TME3 (same order breaks exact ties); then Cells/Function/LR; within each subset average-linkage correlation clustering with optimal leaf ordering",
        "column_order": "preserve TME1/TME2/TME3 membership; within each group Ward Euclidean clustering with optimal leaf ordering on the original 22 whole-cohort cell z-score features",
        "tie_reproducibility": "IDs sorted before linkage; choose lexicographically smaller forward/reverse leaf order; versions recorded",
        "scale": "reuse archived whole-cohort row z-scores, ddof=1; color saturation +/-3 only; saved cells are not clipped",
        "group_summary": "arithmetic means of the same row z-scores, not raw values or extra patients",
        "adjacent_patient_cell_distance": adjacency,
        "checks": {"unique_ids": True, "all_412_patients_preserved": True,
                   "all_61_features_preserved": True, "all_cell_values_preserved": True,
                   "cluster_membership_preserved": True, "shared_column_order": True,
                   "group_means_match_plotted_cells": True, "finite_nonconstant_rows": True},
        "files": {p.name: {"bytes": p.stat().st_size, "sha256": digest(p.read_bytes())}
                  for p in sorted(out.iterdir()) if p.is_file() and p.suffix in {".png", ".svg", ".csv"}},
    }
    (out / "redraw-validation.json").write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    print(json.dumps({k: v for k, v in report.items() if k != "files"}, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
