#!/usr/bin/env python
"""公开数据的真实 bulk RNA-seq 差异表达样例(GEO: GSE60450,小鼠乳腺)。

样本分组不靠记忆:从 GEO 的 series matrix 里读官方 sample_title,再按标题解析细胞类型与生理状态。
比较:luminal lactating vs luminal virgin —— 一个生物学上极明确的对照,
如果流程正确,差异基因顶部应该出现乳蛋白基因(Csn2/Wap/Glycam1 等),这就是结果对不对的自检。
"""
import gzip, io, json, os, re, urllib.request, warnings
warnings.filterwarnings("ignore")
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd

HERE = os.path.dirname(os.path.abspath(__file__))
OUT = os.path.join(HERE, "out_bulk")
os.makedirs(OUT, exist_ok=True)
M = {}

# ---------- 1. 官方样本标题(不靠记忆) ----------
url = "https://ftp.ncbi.nlm.nih.gov/geo/series/GSE60nnn/GSE60450/matrix/GSE60450_series_matrix.txt.gz"
raw = gzip.decompress(urllib.request.urlopen(url, timeout=120).read()).decode("utf-8", "ignore")
titles = [t.strip('"') for t in re.search(r'^!Sample_title\t(.*)$', raw, re.M).group(1).split('\t')]
geo_ids = [t.strip('"') for t in re.search(r'^!Sample_geo_accession\t(.*)$', raw, re.M).group(1).split('\t')]
# 样本短名在 !Sample_description 的第一行("Sample name: MCL1-LA"),分组在 !Sample_source_name_ch1
desc = [t.strip('"') for t in re.findall(r'^!Sample_description\t(.*)$', raw, re.M)[0].split('\t')]
src = [t.strip('"') for t in re.search(r'^!Sample_source_name_ch1\t(.*)$', raw, re.M).group(1).split('\t')]
shorts = [re.sub(r'^Sample name:\s*', '', d).strip().upper() for d in desc]
M["geo_sample_titles"] = dict(zip(geo_ids, titles))

def parse(source):
    t = source.lower()
    ct = "luminal" if "luminal" in t else ("basal" if "basal" in t else "?")
    st = ("lactate" if "lactation" in t else "pregnant" if "pregnan" in t else "virgin" if "virgin" in t else "?")
    return ct, st

short = {sn: parse(sc) for sn, sc in zip(shorts, src)}

# ---------- 2. 计数矩阵 ----------
cnt = pd.read_csv(os.path.join(HERE, "bulk_counts.txt.gz"), sep="\t", index_col=0)
lengths = cnt.pop("Length")
cnt.columns = [c.split("_")[0].upper() for c in cnt.columns]
M["genes"], M["samples"] = int(cnt.shape[0]), int(cnt.shape[1])

meta = pd.DataFrame(
    [{"sample": s, "cell_type": short.get(s, ("?", "?"))[0], "status": short.get(s, ("?", "?"))[1]}
     for s in cnt.columns]).set_index("sample")
M["design_table"] = meta.reset_index().to_dict("records")

sel = meta[(meta.cell_type == "luminal") & (meta.status.isin(["lactate", "virgin"]))]
assert len(sel) >= 4, f"分组解析失败: {meta.to_dict()}"
counts = cnt[sel.index].T                      # pydeseq2 要 样本×基因
counts = counts.loc[:, counts.sum(axis=0) >= 10]
M["compared"] = {"group_a": "luminal lactating", "group_b": "luminal virgin",
                 "n_a": int((sel.status == "lactate").sum()), "n_b": int((sel.status == "virgin").sum()),
                 "genes_tested": int(counts.shape[1])}

# ---------- 3. DESeq2 ----------
from pydeseq2.dds import DeseqDataSet
from pydeseq2.ds import DeseqStats
clin = pd.DataFrame({"condition": ["lactating" if s == "lactate" else "virgin" for s in sel.status]},
                    index=sel.index)
