#!/usr/bin/env python3
# (C) 2023 Cadence Design Systems, Inc. (Cadence)
# All rights reserved.
# TERMS FOR USE OF SAMPLE CODE The software below ("Sample Code") is
# provided to current licensees or subscribers of Cadence products or
# SaaS offerings (each a "Customer").
# Customer is hereby permitted to use, copy, and modify the Sample Code,
# subject to these terms. Cadence claims no rights to Customer's
# modifications. Modification of Sample Code is at Customer's sole and
# exclusive risk. Sample Code may require Customer to have a then
# current license or subscription to the applicable Cadence offering.
# THE SAMPLE CODE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
# EXPRESS OR IMPLIED.  OPENEYE DISCLAIMS ALL WARRANTIES, INCLUDING, BUT
# NOT LIMITED TO, WARRANTIES OF MERCHANTABILITY, FITNESS FOR A
# PARTICULAR PURPOSE AND NONINFRINGEMENT. In no event shall Cadence be
# liable for any damages or liability in connection with the Sample Code
# or its use.


"""Calculates XLogP of set of molecules and visualizes the fragment contributions."""

import argparse
import enum
import os
import pathlib
import sys
from collections.abc import Callable, Iterator
from pathlib import Path

import rich.console
from openeye import oechem, oedepict, oegrapheme, oemedchem, oemolprop, oequacpac
from rich_argparse import HelpPreviewAction, RichHelpFormatter

__SCRIPT_NAME__ = Path(__file__).absolute().stem
__SCRIPT_DESC__ = "Depict XLogP of molecule (fragment base)."
__SCRIPT_TOOLKITS__ = [
    "oechem",
    "oedepict",
    "oegrapheme",
    "oemolprop",
    "oequacpac",
    "oemedchem",
]

__SCRIPT_NAME__ = pathlib.Path(__file__).absolute().stem
__SCRIPT_DESC__ = "Depict XLogP of set of molecules."
__SCRIPT_TOOLKITS__ = ["oechem", "oedepict", "oegrapheme", "oemolprop", "oequacpac"]
__SCRIPT_CATEGORIES__ = ["depiction"]


def parse_options() -> argparse.Namespace:
    """Set up command line options."""
    parser = argparse.ArgumentParser(
        add_help=True,
        formatter_class=RichHelpFormatter,
        description="[yellow]" + __SCRIPT_DESC__ + "[/yellow]",
    )
    parser.add_argument("--help-image", action=HelpPreviewAction)

    input_group = parser.add_argument_group("Input options")
    input_group.add_argument(
        "--mol",
        type=str,
        required=True,
        metavar="MOL-FILE",
        help="input multi-conformer molecule file (oeb, sdf)",
    )

    frag_group = parser.add_argument_group("Fragmentation options")
    frag_group.add_argument(
        "--frag-type",
        "--fragmentation-type",
        type=FragmentationType,
        default=FragmentationType.FunctionalGroup,
        choices=list(FragmentationType),
        help="type of fragmentation to perform (default: %(default)s)",
    )

    report_group = parser.add_argument_group("Report options")
    report_group.add_argument(
        "--report",
        type=str,
        required=True,
        metavar="REPORT-FILE",
        help="output report file (PDF)",
    )
    report_group.add_argument(
        "--rows",
        type=int,
        default=3,
        choices=range(2, 6),
        metavar="N",
        help="number of rows per page (default: %(default)s)",
    )
    report_group.add_argument(
        "--cols",
        type=int,
        default=2,
        choices=range(1, 3),
        metavar="N",
        help="number of columns per page (default: %(default)s)",
    )
    report_group.add_argument(
        "--page-by-page",
        action="store_true",
        help="write pages of report to separate numbered image files",
    )
    return parser.parse_args()


