#!/usr/bin/env python3
"""Redraw the tutorial oncoprint from its archived result tables (no model rerun).

Requires Python 3.10+, pandas, numpy, matplotlib and a Chinese font.
python replot_oncoprint.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


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


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/07_genomics/" + suffix)]
            if len(names) != 1:
                raise ValueError(f"Expected exactly one archived table: {suffix}")
            data = archive.read(names[0])
            sources[names[0]] = digest(data)
            return pd.read_csv(io.BytesIO(data), keep_default_na=False)
        original = read("figures/fig2_oncoprint_top30.csv").set_index("Hugo_Symbol")
        frequency = read("tables/mutation_frequency_all_genes.csv")
        burden = read("tables/sample_mutation_burden.csv").set_index("patient_id")

    assert original.index.is_unique and original.columns.is_unique and burden.index.is_unique
    assert frequency.Hugo_Symbol.is_unique
    assert set(original.columns) == set(burden.index)
    assert set(np.unique(original.values)) <= {"", "missense", "in-frame", "truncating"}
    assert (burden.n_protein_altering_events >= 0).all()
    # Preserve the original 30-gene panel, including its selection at tied cutoff ranks.
    selected_genes = original.index[:30].tolist()
    frequency = frequency[frequency.Hugo_Symbol.isin(selected_genes)].sort_values(
        ["n_mut", "Hugo_Symbol"], ascending=[False, True])
    genes = frequency.Hugo_Symbol.tolist()
    assert len(genes) == 30 and set(genes) == set(selected_genes)
    panel = original.loc[genes]
    presence = panel.ne("").astype(int)
    cohort_n = len(panel.columns)
    observed = presence.sum(axis=1)
    expected = frequency.set_index("Hugo_Symbol").loc[genes, "n_mut"]
    assert observed.equals(expected.rename(observed.name)), "Gene counts differ from source frequency table"
    assert np.allclose(expected / cohort_n, frequency.set_index("Hugo_Symbol").loc[genes, "frac_mut"])

    # One permutation for the whole matrix and every annotation. Never sort rows independently.
    order = sorted(panel.columns, key=lambda patient: (
        *(-int(x) for x in presence[patient]),
        -int(burden.loc[patient, "n_protein_altering_events"]), patient))
    plotted = panel.loc[:, order]
    events = burden.loc[order, "n_protein_altering_events"].astype(int)
    empty_panel = plotted.eq("").all(axis=0)
    zero_all = events.eq(0)
    assert plotted.loc[:, panel.columns].equals(panel), "Reordering changed cells"
    assert not (zero_all & ~empty_panel).any()
    assert (plotted.ne("").sum(axis=0) <= events).all()
    assert cohort_n == 409 and int(empty_panel.sum()) == 22 and int(zero_all.sum()) == 1
    patterns = [tuple(int(x) for x in plotted[p].ne("")) for p in order]
    assert patterns == sorted(patterns, reverse=True)

    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-waterfall-v1"})
    fig = plt.figure(figsize=(16, 9.4), facecolor="white")
    grid = fig.add_gridspec(2, 2, height_ratios=[1.25, 6.7], width_ratios=[12, 2.25],
                           left=.060, right=.975, bottom=.155, top=.865, hspace=.12, wspace=.025)
    top = fig.add_subplot(grid[0, 0])
    matrix = fig.add_subplot(grid[1, 0], sharex=top)
    right = fig.add_subplot(grid[1, 1], sharey=matrix)
    x = np.arange(cohort_n)
    top.bar(x, events.values, width=.92, color="#718078", linewidth=0)
    top.set_xlim(-.5, cohort_n - .5)
    top.set_ylim(0, max(events) * 1.06)
    top.set_yticks([0, 2000, 4000, 6000])
    top.tick_params(axis="both", length=0, labelsize=9, labelbottom=False)
    top.set_ylabel("蛋白改变型\n事件数", fontsize=10)
    top.set_title("每位患者的总事件数（全部纳入基因；与下方患者顺序一致）", loc="left", fontsize=10, pad=9)
    for edge in ["top", "right", "bottom"]:
        top.spines[edge].set_visible(False)
    top.spines["left"].set_color("#CED6D0")

    colors = ["#F2F4F2", "#86BAD8", "#F2B27B", "#AE2742"]
    codes = {"": 0, "missense": 1, "in-frame": 2, "truncating": 3}
    encoded = np.array([[codes[v] for v in row] for row in plotted.values])
    matrix.imshow(encoded, cmap=ListedColormap(colors), vmin=-.5, vmax=3.5,
                  aspect="auto", interpolation="nearest")
    matrix.set_yticks(np.arange(len(genes)), labels=genes, fontsize=10)
    matrix.set_xticks([])
    matrix.tick_params(axis="y", length=0, pad=6)
    matrix.set_yticks(np.arange(-.5, len(genes), 1), minor=True)
    matrix.grid(which="minor", axis="y", color="white", linewidth=1.1)
    matrix.tick_params(which="minor", left=False)
    matrix.set_xlabel("患者按上方基因的突变组合依次分组；每一列始终对应同一位患者", labelpad=12, fontsize=10)
    for edge in matrix.spines.values():
        edge.set_visible(False)
    first_empty = cohort_n - int(empty_panel.sum())
    for axis in [matrix, top]:
        axis.axvline(first_empty - .5, color="#7E8D83", linewidth=.9, linestyle=(0, (3, 3)))

    rates = expected.to_numpy() / cohort_n * 100
    right.barh(np.arange(len(genes)), rates, height=.70, color="#7C9789")
    right.set_xlim(0, 105)
    right.tick_params(axis="y", which="both", left=False, labelleft=False)
    right.set_xticks([0, 25, 50], labels=["0", "25", "50%"], fontsize=9)
    right.tick_params(axis="x", length=0, pad=6)
    right.set_title("突变患者比例（人数）", fontsize=10, loc="left", pad=10)
    for i, (rate, count) in enumerate(zip(rates, expected)):
        right.text(59, i, f"{rate:.1f}% ({int(count)})", va="center", fontsize=9)
    for edge in right.spines.values():
        edge.set_visible(False)

    fig.text(.06, .955, "STAD 突变概览", fontsize=21, weight="bold", color="#243A2E")
    fig.text(.06, .914, "30 个高频基因  ·  409 位患者  ·  按突变组合瀑布式排序", fontsize=12, color="#52645A")
    legend = [Patch(facecolor=colors[3], label="截短 / 剪接 / 起止密码子改变"),
              Patch(facecolor=colors[2], label="框内插入 / 缺失"),
              Patch(facecolor=colors[1], label="错义突变"),
              Patch(facecolor=colors[0], edgecolor="#D7DDD8", label="本图基因中无符合定义的事件")]
    fig.legend(handles=legend, loc="lower left", bbox_to_anchor=(.054, .066), ncol=4,
               frameon=False, fontsize=10, handlelength=1.4, columnspacing=2)
    fig.text(.06, .052, "右侧虚线后 22 位患者在这 30 个基因中无事件；其中 1 位在全部纳入基因中无事件。分母始终为 409 位患者。", fontsize=9, color="#52645A")
    fig.text(.06, .028, "同一基因与患者有多类事件时沿用原图优先级：截短 / 剪接 / 起止密码子改变 > 框内插入 / 缺失 > 错义。仅重排展示，不改变事件与统计结果。", fontsize=9, color="#52645A")
    for ext in ["png", "svg"]:
        metadata = {"Date": None} if ext == "svg" else {}
        fig.savefig(out / f"fig2_oncoprint_top30.{ext}", dpi=180, metadata=metadata)
    plt.close(fig)

    plotted.to_csv(out / "plotted-matrix.csv", lineterminator="\n")
    pd.DataFrame({"position": np.arange(1, cohort_n + 1), "patient_id": order,
                  "n_protein_altering_events": events.values,
                  "no_events_in_displayed_top30": empty_panel.values,
                  "zero_events_all_analyzed_genes": zero_all.values}).to_csv(
                      out / "sample-order.csv", index=False, lineterminator="\n")
    pd.DataFrame({"position": np.arange(1, 31), "gene": genes, "n_mut": expected.values,
                  "denominator": cohort_n, "percent": rates}).to_csv(
                      out / "gene-order.csv", index=False, lineterminator="\n")
    report = {
        "revision": "waterfall figure redraw; no statistical reanalysis",
        "source_table_sha256": sources,
        "script_sha256": digest(Path(__file__).read_bytes()),
        "environment": {"python": platform.python_version(), "pandas": pd.__version__,
                        "numpy": np.__version__, "matplotlib": matplotlib.__version__, "font": font},
        "row_order": "mutated-patient count descending, gene symbol ascending for ties",
        "gene_selection": "same 30 genes as the original figure; cutoff ties do not reselect genes",
        "column_order": "binary lexicographic descending in displayed gene order; ties: total event count descending, patient ID ascending",
        "cohort_n": cohort_n, "genes_n": len(genes),
        "no_events_in_displayed_top30": int(empty_panel.sum()),
        "zero_events_all_analyzed_genes": int(zero_all.sum()),
        "checks": {"unique_ids": True, "cohort_preserved": True, "gene_panel_preserved": True, "cell_values_preserved": True,
                   "frequency_counts_match_source": True, "all_panels_share_patient_order": True,
                   "binary_waterfall_order": 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()