dds = DeseqDataSet(counts=counts, metadata=clin, design="~condition", refit_cooks=True, quiet=True)
dds.deseq2()
st = DeseqStats(dds, contrast=["condition", "lactating", "virgin"], quiet=True)
st.summary()
res = st.results_df.dropna(subset=["padj"]).copy()

M["deg_up"] = int(((res.padj < 0.05) & (res.log2FoldChange > 1)).sum())
M["deg_down"] = int(((res.padj < 0.05) & (res.log2FoldChange < -1)).sum())

# Entrez ID → 基因名(NCBI eutils,官方接口)
top = res[(res.padj < 0.05)].sort_values("log2FoldChange", ascending=False).head(15)
ids = ",".join(str(i) for i in top.index)
names = {}
try:
    q = urllib.request.urlopen(
        "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esummary.fcgi?db=gene&id=" + ids + "&retmode=json",
        timeout=45).read()
    r = json.loads(q).get("result", {})
    for uid in r.get("uids", []):
        names[str(uid)] = r[uid].get("name", str(uid))
except Exception as e:
    print("gene symbol lookup failed:", e)
M["top_up"] = [{"gene": names.get(str(i), str(i)), "log2FC": round(float(r.log2FoldChange), 2),
                "padj": float(f"{r.padj:.3g}")} for i, r in top.iterrows()]
milk = {"Csn2", "Csn1s2a", "Csn1s2b", "Wap", "Glycam1", "Csn3", "Lalba", "Muc15", "Cel", "Csn1s1",
        "Csn1s2", "Mfge8", "Btn1a1", "Xdh", "Lao1", "Spp1"}
M["sanity_check_milk_genes_in_top15"] = sorted(milk & {x["gene"] for x in M["top_up"]})

# ---------- 4. 图 ----------
res["neglog10padj"] = -np.log10(res.padj.clip(lower=1e-300))
sig = (res.padj < 0.05) & (res.log2FoldChange.abs() > 1)
plt.figure(figsize=(5.6, 4.6))
plt.scatter(res.log2FoldChange[~sig], res.neglog10padj[~sig], s=4, c="#b8c2d4", alpha=.5, edgecolors="none")
plt.scatter(res.log2FoldChange[sig], res.neglog10padj[sig], s=5, c="#dc2626", alpha=.7, edgecolors="none")
plt.axvline(0, c="#94a3b8", lw=.6); plt.axhline(-np.log10(0.05), c="#94a3b8", lw=.6, ls="--")
plt.xlabel("log2 fold change (lactating vs virgin)"); plt.ylabel("-log10 adjusted p")
plt.title("Differential expression", fontsize=11)
plt.tight_layout(); plt.savefig(f"{OUT}/volcano.png", dpi=140); plt.close()

norm = np.log2(dds.layers["normed_counts"] + 1)
v = norm[:, np.argsort(norm.var(axis=0))[::-1][:500]]
v = (v - v.mean(0)) / (v.std(0) + 1e-9)
u, s_, _ = np.linalg.svd(v - v.mean(0), full_matrices=False)
pc = u[:, :2] * s_[:2]
plt.figure(figsize=(5, 4.2))
for cond, col in [("lactating", "#dc2626"), ("virgin", "#0a6ed1")]:
    m = (clin.condition == cond).values
    plt.scatter(pc[m, 0], pc[m, 1], s=70, c=col, label=cond, edgecolors="white")
plt.xlabel("PC1"); plt.ylabel("PC2"); plt.legend(); plt.title("Sample clustering (top 500 variable genes)", fontsize=11)
plt.tight_layout(); plt.savefig(f"{OUT}/pca.png", dpi=140); plt.close()

with open(f"{OUT}/metrics.json", "w") as f:
    json.dump(M, f, indent=2)
print(json.dumps({k: M[k] for k in ("genes", "samples", "compared", "deg_up", "deg_down",
                                     "sanity_check_milk_genes_in_top15")}, indent=2))
print("top up:", [x["gene"] for x in M["top_up"][:8]])