def main() -> int:
    """Visualizes XLogP of set of molecules."""
    args = parse_options()

    _check_report_file(args)

    mols: list[oechem.OEMolBase] = _read_molecules(args.mol)
    console = rich.console.Console()
    console.print(f"Imported {len(mols)} molecules from {args.mol}")

    # initialize multi-page report

    report_opts = oedepict.OEReportOptions(args.rows, args.cols)
    report_opts.SetHeaderHeight(35)
    report_opts.SetFooterHeight(45)
    report_opts.SetPageMargins(10)
    report_opts.SetCellGap(5)
    report = oedepict.OEReport(report_opts)

    # setup depiction options

    width, height = report.GetCellWidth(), report.GetCellHeight()
    opts = oedepict.OE2DMolDisplayOptions(width, height, oedepict.OEScale_AutoScale)
    opts.SetAtomColorStyle(oedepict.OEAtomColorStyle_WhiteMonochrome)

    frag_func = _get_fragmentation_function(args.frag_type)

    depict_molecules_fragment_xlogp(report, mols, frag_func, opts)

    if args.page_by_page:
        oedepict.OEWriteReportPageByPage(args.report, report)
    else:
        oedepict.OEWriteReport(args.report, report)

    return os.EX_OK


def set_atom_properties(mol: oechem.OEMolBase, data_tag: int) -> None:
    """Attach the XLogP atom contribution to each atom with the given tag."""
    oequacpac.OERemoveFormalCharge(mol)

    atom_values = oechem.OEFloatArray(mol.GetMaxAtomIdx())
    xlogp = oemolprop.OEGetXLogP(mol, atom_values)

    mol.SetTitle(f"{mol.GetTitle()} -- OEXLogP = {xlogp:.2f}")

    for atom in mol.GetAtoms():
        val = atom_values[atom.GetIdx()]
        atom.SetData(data_tag, val)


def fragment_molecule(
    mol: oechem.OEMolBase,
    frag_func: Callable[[oechem.OEMolBase], Iterator[oechem.OEAtomBondSet]],
    group_tag: int,
) -> None:
    """Fragments the molecule and stores each fragment as a group on the molecule."""
    for frag in frag_func(mol):
        atoms = oechem.OEAtomVector()
        for atom in frag.GetAtoms():
            atoms.append(atom)
        bonds = oechem.OEBondVector()
        for bond in frag.GetBonds():
            bonds.append(bond)

        mol.NewGroup(group_tag, atoms, bonds)  # type: ignore[attr-defined]


def set_fragment_properties(
    mol: oechem.OEMolBase,
    data_tag: int,
    group_tag: int,
    min_value: float,
    max_value: float,
) -> tuple[float, float]:
    """Calculate the fragment contribution based on attached atom properties for pre-generated fragments."""
    for group in mol.GetGroups(oechem.OEHasGroupType(group_tag)):
        sum_prop = 0.0
        for atom in group.GetAtoms():
            sum_prop += atom.GetData(data_tag)
        group.SetData(data_tag, sum_prop)

        min_value = min(min_value, sum_prop)
        max_value = max(max_value, sum_prop)

    return min_value, max_value


