#!/usr/bin/env python3
"""校正版：扣掉 MeSH 词表版本与索引完整度两个混杂，再看还剩多少真变化。

两处校正：
  1) 只保留 2018 年前就在语料里出现过的词 —— 排除 MeSH 新增词条
  2) 用「该词占本篇 MeSH 总数的份额」而不是「出现率」—— 抵消索引完整度漂移
另加一条对照：把人口学检索标签单独拎出来当"内部对照"。若校正有效，
这些与研究内容无关的标签在校正后应当不再显著。
"""
import json, collections
import numpy as np
from scipy.stats import mannwhitneyu

CHECK_TAGS = {"Male", "Female", "Humans", "Adult", "Aged", "Middle Aged", "Young Adult",
              "Adolescent", "Aged, 80 and over", "Child", "Animals", "Retrospective Studies",
              "Prospective Studies", "Treatment Outcome", "Survival Rate", "Follow-Up Studies"}


def bh(p):
    p = np.asarray(p, float); n = len(p); o = np.argsort(p); adj = np.empty(n); prev = 1.0
    for r, i in enumerate(o[::-1]):
        prev = min(prev, p[i] * n / (n - r)); adj[i] = prev
    return adj


recs = [r for r in json.load(open("corpus.json"))["records"] if 2015 <= r["year"] <= 2026]
first = {}
for r in sorted(recs, key=lambda x: x["year"]):
    for m in r["mesh"]:
        first.setdefault(m, r["year"])
cnt = collections.Counter(m for r in recs for m in r["mesh"])
terms = [t for t, c in cnt.items() if c >= 20 and first[t] < 2018]
print(f"  校正后参与检验的词: {len(terms)} 个（原 115 个，剔除 MeSH 新增词条）")

late = np.array([r["year"] >= 2020 for r in recs])
share = {t: np.array([(1.0 / len(r["mesh"])) if t in r["mesh"] else 0.0 for r in recs])
         for t in terms}
ps = []
for t in terms:
    a, b = share[t][late], share[t][~late]
    ps.append(mannwhitneyu(a, b, alternative="two-sided")[1])
qs = bh(ps)
sig = [(t, float(q), float(share[t][~late].mean()), float(share[t][late].mean()))
       for t, q in zip(terms, qs) if q < 0.05]
sig.sort(key=lambda x: -(x[3] - x[2]))

tag_sig = [s for s in sig if s[0] in CHECK_TAGS]
content_sig = [s for s in sig if s[0] not in CHECK_TAGS]
print(f"  校正后显著: {len(sig)} 个｜其中内容词 {len(content_sig)}｜检索标签 {len(tag_sig)}")
print(f"  内部对照：检索标签{'仍有 %d 个显著 ⚠️ 校正不彻底' % len(tag_sig) if tag_sig else '全部不再显著 ✅ 校正有效'}")
print("\n  校正后真正上升的内容词：")
for t, q, e, l in content_sig[:8]:
    if l > e: print(f"    {t[:44]:<46} 份额 {e:.4f} → {l:.4f}  (FDR {q:.1e})")
print("  校正后真正下降的内容词：")
for t, q, e, l in sorted(content_sig, key=lambda x: x[3] - x[2])[:6]:
    if l < e: print(f"    {t[:44]:<46} 份额 {e:.4f} → {l:.4f}  (FDR {q:.1e})")

json.dump({"n_terms_after_filter": len(terms), "n_sig": len(sig),
           "n_content_sig": len(content_sig), "n_tag_sig": len(tag_sig),
           "correction_worked": len(tag_sig) == 0,
           "rising": [{"term": t, "early_share": e, "late_share": l, "fdr": q}
                      for t, q, e, l in content_sig if l > e][:12],
           "falling": [{"term": t, "early_share": e, "late_share": l, "fdr": q}
                       for t, q, e, l in sorted(content_sig, key=lambda x: x[3] - x[2]) if l < e][:12]},
          open("corrected.json", "w"), ensure_ascii=False, indent=1)
