#!/home/ubuntu/.hermes/hermes-agent/venv/bin/python3
"""
THESEUS v1.0 — Toolkit for Homologous Editome & Syntactic Evolutionary Unified Screening

A standardized pipeline for PPR/DYW annotation, cis-regulatory classification,
and isoform-resolved ω estimation in non-model perennial monocots.

Usage:
    python3 theseus.py annotate --proteome proteins.faa --outdir results/
    python3 theseus.py cis-scan --promoters promoters.fa --species coconut
    python3 theseus.py omega --gff genes.gff3 --cds cds.fna --proteome proteins.faa --query nypa_dyw.faa

Modules:
    1. annotate — Dual-strategy DYW quantification (CxxC+HxE + Pfam)
    2. cis-scan — Multi-layer cis-regulatory classification
    3. omega — GFF-CDS-protein full-chain isoform validation + codeml ω

Dependencies: HMMER3, MAFFT, pal2nal, PAML (codeml), samtools
"""

import argparse, sys, os, re, gzip, subprocess, json
from collections import defaultdict

# ═══════════════════════════════════════════════════════════
# CONSTANTS
# ═══════════════════════════════════════════════════════════

DYW_MOTIF = re.compile(r'C[A-Z]{2}C.{20,150}H[A-Z]E')  # CxxC...HxE within 150aa

C4_MOTIFS = {
    'ACGTG_core':    r'ACGTG',
    'G-box':          r'CACGTG',
    'ABRE':           r'[ACT]ACGTG[GT][AC]',
    'MESP-1':         r'TTTTCT[AT]',
    'DOF':            r'AAAG',
    'GT-1':           r'G[AT]TTA[AT][AT]',
    'I-box':          r'GATAAG',
    'TATA':           r'TATA[AT]A[AT]',
    'CAAT':           r'CAAT',
}

CODON = {
    'TTT':'F','TTC':'F','TTA':'L','TTG':'L','TCT':'S','TCC':'S','TCA':'S','TCG':'S',
    'TAT':'Y','TAC':'Y','TGT':'C','TGC':'C','TGG':'W','CTT':'L','CTC':'L','CTA':'L','CTG':'L',
    'CCT':'P','CCC':'P','CCA':'P','CCG':'P','CAT':'H','CAC':'H','CAA':'Q','CAG':'Q',
    'CGT':'R','CGC':'R','CGA':'R','CGG':'R','ATT':'I','ATC':'I','ATA':'I','ATG':'M',
    'ACT':'T','ACC':'T','ACA':'T','ACG':'T','AAT':'N','AAC':'N','AAA':'K','AAG':'K',
    'AGT':'S','AGC':'S','AGA':'R','AGG':'R','GTT':'V','GTC':'V','GTA':'V','GTG':'V',
    'GCT':'A','GCC':'A','GCA':'A','GCG':'A','GAT':'D','GAC':'D','GAA':'E','GAG':'E',
    'GGT':'G','GGC':'G','GGA':'G','GGG':'G','TAA':'*','TAG':'*','TGA':'*',
}

# ═══════════════════════════════════════════════════════════
# FASTA I/O
# ═══════════════════════════════════════════════════════════

def read_fasta(path, gz=False):
    """Read FASTA, return {id: seq}."""
    idx = {}
    opener = gzip.open if gz else open
    cur_id, cur_seq = None, []
    with opener(path, 'rt') as f:
        for line in f:
            if line.startswith('>'):
                if cur_id:
                    idx[cur_id] = ''.join(cur_seq)
                cur_id = line[1:].strip().split()[0]
                cur_seq = []
            else:
                cur_seq.append(line.strip())
        if cur_id:
            idx[cur_id] = ''.join(cur_seq)
    return idx

def translate(dna):
    prot = []
    for i in range(0, len(dna)-2, 3):
        codon = dna[i:i+3].upper()
        aa = CODON.get(codon, 'X')
        if aa == '*':
            break
        prot.append(aa)
    return ''.join(prot)

# ═══════════════════════════════════════════════════════════
# MODULE 1: DYW Annotation
# ═══════════════════════════════════════════════════════════

