#!/usr/bin/env python3
"""Prune each per-host HA tree to exactly the tips present in the (frame-2) HA alignment.

The GISAID per-host .nwk trees are SUPERSETS of the alignments (all alignment seqs present, plus
extra tips). Pruning a superset to the alignment tips preserves the authors' inferred topology and
branch lengths for the retained taxa — it does NOT re-estimate anything. Output feeds tree-based
`hyphaeon meme -t` (R1c), replacing the --no-tree/TN93 pass whose power was in question.

Pure-Python Newick prune (no deps): parse -> drop tips not in keep-set -> suppress unifurcations,
summing branch lengths through removed internal nodes so path lengths are preserved.
Writes <host>.pruned.nwk next to the tree; prints tip-count + concordance report.
"""
import sys, os, re, glob

def read_fasta_ids(fa):
    return {l[1:].strip() for l in open(fa) if l.startswith('>')}

# --- minimal Newick parser -> nested (name, length, [children]) ---
class Node:
    __slots__=('name','length','children')
    def __init__(s,name=None,length=None): s.name=name; s.length=length; s.children=[]

def parse_newick(txt):
    txt=txt.strip().rstrip(';')
    pos=0
    def parse_clade():
        nonlocal pos
        node=Node()
        if txt[pos]=='(':
            pos+=1
            while True:
                node.children.append(parse_clade())
                if txt[pos]==',': pos+=1; continue
                if txt[pos]==')': pos+=1; break
        # label
        m=re.match(r"([^,():;]+)?(?::(-?[0-9.eE+]+))?", txt[pos:])
        lab,ln=m.group(1),m.group(2)
        if lab is not None: node.name=lab.strip().strip("'\"")
        if ln is not None: node.length=float(ln)
        pos+=m.end()
        return node
    return parse_clade()

def prune(node, keep):
    # returns pruned node or None; collapses unifurcations summing lengths
    if not node.children:
        return node if (node.name in keep) else None
    kids=[c for c in (prune(c,keep) for c in node.children) if c is not None]
    if not kids: return None
    if len(kids)==1:
        c=kids[0]
        c.length=(c.length or 0.0)+(node.length or 0.0)
        return c
    node.children=kids
    return node

def to_newick(node):
    if not node.children:
        s=node.name or ''
    else:
        s='('+','.join(to_newick(c) for c in node.children)+')'+(node.name or '')
    if node.length is not None: s+=f":{node.length:.8g}"
    return s

def count_tips(node):
    return 1 if not node.children else sum(count_tips(c) for c in node.children)

def main(ha_dir):
    for fa in sorted(glob.glob(os.path.join(ha_dir,"*.inframe.fasta"))):
        host=os.path.basename(fa).replace('.inframe.fasta','')
        nwk=os.path.join(ha_dir,f"{host}.nwk")
        if not os.path.exists(nwk): print(f"{host}: NO TREE"); continue
        keep=read_fasta_ids(fa)
        root=parse_newick(open(nwk).read())
        before=count_tips(root)
        pr=prune(root,keep)
        after=count_tips(pr) if pr else 0
        out=os.path.join(ha_dir,f"{host}.pruned.nwk")
        open(out,'w').write(to_newick(pr)+';\n')
        missing=len(keep)-after
        print(f"{host:16s} aln={len(keep):5d} tree_before={before:5d} tree_after={after:5d} "
              f"aln_not_in_tree={missing:3d} -> {os.path.basename(out)}")

if __name__=="__main__":
    main(sys.argv[1] if len(sys.argv)>1 else ".")
