#!/usr/bin/env python3
"""Summarize R1 MEME sweep: episodic-selection sites per host x gene.

Tests PEER_REVIEW M2: the paper's Bhatt/MK estimator finds ~no directional selection in
birds. If site-level MEME ALSO finds ~no episodic sites in avian HA/NA (comparable to the
PB1 negative control), M2 is weakened (birds really are quiescent). If avian HA/NA show
episodic sites comparable to mammals, the paper's "no avian adaptation" is an artifact of
its persistence-dependent estimator.

Emits a markdown table to results/R1_SUMMARY.md and prints it.
"""
import csv, glob, os, collections

RES = os.path.join(os.path.dirname(__file__), "..", "results")
RES = os.path.abspath(RES)
HOSTS = ["na_avian","eurasian_avian","swine","human","equine","canine"]
HOST_LABEL = {"na_avian":"N.Am avian","eurasian_avian":"Eurasian avian","swine":"swine",
              "human":"human","equine":"equine","canine":"canine"}
GENES = ["ha","na","pb1"]
CLASS = {"na_avian":"avian","eurasian_avian":"avian","swine":"mammal","human":"mammal",
         "equine":"mammal","canine":"mammal"}

def summarize(path):
    n=var=p05=q05=0
    for r in csv.DictReader(open(path)):
        n+=1
        inv = str(r["is_invariable"]).strip().lower()=="true"
        if not inv: var+=1
        try:
            if float(r["p_value"])<=0.05: p05+=1
            if float(r["q_value"])<=0.05: q05+=1
        except (ValueError,KeyError): pass
    return n,var,p05,q05

rows=[]
for gene in GENES:
    for host in HOSTS:
        f=os.path.join(RES,f"meme_{gene}_{host}.csv")
        if not os.path.exists(f):
            rows.append((gene,host,None)); continue
        rows.append((gene,host,summarize(f)))

lines=[]
lines.append("# R1 — MEME episodic-selection sites per host × gene\n")
lines.append("`p05` = sites p≤0.05 · `q05` = FDR-significant sites q≤0.05 · "
             "`q05/100var` = FDR-sig per 100 variable codons (length-normalized).\n")
lines.append("| gene | host | class | codons | variable | p≤0.05 | q≤0.05 | q05/100var |")
lines.append("|------|------|-------|-------:|---------:|-------:|-------:|-----------:|")
for gene,host,s in rows:
    if s is None:
        lines.append(f"| {gene} | {HOST_LABEL[host]} | {CLASS[host]} | — | — | — | — | — |"); continue
    n,var,p05,q05=s
    rate = f"{100*q05/var:.2f}" if var else "—"
    lines.append(f"| {gene.upper()} | {HOST_LABEL[host]} | {CLASS[host]} | {n} | {var} | {p05} | {q05} | {rate} |")

# avian-vs-mammal roll-up for HA and NA (the M2-relevant genes)
lines.append("\n## M2 roll-up: avian vs mammal, q≤0.05 FDR-significant sites\n")
lines.append("| gene | avian hosts (q05 each) | mammal hosts (q05 each) |")
lines.append("|------|------------------------|-------------------------|")
by=collections.defaultdict(dict)
for gene,host,s in rows:
    if s: by[gene][host]=s[3]
for gene in ["ha","na","pb1"]:
    av=", ".join(f"{HOST_LABEL[h]}={by[gene].get(h,'—')}" for h in HOSTS if CLASS[h]=="avian")
    ma=", ".join(f"{HOST_LABEL[h]}={by[gene].get(h,'—')}" for h in HOSTS if CLASS[h]=="mammal")
    lines.append(f"| {gene.upper()} | {av} | {ma} |")

out="\n".join(lines)+"\n"
open(os.path.join(RES,"R1_SUMMARY.md"),"w").write(out)
print(out)
