#!/usr/bin/env python
"""公开数据的真实单细胞分析样例 —— 既是给客户看的交付样品,也是我方产能的自证。

数据:10x Genomics 公开的 PBMC 3k(健康人外周血单个核细胞),scanpy 内置下载。
流程:质控 → 过滤 → 标准化 → 高变基因 → PCA → 近邻图 → UMAP → Leiden 聚类 →
      每群 marker 基因 → 依据经典 marker 做细胞类型标注 → 出图 + 指标 JSON。
所有数字都由本脚本真实算出,可复算;不是示意图。
"""
import json, os, sys, warnings
warnings.filterwarnings("ignore")
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import scanpy as sc
import numpy as np

OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "out")
os.makedirs(OUT, exist_ok=True)
sc.settings.figdir = OUT
sc.settings.verbosity = 1
sc.settings.set_figure_params(dpi=140, frameon=False, figsize=(5, 4.2), facecolor="white")

M = {}   # 交付指标

adata = sc.datasets.pbmc3k()
M["cells_raw"], M["genes_raw"] = int(adata.n_obs), int(adata.n_vars)

# ---------- 质控 ----------
adata.var["mt"] = adata.var_names.str.startswith("MT-")
sc.pp.calculate_qc_metrics(adata, qc_vars=["mt"], percent_top=None, log1p=False, inplace=True)
M["median_genes_per_cell"] = float(np.median(adata.obs["n_genes_by_counts"]))
M["median_umi_per_cell"] = float(np.median(adata.obs["total_counts"]))
M["median_pct_mito"] = round(float(np.median(adata.obs["pct_counts_mt"])), 2)

fig, axes = plt.subplots(1, 3, figsize=(12, 3.6))
for ax, key, title in zip(axes,
                          ["n_genes_by_counts", "total_counts", "pct_counts_mt"],
                          ["Genes per cell", "UMIs per cell", "Mitochondrial %"]):
    ax.violin = ax.violinplot(adata.obs[key].values, showextrema=False)
    ax.set_title(title, fontsize=11); ax.set_xticks([])
plt.tight_layout(); plt.savefig(f"{OUT}/qc_violin.png", dpi=140, bbox_inches="tight"); plt.close()

# ---------- 过滤 ----------
sc.pp.filter_cells(adata, min_genes=200)
sc.pp.filter_genes(adata, min_cells=3)
adata = adata[(adata.obs.n_genes_by_counts < 2500) & (adata.obs.pct_counts_mt < 5), :].copy()
M["cells_after_qc"], M["genes_after_qc"] = int(adata.n_obs), int(adata.n_vars)
M["cells_removed_pct"] = round(100 * (1 - M["cells_after_qc"] / M["cells_raw"]), 1)

# ---------- 标准化 + 高变基因 ----------
adata.layers["counts"] = adata.X.copy()
sc.pp.normalize_total(adata, target_sum=1e4)
sc.pp.log1p(adata)
sc.pp.highly_variable_genes(adata, min_mean=0.0125, max_mean=3, min_disp=0.5)
M["hvg"] = int(adata.var.highly_variable.sum())
adata.raw = adata
adata = adata[:, adata.var.highly_variable].copy()
sc.pp.regress_out(adata, ["total_counts", "pct_counts_mt"])
sc.pp.scale(adata, max_value=10)

# ---------- 降维聚类 ----------
sc.tl.pca(adata, svd_solver="arpack", n_comps=50)
sc.pp.neighbors(adata, n_neighbors=10, n_pcs=30)
sc.tl.umap(adata)
sc.tl.leiden(adata, resolution=0.5, key_added="leiden", flavor="igraph", n_iterations=2, directed=False)
M["clusters"] = int(adata.obs["leiden"].nunique())

# ---------- marker 基因 ----------
sc.tl.rank_genes_groups(adata, "leiden", method="wilcoxon")
top = {g: [str(x) for x in adata.uns["rank_genes_groups"]["names"][g][:8]]
       for g in adata.obs["leiden"].cat.categories}
M["top_markers"] = top

# ---------- 按经典 marker 标注细胞类型 ----------
SIG = {
    "CD4+ T": ["IL7R", "CCR7", "CD3D"], "CD8+ T": ["CD8A", "CD8B", "GZMK"],
    "NK": ["GNLY", "NKG7", "KLRD1"], "B": ["MS4A1", "CD79A", "CD79B"],
    "CD14+ Monocyte": ["CD14", "LYZ", "S100A9"], "FCGR3A+ Monocyte": ["FCGR3A", "MS4A7"],
    "Dendritic": ["FCER1A", "CST3"], "Platelet": ["PPBP", "PF4"],
}
label = {}
for cl, marks in top.items():
    best, hit = "Unassigned", 0
    for name, sig in SIG.items():
        n = len(set(sig) & set(marks))
        if n > hit: best, hit = name, n
    label[cl] = best if hit else "Unassigned"
adata.obs["cell_type"] = adata.obs["leiden"].map(label).astype("category")
M["cell_types"] = {k: int(v) for k, v in adata.obs["cell_type"].value_counts().items()}
M["unassigned_pct"] = round(100 * M["cell_types"].get("Unassigned", 0) / adata.n_obs, 1)

# ---------- 出图 ----------
sc.pl.umap(adata, color="leiden", legend_loc="on data", title="Leiden clusters",
           save="_clusters.png", show=False)
sc.pl.umap(adata, color="cell_type", title="Annotated cell types", save="_celltypes.png", show=False)
sc.pl.rank_genes_groups_dotplot(adata, n_genes=3, save="_markers.png", show=False)

with open(f"{OUT}/metrics.json", "w") as f:
    json.dump(M, f, indent=2, ensure_ascii=False)
print(json.dumps({k: v for k, v in M.items() if k != "top_markers"}, indent=2))
print("图与指标写入:", OUT)