def module_annotate(args):
    """Dual-strategy DYW quantification."""
    proteome = read_fasta(args.proteome, args.proteome.endswith('.gz'))
    
    ppr_ids = set()
    dyw_motif = set()
    
    # Motif-based scan (primary)
    for pid, seq in proteome.items():
        if DYW_MOTIF.search(seq):
            dyw_motif.add(pid)
    
    # Pfam DYW scan (control) — requires external HMM
    dyw_pfam = set()
    if args.pfam_hmm and os.path.exists(args.pfam_hmm):
        result = subprocess.run(['hmmsearch', '--noali', '-E', '10', '--cpu', '2',
                                args.pfam_hmm, args.proteome],
                               capture_output=True, text=True)
        for line in result.stdout.split('\n'):
            if line.startswith('>>'):
                dyw_pfam.add(line.split()[1])
    
    # PPR scan — requires external HMM
    if args.ppr_hmm and os.path.exists(args.ppr_hmm):
        result = subprocess.run(['hmmsearch', '--noali', '-E', '100', '--cpu', '2',
                                args.ppr_hmm, args.proteome],
                               capture_output=True, text=True)
        for line in result.stdout.split('\n'):
            if line.startswith('>>'):
                ppr_ids.add(line.split()[1])
    
    # Classification
    high_conf = dyw_motif & dyw_pfam & ppr_ids  # Both strategies + PPR+
    candidate = (dyw_motif | dyw_pfam) & ppr_ids - high_conf
    
    print(f"=== THESEUS Module 1: DYW Annotation ===")
    print(f"Total proteins:           {len(proteome)}")
    print(f"PPR (PF01535):            {len(ppr_ids)}")
    print(f"DYW motif (CxxC+HxE):     {len(dyw_motif)}")
    print(f"DYW Pfam (PF14432):       {len(dyw_pfam)}")
    print(f"High-confidence editing:   {len(high_conf)}")
    print(f"Candidate editing:         {len(candidate)}")
    print(f"DYW/PPR ratio (high-conf): {len(high_conf)/max(1,len(ppr_ids))*100:.1f}%")
    
    return {'ppr': len(ppr_ids), 'dyw_motif': len(dyw_motif), 
            'dyw_pfam': len(dyw_pfam), 'high_conf': len(high_conf),
            'candidate': len(candidate), 'proteome': len(proteome)}

# ═══════════════════════════════════════════════════════════
# MODULE 2: Cis-element Scanning
# ═══════════════════════════════════════════════════════════

def module_cis_scan(args):
    """Multi-layer cis-regulatory classification."""
    promoters = read_fasta(args.promoters)
    
    print(f"=== THESEUS Module 2: Cis-Regulatory Scan ===")
    print(f"Promoters: {len(promoters)}")
    print(f"{'Motif':<12s} {'Total':>6s} {'Per-gene':>8s} {'Layer':>12s}")
    print("-" * 42)
    
    results = {}
    for name, pat in C4_MOTIFS.items():
        total = sum(len(re.findall(pat, s, re.IGNORECASE)) for s in promoters.values())
        pg = total / max(1, len(promoters))
        
        # Layer classification
        if name in ('ACGTG_core', 'CAAT', 'TATA'):
            layer = 'basal'
        elif name == 'G-box':
            layer = 'intermediate'
        elif name in ('ABRE', 'MESP-1', 'DOF', 'GT-1', 'I-box'):
            layer = 'C4-specialized'
        else:
            layer = 'other'
        
        results[name] = {'total': total, 'per_gene': round(pg, 2), 'layer': layer}
        print(f"{name:<12s} {total:>6d} {pg:>8.2f} {layer:>12s}")
    
    # Check for C4 combinatorial module — SPATIAL CLUSTERING required
    # C4 functional syntax requires:
    #   1. MESP-1 within −200 to −100 bp regulatory window
    #   2. DOF, GT-1, ABRE within ≤80 bp of the MESP-1 site
    #   3. ≥3 tandem repeats of the clustered unit
    # Scattered isolated motifs do NOT constitute functional C4 syntax.
    
    mesp_pattern = re.compile('TTTTCT[AT]')
    dof_pattern  = re.compile('AAAG')
    gt1_pattern  = re.compile(r'G[AT]TTA[AT][AT]')
    abre_pattern = re.compile(r'[ACT]ACGTG[GT][AC]')
    
    clusters_found = 0
    for name, seq in promoters.items():
        for mesp_match in mesp_pattern.finditer(seq):
            pos = mesp_match.start()
            # Check if MESP-1 is in the −200 to −100 regulatory window
            # (Promoter sequences are 2kb, so −200 = position 1800, −100 = position 1900)
            window_start = max(0, len(seq) - 200)
            window_end   = max(0, len(seq) - 100)
            if window_start > 0 and window_end > window_start:
                if window_start <= pos <= window_end:
                    # Check for co-localized DOF, GT-1, ABRE within 80 bp
                    region_start = max(0, pos - 80)
                    region_end   = min(len(seq), pos + 80)
                    region = seq[region_start:region_end]
                    
                    has_dof_near = bool(dof_pattern.search(region))
                    has_gt1_near = bool(gt1_pattern.search(region))
                    has_abre_near = bool(abre_pattern.search(region))
                    
                    if has_dof_near and has_gt1_near and has_abre_near:
                        clusters_found += 1
    
    if clusters_found >= 3:
        print(f"\nC4 combinatorial module (spatial clustering, ≥3 tandem units): PRESENT ✅ ({clusters_found} clusters)")
    else:
        print(f"\nC4 combinatorial module: Individual motifs detected, no functional clustered syntax ❌")
        if clusters_found > 0:
            print(f"  ({clusters_found} candidate clusters found, <3 tandem repeats — insufficient for functional C4 syntax)")
    
    return results