def depict_molecules_fragment_xlogp(
    report: oedepict.OEReport,
    mols: list[oechem.OEMolBase],
    frag_func: Callable[[oechem.OEMolBase], Iterator[oechem.OEAtomBondSet]],
    opts: oedepict.OE2DMolDisplayOptions,
) -> None:
    """Generate a report of molecules depicting the fragment contribution of XLogP."""
    # calculate atom contributions of XLogP
    data_tag: int = oechem.OEGetTag("XLogP")

    for mol in mols:
        set_atom_properties(mol, data_tag)

    # fragment molecules
    group_tag: int = oechem.OEGetTag("fragment")

    for mol in mols:
        fragment_molecule(mol, frag_func, group_tag)

    # calculate fragment contributions
    min_value, max_value = float("inf"), float("-inf")
    for mol in mols:
        min_value, max_value = set_fragment_properties(
            mol, data_tag, group_tag, min_value, max_value
        )

    # initialize color gradient
    lightgrey = oechem.OEColor(240, 240, 240)
    color_gradient = oechem.OELinearColorGradient(oechem.OEColorStop(0.0, lightgrey))
    color_gradient.AddStop(oechem.OEColorStop(min_value, oechem.OEDarkGreen))
    color_gradient.AddStop(oechem.OEColorStop(max_value, oechem.OEDarkPurple))

    # initialize highlighting style
    highlight = oedepict.OEHighlightByLasso(oechem.OEWhite)
    highlight.SetConsiderAtomLabelBoundingBox(True)

    for mol in mols:
        # generate image frames
        cell = report.NewCell()
        cell_width, cell_height = cell.GetWidth(), cell.GetHeight()
        mol_frame = oedepict.OEImageFrame(
            cell, cell_width, cell_height * 0.8, oedepict.OE2DPoint(0.0, 0.0)
        )
        color_frame = oedepict.OEImageFrame(
            cell,
            cell_width,
            cell_height * 0.2,
            oedepict.OE2DPoint(0.0, cell_height * 0.8),
        )

        # initialize molecule display

        opts.SetDimensions(
            mol_frame.GetWidth(), mol_frame.GetHeight(), oedepict.OEScale_AutoScale
        )

        oedepict.OEPrepareDepiction(mol)
        disp = oedepict.OE2DMolDisplay(mol, opts)

        color_gradient_opts = oegrapheme.OEColorGradientDisplayOptions()

        for group in mol.GetGroups(oechem.OEHasGroupType(group_tag)):
            group_value = group.GetData(data_tag)
            color_gradient_opts.AddMarkedValue(group_value)

            # depict fragment contribution
            color = color_gradient.GetColorAt(group_value)
            highlight.SetColor(color)

            ab_set = oechem.OEAtomBondSet(group.GetAtoms(), group.GetBonds())
            oedepict.OEAddHighlighting(disp, highlight, ab_set)

        # render molecule and color gradient

        oedepict.OERenderMolecule(mol_frame, disp)
        oegrapheme.OEDrawColorGradient(color_frame, color_gradient, color_gradient_opts)


def _check_report_file(args: argparse.Namespace) -> bool:
    ext = pathlib.Path(args.report).suffix[1:]
    if not oedepict.OEIsRegisteredImageFile(ext):
        oechem.OEThrow.Fatal("Unknown image outout type!")

    if not args.page_by_page and not oedepict.OEIsRegisteredMultiPageImageFile(ext):
        oechem.OEThrow.Warning("Report will be generated into separate pages!")
        args.page_by_page = True

    return True


def _read_molecules(filename: str) -> list[oechem.OEMolBase]:
    ifs = oechem.oemolistream()
    if not ifs.open(filename):
        oechem.OEThrow.Fatal(f"Cannot open {filename} input file!")

    mols: list[oechem.OEMolBase] = [oechem.OEGraphMol(m) for m in ifs.GetOEGraphMols()]
    if not mols:
        oechem.OEThrow.Fatal(f"No molecules could be read from {filename}")
    return mols


class FragmentationType(enum.Enum):
    """Molecule fragmentation type."""

    FunctionalGroup = "func-group"
    RingChain = "ring-chain"
    RingLinkerSideChain = "ring-linker-sidechain"

    def __str__(self) -> str:
        """Convert to string representation."""
        return self.value


def _get_fragmentation_function(
    frag_type: FragmentationType,
) -> Callable[[oechem.OEMolBase], Iterator[oechem.OEAtomBondSet]]:
    match frag_type:
        case FragmentationType.RingChain:
            return oemedchem.OEGetRingChainFragments
        case FragmentationType.RingLinkerSideChain:
            return oemedchem.OEGetRingLinkerSideChainFragments
    return oemedchem.OEGetFuncGroupFragments


setattr(main, "__SCRIPT_NAME__", __SCRIPT_NAME__)
setattr(main, "__SCRIPT_DESC__", __SCRIPT_DESC__)
setattr(main, "__SCRIPT_TOOLKITS__", __SCRIPT_TOOLKITS__)
setattr(main, "__SCRIPT_CATEGORIES__", __SCRIPT_CATEGORIES__)

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