# Reproduction code for 03_mcr_metabolism
# Environment: archaea-bio

import urllib.request
import urllib.parse
import json
import time
import subprocess
import numpy as np
import pandas as pd
from io import StringIO
from collections import Counter
from Bio import SeqIO, Phylo

# Load feature matrix for taxonomy
fm = pd.read_csv("/Users/ajayvellanki/.claude-science/orgs/a85da623-e9bf-4a84-bae1-c872a64348e0/artifacts/proj_aa6c67020290/bc0e3352-d2ec-4e97-91a6-ae3a4c4a7c27/vd4aee24f_feature_matrix.csv")
tax = fm[['accession','order','family','genus','organism_name']].drop_duplicates('accession')

# Load detected mcrA sequences
det = list(SeqIO.parse("/Users/ajayvellanki/.claude-science/orgs/a85da623-e9bf-4a84-bae1-c872a64348e0/artifacts/proj_aa6c67020290/8a0141ef-e4e7-44c7-92fd-3379036a4f90/v0ef8ef2f_mcr_detected.faa", 'fasta'))
det_full = [r for r in det if len(r.seq) >= 400]

# Fetch reference sequences
def uniprot_search(query, size=3):
    base = "https://rest.uniprot.org/uniprotkb/search?"
    params = {'query': query, 'format': 'fasta', 'size': size}
    url = base + urllib.parse.urlencode(params)
    try:
        return urllib.request.urlopen(url, timeout=40).read().decode()
    except Exception as e:
        print("  err", repr(e)[:80])
        return ""

def ncbi_esearch(term, db='protein', retmax=5):
    url = "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esearch.fcgi?" + urllib.parse.urlencode(
        {'db': db, 'term': term, 'retmax': retmax, 'retmode': 'json'})
    try:
        return json.loads(urllib.request.urlopen(url, timeout=30).read())['esearchresult']['idlist']
    except Exception as e:
        print("esearch err", repr(e)[:80])
        return []

def ncbi_efetch(ids, db='protein'):
    if not ids:
        return ""
    url = "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/efetch.fcgi?" + urllib.parse.urlencode(
        {'db': db, 'id': ','.join(ids), 'rettype': 'fasta', 'retmode': 'text'})
    try:
        return urllib.request.urlopen(url, timeout=40).read().decode()
    except Exception as e:
        print("efetch err", repr(e)[:80])
        return ""

seen = set()
valid_refs = []

def add_fasta(fa_text, clade):
    for r in SeqIO.parse(StringIO(fa_text), 'fasta'):
        uid = r.id.split('|')[1] if '|' in r.id else r.id
        if uid in seen:
            continue
        seen.add(uid)
        r.id = f"REF_{clade}_{uid}"
        r.description = ""
        valid_refs.append(r)

# Fetch canonical methanogen references
canon_queries = [
    ('canonical_methanogen', 'gene:mcrA AND organism_name:"Methanocaldococcus jannaschii"'),
    ('canonical_methanogen', 'gene:mcrA AND organism_name:"Methanopyrus kandleri"'),
    ('canonical_methanogen', 'gene:mcrA AND organism_name:"Methanomassiliicoccus"'),
]
broad_queries = [
    ('alkane_acr', 'gene:mcrA AND organism_name:"Candidatus Syntrophoarchaeum"'),
    ('alkane_acr', 'organism_name:"Candidatus Argoarchaeum ethanivorans"'),
    ('alkane_acr', 'organism_name:"Candidatus Methanoliparum"'),
    ('alkane_acr', 'taxonomy_name:"Syntrophoarchaeia"'),
    ('anme_methanotroph', 'gene:mcrA AND organism_name:"Candidatus Methanophagales"'),
    ('anme_methanotroph', 'taxonomy_name:"ANME-1" AND gene:mcrA'),
    ('anme_methanotroph', 'organism_name:"Candidatus Methanoperedens nitroreducens" AND gene:mcrA'),
]

for clade, q in canon_queries + broad_queries:
    fa = uniprot_search(q, size=3)
    if fa.count('>'):
        add_fasta(fa, clade)
    time.sleep(0.4)

# Fetch from NCBI
ncbi_terms = [
    ('alkane_acr', 'Syntrophoarchaeum methyl-coenzyme M reductase'),
    ('alkane_acr', 'Methanoliparum methyl-coenzyme M reductase alpha'),
    ('alkane_acr', 'Argoarchaeum ethanivorans methyl-coenzyme M reductase'),
    ('anme_methanotroph', 'Methanoperedens methyl-coenzyme M reductase alpha'),
]
for clade, t in ncbi_terms:
    ids = ncbi_esearch(t, retmax=6)
    fa = ncbi_efetch(ids[:6]) if ids else ""
    for r in SeqIO.parse(StringIO(fa), 'fasta'):
        if len(r.seq) >= 420 and r.id not in seen:
            seen.add(r.id)
            r.id = f"REF_{clade}_{r.id.replace('.','_')}"
            r.description = ""
            valid_refs.append(r)
    time.sleep(0.5)