# ═══════════════════════════════════════════════════════════
# MODULE 3: Omega Estimation (GFF-CDS-protein full-chain)
# ═══════════════════════════════════════════════════════════

def module_omega(args):
    """GFF-CDS-protein isoform validation + codeml ω."""
    print(f"=== THESEUS Module 3: Isoform-Resolved ω Estimation ===")
    
    # Load query proteins
    queries = read_fasta(args.query)
    proteome = read_fasta(args.proteome, args.proteome.endswith('.gz'))
    cds = read_fasta(args.cds, args.cds.endswith('.gz'))
    
    validated = 0
    total = len(queries)
    
    for qid, qseq in queries.items():
        print(f"\n  Query: {qid} ({len(qseq)}aa)")
        
        # 1. BLAST against proteome to find ortholog
        with open('/tmp/q.faa', 'w') as f:
            f.write(f'>{qid}\n{qseq}\n')
        
        subprocess.run(['makeblastdb', '-in', args.proteome, '-dbtype', 'prot',
                       '-out', '/tmp/target_db'], capture_output=True, timeout=10)
        
        result = subprocess.run(['blastp', '-query', '/tmp/q.faa', '-db', '/tmp/target_db',
                                '-outfmt', '6 sseqid', '-evalue', '1e-30', '-max_target_seqs', '1'],
                               capture_output=True, text=True, timeout=10)
        
        hit = result.stdout.strip().split('\n')[0] if result.stdout.strip() else ''
        
        if not hit:
            print(f"    ❌ No ortholog found in proteome")
            continue
        
        # 2. Verify CDS exists for the ortholog
        if hit not in proteome:
            print(f"    ❌ Ortholog {hit} not in proteome index")
            continue
        
        if hit not in cds:
            print(f"    ❌ Ortholog {hit} has no CDS entry")
            continue
        
        # 3. Translate CDS and verify >99% match
        cds_seq = cds[hit]
        cds_prot = translate(cds_seq)
        prot_seq = proteome[hit]
        
        min_len = min(len(cds_prot), len(prot_seq))
        matches = sum(1 for i in range(min_len) if cds_prot[i] == prot_seq[i])
        identity = matches / min_len * 100 if min_len > 0 else 0
        
        if identity < 99.0:
            print(f"    ❌ CDS↔protein mismatch: {identity:.1f}% ({matches}/{min_len})")
            continue
        
        print(f"    ✅ Full-chain validated: {hit} ({len(cds_seq)}nt, {identity:.1f}%)")
        
        # 4. pal2nal + codeml
        og = os.path.join(args.outdir, qid)
        os.makedirs(og, exist_ok=True)
        
        # Write protein alignment
        with open(f'{og}/prot.faa', 'w') as f:
            f.write(f'>query\n{qseq}\n')
            f.write(f'>target\n{prot_seq}\n')
        
        subprocess.run(['mafft', '--auto', f'{og}/prot.faa'],
                      stdout=open(f'{og}/prot_aln.faa', 'w'), stderr=subprocess.DEVNULL, timeout=30)
        
        # Clean CDS (remove stop codon)
        if cds_seq[-3:].upper() in ('TAG', 'TGA', 'TAA'):
            cds_seq = cds_seq[:-3]
        qcds = ''
        for k, v in cds.items():
            if qid in k:
                qcds = v
                break
        if qcds and qcds[-3:].upper() in ('TAG', 'TGA', 'TAA'):
            qcds = qcds[:-3]
        
        if qcds:
            ml = min(len(qcds), len(cds_seq))
            ml = (ml // 3) * 3
            with open(f'{og}/cds.fna', 'w') as f:
                f.write(f'>query\n{qcds[:ml]}\n')
                f.write(f'>target\n{cds_seq[:ml]}\n')
            
            result = subprocess.run(['pal2nal.pl', f'{og}/prot_aln.faa', f'{og}/cds.fna',
                                    '-output', 'paml'], capture_output=True, text=True, timeout=10)
            
            if result.returncode == 0 and 'ERROR' not in result.stderr:
                with open(f'{og}/codon.phy', 'w') as f:
                    f.write(result.stdout)
                
                with open(f'{og}/codeml.ctl', 'w') as f:
                    f.write(f'seqfile = codon.phy\noutfile = mlc.txt\n')
                    f.write('noisy = 0; verbose = 0; runmode = -2\n')
                    f.write('seqtype = 1; CodonFreq = 2; clock = 0\n')
                    f.write('model = 0; NSsites = 0; icode = 0\n')
                    f.write('fix_omega = 0; omega = 1; cleandata = 1\n')
                
                subprocess.run(['codeml', 'codeml.ctl'], cwd=og, capture_output=True, timeout=10)
                
                if os.path.exists(f'{og}/rst'):
                    with open(f'{og}/rst') as f:
                        for line in f:
                            if line.startswith(' YN:'):
                                parts = line.split()
                                if len(parts) >= 7:
                                    w = float(parts[3])
                                    print(f"    ω = {w:.4f}  ({ml}nt, {ml//3} codons)")
                                    break
                
                validated += 1
    
    print(f"\n  Validated: {validated}/{total} ortholog pairs")
    print(f"  THESEUS QC: {'PASS ✅' if validated > 0 else 'FAIL ❌ — no pairs survived isoform audit'}")
    
    return {'validated': validated, 'total': total}

# ═══════════════════════════════════════════════════════════
# CLI
# ═══════════════════════════════════════════════════════════

def main():
    parser = argparse.ArgumentParser(description='THESEUS v1.0 — Evolutionary Constraint Audit')
    sub = parser.add_subparsers(dest='module')
    
    # annotate
    p = sub.add_parser('annotate', help='Dual-strategy DYW annotation')
    p.add_argument('--proteome', required=True, help='Protein FASTA')
    p.add_argument('--ppr-hmm', help='Pfam PF01535 HMM')
    p.add_argument('--pfam-hmm', help='Pfam PF14432 HMM')
    p.add_argument('--outdir', default='.')
    
    # cis-scan
    p = sub.add_parser('cis-scan', help='Multi-layer cis-element scan')
    p.add_argument('--promoters', required=True, help='Promoter FASTA')
    p.add_argument('--species', default='unknown')
    
    # omega
    p = sub.add_parser('omega', help='Isoform-resolved ω estimation')
    p.add_argument('--query', required=True, help='Query DYW PPR FASTA')
    p.add_argument('--proteome', required=True, help='Target proteome FASTA')
    p.add_argument('--cds', required=True, help='Target CDS FASTA')
    p.add_argument('--outdir', default='.')
    
    args = parser.parse_args()
    
    if args.module == 'annotate':
        module_annotate(args)
    elif args.module == 'cis-scan':
        module_cis_scan(args)
    elif args.module == 'omega':
        module_omega(args)
    else:
        parser.print_help()

if __name__ == '__main__':
    main()
