#!/usr/bin/env python

# python libraries
import os, sys, optparse
import time
import random
import gzip

# geometry libraries
from shapely import geometry

# dlcpar libraries
import dlcpar
from dlcpar import common
from dlcpar import reconlib
import dlcpar.reconscape

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

# numpy libraries
import numpy as np

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

VERSION = dlcpar.PROGRAM_VERSION_TEXT

def parse_args():
    """parse input arguments"""

    # callback to parse cost ranges
    def parse_range(option, opt_str, value, parser):
        try:
            if "-" not in value:
                a = float(value)
                b = 1/a
                if a < b:
                    low, high = a, b
                else:
                    low, high = b, a
            else:
                low, high = map(float, value.split("-"))
            assert (low > 0) and (high > 0) and (low < high)
        except:
            raise optparse.OptionValueError("invalid %s: %s" % (option, value))
        setattr(parser.values, option.dest, (low, high))

    parser = optparse.OptionParser(
        usage = "usage: %prog [options] <gene tree> ...",

        version = "%prog " + VERSION,

        description =
        "%prog is a program for finding landscapes of optimal (most parsimonious) " +
        "gene tree-species tree reconciliations.  " +
        "See http://www.cs.hmc.edu/~yjw/software/dlcpar for details.",

        epilog =
        "Written by Yi-Chieh Wu (yjw@cs.hmc.edu), Harvey Mudd College. " +
        "(c) 2014-2019. Released under the terms of the GNU General Public License.")

    grp_io = optparse.OptionGroup(parser, "Input/Output")
    grp_io.add_option("-s", "--stree", dest="stree",
                      metavar="<species tree>",
                      help="species tree file in newick format")
    grp_io.add_option("-S", "--smap", dest="smap",
                      metavar="<species map>",
                      help="gene to species map")
    grp_io.add_option("--lmap", dest="lmap",
                      metavar="<locus map>",
                      help="gene to locus map (species-specific)")
    grp_io.add_option("--events", dest="events",
                      metavar="<output events>",
                      choices=[None,"I","U"], default=None,
                      help="set to output (I)ntersection or (U)nion of events (default: None)")
    grp_io.add_option("--draw_regions", dest="draw_regions",
                      metavar="<draw regions>",
                      default=False,
                      action="store_true",
                      help="set to draw regions to screen")
    parser.add_option_group(grp_io)

    grp_ext = optparse.OptionGroup(parser, "File Extensions")
    grp_ext.add_option("-I","--inputext", dest="inext",
                       metavar="<input file extension>",
                       default="",
                       help="input file extension (default: \"\")")
    grp_ext.add_option("-O", "--outputext", dest="outext",
                       metavar="<output file extension>",
                       default=".dlcscape",
                       help="output file extension (default: \".dlcscape\")")
    parser.add_option_group(grp_ext)

    grp_costs = optparse.OptionGroup(parser, "Costs")
    grp_costs.add_option("-D", "--duprange", dest="duprange",
                         metavar="<dup range>",
                         default=dlcpar.reconscape.DEFAULT_RANGE, type="string",
                         action="callback", callback=parse_range,
                         help="duplication range (default: %s)" % dlcpar.reconscape.DEFAULT_RANGE_STR)
    grp_costs.add_option("-L", "--lossrange", dest="lossrange",
                         metavar="<loss range>",
                         default=dlcpar.reconscape.DEFAULT_RANGE, type="string",
                         action="callback", callback=parse_range,
                         help="loss range (default: %s)" % dlcpar.reconscape.DEFAULT_RANGE_STR)
    parser.set_defaults(explicit=False)
    parser.add_option_group(grp_costs)

    grp_heur = optparse.OptionGroup(parser, "Heuristics")
    parser.set_defaults(delay=False)
    grp_heur.add_option("--max_loci", dest="max_loci",
                        metavar="<max # of loci>",
                        default=-1, type="int",
                        help="maximum # of co-existing loci (in each ancestral species), " +\
                             "set to -1 for no limit (default: -1)")
    grp_heur.add_option("--max_dups", dest="max_dups",
                        metavar="<max # of dups>",
                        default=4, type="int",
                        help="maximum # of duplications (in each ancestral species), " +\
                             "set to -1 for no limit (default: 4)")
    grp_heur.add_option("--max_losses", dest="max_losses",
                        metavar="<max # of losses>",
                        default=4, type="int",
                        help="maximum # of losses (in each ancestral species), " +\
                             "set to -1 for no limit (default: 4)")
    parser.add_option_group(grp_heur)

    grp_misc = optparse.OptionGroup(parser, "Miscellaneous")
    grp_misc.add_option("-x", "--seed", dest="seed",
                        metavar="<random seed>",
                        type="int", default=None,
                        help="random number seed")
    parser.add_option_group(grp_misc)

    grp_info = optparse.OptionGroup(parser, "Information")
    common.move_option(parser, "--version", grp_info)
    common.move_option(parser, "--help", grp_info)
    grp_info.add_option("-l", "--log", dest="log",
                        action="store_true",
                        help="if given, output debugging log")
    parser.add_option_group(grp_info)

    options, treefiles = parser.parse_args()

    #=============================
    # check arguments

    # input gene tree files
    if len(treefiles) == 0:
        parser.error("must specify input file(s)")

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

    # max loci/dups/losses options
    if options.max_loci == -1:
        options.max_loci = util.INF
    elif options.max_loci < 1:
        parser.error("--max_loci must be > 0")

    if options.max_dups == -1:
        options.max_dups = util.INF
    elif options.max_dups < 1:
        parser.error("--max_dups must be > 0")

    if options.max_losses == -1:
        options.max_losses = util.INF
    elif options.max_losses < 1:
        parser.error("--max_losses must be > 0")

    return options, treefiles


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