# Known accessions
fetch_final = {
    'alkane_acr': ['RJS73394.1', 'BDC35338.1', 'WP_229234681.1'],
    'anme_methanotroph': ['WP_097300250.1', 'WP_096206223.1'],
}
for clade, accs in fetch_final.items():
    fa = ncbi_efetch(accs)
    for r in SeqIO.parse(StringIO(fa), 'fasta'):
        if len(r.seq) >= 500 and r.id not in seen:
            seen.add(r.id)
            r.id = f"REF_{clade}_{r.id.replace('.','_').replace('|','_')}"
            r.description = ""
            valid_refs.append(r)
    time.sleep(0.4)

# Additional canonical diversity anchors
canon_more = ['WP_048196255.1', 'WP_013295740.1', 'WP_048061141.1']
fa = ncbi_efetch(canon_more)
for r in SeqIO.parse(StringIO(fa), 'fasta'):
    if len(r.seq) >= 500 and r.id not in seen:
        seen.add(r.id)
        r.id = f"REF_canonical_methanogen_{r.id.replace('.','_')}"
        r.description = ""
        valid_refs.append(r)

SeqIO.write(valid_refs, 'data/mcr_references.faa', 'fasta')

# Build tree input
allseqs = det_full + valid_refs
SeqIO.write(allseqs, 'data/mcr_tree_input.faa', 'fasta')

# Align
r = subprocess.run(['mafft', '--auto', '--thread', '6', '--quiet', 'data/mcr_tree_input.faa'],
                   capture_output=True, text=True)
open('data/mcr_aln.faa', 'w').write(r.stdout)

# Build tree
with open('data/mcr_reference_tree.nwk', 'w') as tree_out, open('data/fasttree.log', 'w') as log_out:
    subprocess.run(['FastTree', '-wag', '-gamma', '-quiet', 'data/mcr_aln.faa'],
                   stdout=tree_out, stderr=log_out)

# Load and root tree
tree = Phylo.read('data/mcr_reference_tree.nwk', 'newick')
tree.root_at_midpoint()

# Reference tip -> clade mapping
ref_clade = {}
for r in valid_refs:
    ref_clade[r.id] = '_'.join(r.id.split('_')[1:3])

terms = {t.name: t for t in tree.get_terminals()}
ref_tips = [n for n in terms if n in ref_clade]
det_tips = [n for n in terms if n not in ref_clade]

# Build clade reference objects
clade_refs = {}
for rn, cl in ref_clade.items():
    clade_refs.setdefault(cl, []).append(terms[rn])

# Compute distances from each detected tip to all clade refs
detrows = []
for i, dt in enumerate(det_tips):
    dobj = terms[dt]
    dists = {rn: tree.distance(dobj, terms[rn]) for rn in ref_tips}
    clade_min = {cl: min([dists[r.name] for r in objs]) for cl, objs in clade_refs.items()}
    best = min(clade_min, key=clade_min.get)
    detrows.append({
        'tip': dt, 'best_clade': best,
        'd_canonical': round(clade_min.get('canonical_methanogen', 999), 4),
        'd_anme': round(clade_min.get('anme_methanotroph', 999), 4),
        'd_alkane': round(clade_min.get('alkane_acr', 999), 4)
    })
    if i % 300 == 0:
        print(" processed", i)

dd = pd.DataFrame(detrows)
dd['protein_id'] = dd['tip'].str.split('|').str[0]
dd['accession'] = dd['tip'].str.split('|').str[1]

# Known alkane/ANME genera
alkane_genera = {'Candidatus Alkanophaga', 'Candidatus Methanoliparum', 'Candidatus Syntrophoarchaeum',
                 'Candidatus Argoarchaeum', 'Candidatus Methanolliviera', 'Candidatus Alkanophagales'}
anme_genera = {'Candidatus Methanoperedens', 'Candidatus Methanogaster', 'Candidatus Methanovorans'}
anme_orders = {'Candidatus Methanophagales', 'Methanophagales'}

dd = dd.merge(tax, on='accession', how='left')

def final_call(row):
    if row['best_clade'] == 'alkane_acr' and row['d_alkane'] < row['d_canonical']:
        return 'alkane_ACR'
    if row['genus'] in alkane_genera:
        return 'alkane_ACR'
    if row['genus'] in anme_genera or row['order'] in anme_orders:
        return 'ANME_methanotroph'
    return 'canonical_methanogenesis'

dd['metabolism'] = dd.apply(final_call, axis=1)

# Final classification table
mcr_final = dd[['protein_id', 'accession', 'organism_name', 'order', 'family', 'genus',
                'metabolism', 'best_clade', 'd_canonical', 'd_anme', 'd_alkane']].copy()
mcr_final.to_csv('mcr_classification.csv', index=False)
print("Saved mcr_classification.csv; rows:", len(mcr_final), "genomes:", mcr_final['accession'].nunique())
print("\nMetabolism calls:", mcr_final['metabolism'].value_counts().to_dict())