#!/usr/bin/env python

# python libraries
import os, sys, optparse

# dlcpar libraries
import dlcpar

# rasmus, compbio libraries
from rasmus import util, treelib, tablelib
from compbio import phylo, phyloDLC

#=============================================================================
# parser

def parse_args():
    parser = optparse.OptionParser()
    parser.add_option("-s", "--stree", dest="stree",
                      metavar="<species tree>",
                      help="species tree (newick format)")
    parser.add_option("-S", "--smap", dest="smap",
                      metavar="<gene2species map>",
                      help="mapping of gene names to species names")
    parser.add_option("-T", "--treeext", dest="treeext",
                      metavar="<coalescent tree file extension>",
                      default=".coal.tree",
                      help="tree file extension (default: \".coal.tree\")")
    parser.add_option("-R", "--reconext", dest="reconext",
                      metavar="<coalescent recon file extension>",
                      default=".coal.recon",
                      help="tree file extension (default: \".coal.recon\")")
    parser.add_option("--by-fam", dest="by_fam", action="store_true")
    parser.add_option("--use-famid", dest="use_famid", action="store_true")
    parser.add_option("--explicit", dest="explicit",
                      action="store_true", default=False,
                      help="set to ignore extra lineages at implied speciation nodes")
    parser.add_option("--use-locus-recon", dest="use_locus_recon",
                      action="store_true", default=False,
                      help="if set, use locus recon rather than MPR")

    options, args = parser.parse_args()

    if (not options.stree) or (not options.smap):
        parser.error("-s/--stree and -S/--smap required")

    return options, args

#=============================================================================
# main

def count_all_events(options, args, exts):

    stree = treelib.read_tree(options.stree)
    gene2species = phylo.read_gene2species(options.smap)

    treefiles = map(lambda line: line.rstrip(), util.read_strings(sys.stdin))
    coal_trees = []
    extras = []

    for treefile in treefiles:
        prefix = util.replace_ext(treefile, options.treeext, "")
        coal_tree, extra = phyloDLC.read_dlcoal_recon(prefix, stree, exts)
        coal_trees.append(coal_tree)
        extras.append(extra)

    etree = phyloDLC.count_dup_loss_coal_trees(coal_trees, extras, stree, gene2species,
                                               implied=not options.explicit,
                                               locus_mpr=not options.use_locus_recon)

    # make table
    headers = ["genes", "dup", "loss", "coal", "appear"]
    ptable = treelib.tree2parent_table(etree, headers)

    # sort by post order
    lookup = util.list2lookup(x.name for x in stree.postorder())
    ptable.sort(key=lambda x: lookup[x[0]])

    ptable = [[str(row[0]), str(row[1]), float(row[2])] + row[3:]
              for row in ptable]

    tab = tablelib.Table(ptable,
                         headers=["nodeid", "parentid", "dist"] + headers)
    tab.write()

    return 0


def count_by_fam(options, args, exts):

    stree = treelib.read_tree(options.stree)
    gene2species = phylo.read_gene2species(options.smap)

    treefiles = map(lambda line: line.rstrip(), util.read_strings(sys.stdin))

    # write header
    lookup = util.list2lookup(x.name for x in stree.postorder())
    headers = ["genes", "dup", "loss", "coal", "appear"]
    print "\t".join(["famid", "nodeid", "parentid", "dist"] + headers)

    for treefile in treefiles:
        if options.use_famid:
            famid = os.path.basename(os.path.dirname(treefile))
        else:
            famid = treefile

        # read files and events
        prefix = util.replace_ext(treefile, options.treeext, "")
        coal_tree, extra = phyloDLC.read_dlcoal_recon(prefix, stree, exts)

        etree = phyloDLC.count_dup_loss_coal_trees([coal_tree], [extra], stree, gene2species,
                                                   implied=not options.explicit,
                                                   locus_mpr=not options.use_locus_recon)
        ptable = treelib.tree2parent_table(etree, headers)

        # sort by post order
        ptable.sort(key=lambda x: lookup[x[0]])

        # write table
        for row in ptable:
            print "\t".join(map(str, [famid] + row))

    return 0


def main():

    options, args = parse_args()

    exts = {"coal_tree": options.treeext,
            "coal_recon": options.reconext,
            "locus_tree": ".locus.tree",
            "locus_recon": ".locus.recon",
            "daughters": ".daughters"}

    if not options.by_fam:
        count_all_events(options, args, exts)

    else:
        count_by_fam(options, args, exts)


if __name__ == "__main__":
    sys.exit(main())