def main():
    # parse arguments
    options, treefiles = parse_args()

    # read species tree and species map (and possibly locus map)
    stree = treelib.read_tree(options.stree)
    common.check_tree(stree, options.stree)
    gene2species = phylo.read_gene2species(options.smap)
    if options.lmap:
        gene2locus = phylo.read_gene2species(options.lmap)
    else:
        gene2locus = None

    # random seed if not specified
    if options.seed is None:
        # note: numpy requires 32-bit unsigned integer
        options.seed = int(time.time() * 100) % (2^32-1)

    # process genes trees
    for treefile in treefiles:
        # general output path
        out = util.replace_ext(treefile, options.inext, options.outext)

        # info file
        out_info = util.open_stream(out + ".info", "w")

        # log file
        if options.log:
            out_log = gzip.open(out + ".log.gz", 'w')
        else:
            out_log = common.NullLog()

        # command
        cmd = "%s %s\n\n" % (os.path.basename(sys.argv[0]),
                             ' '.join(map(lambda x: x if x.find(' ') == -1 else "\"%s\"" % x,
                                          sys.argv[1:])))
        out_info.write("DLCscape version:\t%s\n" % VERSION)
        out_info.write("DLCscape command:\t%s\n\n" % cmd)
        out_log.write("DLCscape version: %s\n" % VERSION)
        out_log.write("DLCscape executed with the following arguments:\n")
        out_log.write("%s\n\n" % cmd)

        # read and prepare coal tree
        coal_trees = list(treelib.iter_trees(treefile))

        # multiple coal trees not supported
        if len(coal_trees) > 1:
            raise Exception("unsupported: multiple coal trees per file")

        # get coal tree
        coal_tree = coal_trees[0]
        common.check_tree(coal_tree, treefile)

        # remove bootstrap and distances if they exist
        coal_tree_top = coal_tree.copy()
        for node in coal_tree_top:
            if "boot" in node.data:
                del node.data["boot"]
            node.dist = 0
        coal_tree_top.default_data.clear()

        # set random seed
        random.seed(options.seed)
        np.random.seed(options.seed)
        out_log.write("seed: %d\n\n" % options.seed)

        compute_events = (options.events is not None)

        # perform reconciliation
        return_vals = dlcpar.reconscape.dlcscape_recon(
            coal_tree_top, stree, gene2species, gene2locus,
            duprange=options.duprange, lossrange=options.lossrange,
            max_loci=options.max_loci, max_dups=options.max_dups, max_losses=options.max_losses,
            compute_events=compute_events,
            log=out_log)
        out_log.write("\n")

        # infeasible
        if return_vals[1] is None:
            out_info.write("Feasibility:\tinfeasible\n")
            out_info.close()
            if options.log:
                out_log.close()
            break

        # feasible
        cvs, runtime = return_vals

        # determine landscape
        regions = dlcpar.reconscape.get_regions(
            cvs, duprange=options.duprange, lossrange=options.lossrange,
            log=out_log)

        # write the runtime to the logging files
        out_info.write("Runtime:\t%f sec\n" % runtime)

        # write regions
        dlcpar.reconscape.write_regions(out + ".regions", regions,
                                        duprange=options.duprange, lossrange=options.lossrange,
                                        close=True)

        # write events
        if compute_events:
            dlcpar.reconscape.write_events(out + ".events", regions,
                                           intersect=(options.events=="I"),
                                           close=True)

        # draw regions if desired
        if options.draw_regions:
            dlcpar.reconscape.draw_landscape(regions, duprange=options.duprange, lossrange=options.lossrange)

        # end logging
        if options.log:
            out_log.close()

        # end info
        out_info.close()


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

