#!/usr/bin/env python3
"""登革热四型保守表位 vs ADE 区域排除：覆盖率被吃掉多少。

判定规则写在 ../claims.json，跑之前定死。零模型是排除同等总长度的随机区域 1000 次。
"""
import json, collections
import numpy as np

KS = (9, 15)
FUSION_MOTIF = "DRGWGNGCGLFGK"      # E 蛋白融合环，跨黄病毒高度保守
FUSION_PAD = 6                       # 融合环两侧各多切 6 残基，覆盖构象表位边缘
N_PERM, SEED = 1000, 20260827


def kmers(seq, k):
    return {seq[i:i + k]: i for i in range(len(seq) - k + 1)}


def main():
    seqs = json.load(open("sequences.json"))
    names = sorted(seqs)
    rng = np.random.default_rng(SEED)
    print(f"[data] 血清型 {names}｜长度 {[seqs[n]['length'] for n in names]}")

    # ADE 区域：prM 全链 + E 融合环（含两侧 padding），坐标全部从注释/基序定位得到
    ade = {}
    for n in names:
        s = seqs[n]["sequence"]; spans = []
        for nm, (a, b) in seqs[n]["chains"].items():
            if nm.strip() in ("Protein prM", "Peptide pr", "Small envelope protein M"):
                spans.append((a - 1, b))
        i = s.find(FUSION_MOTIF)
        if i < 0:
            raise SystemExit(f"{n} 里找不到融合环基序，需人工确认")
        spans.append((max(0, i - FUSION_PAD), i + len(FUSION_MOTIF) + FUSION_PAD))
        # 合并重叠
        spans.sort(); merged = []
        for a, b in spans:
            if merged and a <= merged[-1][1]:
                merged[-1][1] = max(merged[-1][1], b)
            else:
                merged.append([a, b])
        ade[n] = merged
        tot = sum(b - a for a, b in merged)
        print(f"[ADE] {n} 排除区共 {tot} 残基（占 {100*tot/len(s):.1f}%）；融合环位于 {i+1}")

    results = {}
    for k in KS:
        maps = {n: kmers(seqs[n]["sequence"], k) for n in names}
        conserved = set(maps[names[0]])
        for n in names[1:]:
            conserved &= set(maps[n])
        base = len(conserved)

        def lost_by(spans_by_name):
            """落在被排除区域里的保守 k-mer 数（任一血清型命中即算被排除）。"""
            lost = set()
            for n in names:
                for km in conserved:
                    i = maps[n][km]
                    for a, b in spans_by_name[n]:
                        if i < b and i + k > a:
                            lost.add(km); break
            return len(lost)

        lost_ade = lost_by(ade)
        # 零模型：同样总长度、同样片段数的随机区域
        null = []
        for _ in range(N_PERM):
            rnd = {}
            for n in names:
                L = len(seqs[n]["sequence"]); spans = []
                for a, b in ade[n]:
                    w = b - a
                    st = int(rng.integers(0, max(1, L - w)))
                    spans.append([st, st + w])
                rnd[n] = spans
            null.append(lost_by(rnd))
        null = np.array(null)
        p95 = float(np.percentile(null, 95))
        emp_p = float((null >= lost_ade).mean())
        # H1：随机序列下四型一致 k-mer 的期望（按各型氨基酸组成独立抽样）
        exp_random = 0.0
        aa = collections.Counter("".join(seqs[n]["sequence"] for n in names))
        tot = sum(aa.values()); freqs = np.array([v / tot for v in aa.values()])
        p_same = float((freqs ** 4).sum())          # 四条序列同一位置同氨基酸的概率
        exp_random = (len(seqs[names[0]]["sequence"]) - k + 1) * (p_same ** k)

        results[k] = {"n_conserved": base, "lost_ade": lost_ade,
                      "loss_frac": lost_ade / base if base else 0.0,
                      "null_mean": float(null.mean()), "null_p95": p95,
                      "null_max": int(null.max()), "empirical_p": emp_p,
                      "expected_by_chance": exp_random,
                      "H1_passes": bool(base > max(exp_random * 10, 10)),
                      "H2_passes": bool(lost_ade > p95),
                      "remaining": base - lost_ade}
        r = results[k]
        print(f"\n[k={k}] 四型完全一致的保守表位: {base} 个（随机期望 {exp_random:.2g}）")
        print(f"[k={k}] 排除 ADE 区后损失 {lost_ade} 个（{100*r['loss_frac']:.1f}%），剩余 {r['remaining']} 个")
        print(f"[k={k}] 零模型（随机区域）损失 均值 {null.mean():.1f}｜95分位 {p95:.1f}｜最大 {null.max()}")
        print(f"[k={k}] H1 保守性真实 → {'✅' if r['H1_passes'] else '❌'}｜"
              f"H2 ADE 区损失显著更大 → {'✅' if r['H2_passes'] else '❌'}（经验 p={emp_p:.3f}）")

    json.dump({"seed": SEED, "n_perm": N_PERM, "serotypes": names,
               "fusion_motif": FUSION_MOTIF, "fusion_pad": FUSION_PAD,
               "ade_spans": ade, "by_k": results},
              open("results.json", "w"), ensure_ascii=False, indent=1)


if __name__ == "__main__":
    main()
