#!/usr/bin/env python3
# (C) 2026 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.

"""Enumerate connected fragment combinations."""

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

from openeye import oechem, oemedchem
from rich.console import Console
from rich_argparse import HelpPreviewAction, RichHelpFormatter

__SCRIPT_NAME__ = Path(__file__).absolute().stem
__SCRIPT_DESC__ = "Enumerate connected fragment combinations."
__SCRIPT_TOOLKITS__ = ["oechem", "oemedchem"]
__SCRIPT_CATEGORIES__ = ["cheminformatics"]


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

    io_group = parser.add_argument_group("Input/output options")
    io_group.add_argument(
        "--in",
        dest="input",
        type=str,
        required=True,
        metavar="MOL-FILE",
        help="input molecule file",
    )
    io_group.add_argument(
        "--out",
        dest="output",
        type=str,
        required=True,
        metavar="MOL-FILE",
        help="output molecule file",
    )

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

    parser.add_argument("--help-image", action=HelpPreviewAction)
    parser.add_argument(
        "--save-console-svg",
        default=False,
        action="store_true",
        help=f"run command and capture console output in {__SCRIPT_NAME__}.svg file",
    )

    return parser.parse_args()


def main() -> int:
    """Enumerate connected fragment combinations."""
    args = parse_options()
    console = Console(record=args.save_console_svg)

    ifs = oechem.oemolistream()
    if not ifs.open(args.input):
        oechem.OEThrow.Fatal(f"Cannot open input file: {args.input}")

    ofs = oechem.oemolostream()
    if not ofs.open(args.output):
        oechem.OEThrow.Fatal(f"Cannot open output file: {args.output}")

    mol = oechem.OEGraphMol()
    if not oechem.OEReadMolecule(ifs, mol):
        oechem.OEThrow.Fatal(f"Cannot read molecule from: {args.input}")
    console.print(
        f"Input molecule: {oechem.OEMolToSmiles(mol)}", highlight=False, markup=False
    )

    frag_func = _get_fragmentation_function(args.frag_type)

    frags = list(frag_func(mol))
    console.print(f"{len(frags)} fragments generated")

    frag_combs = get_fragment_combinations(mol, frags)
    console.print(f"{len(frag_combs)} fragment combinations generated")

    for frag in frag_combs:
        oechem.OEWriteMolecule(ofs, frag)

    if args.save_console_svg:
        console.save_svg(f"{__SCRIPT_NAME__}.svg", title="output")

    return os.EX_OK


def is_adjacent_atom_bond_sets(
    frag_a: oechem.OEAtomBondSet, frag_b: oechem.OEAtomBondSet
) -> bool:
    """Check if two atom/bond sets share an adjacent bond."""
    for atom_a in frag_a.GetAtoms():
        for atom_b in frag_b.GetAtoms():
            if atom_a.GetBond(atom_b) is not None:
                return True
    return False


def is_adjacent_atom_bond_set_combination(
    frag_list: list[oechem.OEAtomBondSet],
) -> bool:
    """Check if a list of fragments forms a single connected component."""
    parts = [0] * len(frag_list)
    num_parts = 0

    for idx, frag in enumerate(frag_list):
        if parts[idx] != 0:
            continue

        num_parts += 1
        parts[idx] = num_parts
        traverse_fragments(frag, frag_list, parts, num_parts)

    return num_parts == 1


def traverse_fragments(
    act_frag: oechem.OEAtomBondSet,
    frag_list: list[oechem.OEAtomBondSet],
    parts: list[int],
    num_parts: int,
) -> None:
    """Recursively traverse adjacent fragments to assign connected components."""
    for idx, frag in enumerate(frag_list):
        if parts[idx] != 0:
            continue

        if not is_adjacent_atom_bond_sets(act_frag, frag):
            continue

        parts[idx] = num_parts
        traverse_fragments(frag, frag_list, parts, num_parts)


def combine_and_connect_atom_bond_sets(
    frag_list: tuple[oechem.OEAtomBondSet, ...],
) -> oechem.OEAtomBondSet:
    """Combine fragment atom/bond sets and add connecting bonds."""
    combined = oechem.OEAtomBondSet()
    for frag in frag_list:
        for atom in frag.GetAtoms():
            combined.AddAtom(atom)
        for bond in frag.GetBonds():
            combined.AddBond(bond)

    for atom_a in combined.GetAtoms():
        for atom_b in combined.GetAtoms():
            if atom_a.GetIdx() < atom_b.GetIdx():
                continue

            bond = atom_a.GetBond(atom_b)
            if bond is None:
                continue
            if combined.HasBond(bond):
                continue

            combined.AddBond(bond)

    return combined


def get_fragment_atom_bond_set_combinations(
    frag_list: list[oechem.OEAtomBondSet],
) -> list[oechem.OEAtomBondSet]:
    """Generate all valid adjacent fragment combinations."""
    frag_combs: list[oechem.OEAtomBondSet] = []

    num_frags = len(frag_list)
    for n in range(2, num_frags):
        for frag_comb in combinations(frag_list, n):
            if is_adjacent_atom_bond_set_combination(list(frag_comb)):
                frag = combine_and_connect_atom_bond_sets(frag_comb)
                frag_combs.append(frag)

    return frag_combs


def get_fragment_combinations(
    mol: oechem.OEMolBase, frag_list: list[oechem.OEAtomBondSet]
) -> list[oechem.OEGraphMol]:
    """Generate all connected fragment combination molecules from a fragmented molecule."""
    fragments: list[oechem.OEGraphMol] = []
    frag_combs = get_fragment_atom_bond_set_combinations(frag_list)

    for f in frag_combs:
        frag_atom_pred = oechem.OEIsAtomMember(f.GetAtoms())
        frag_bond_pred = oechem.OEIsBondMember(f.GetBonds())

        fragment = oechem.OEGraphMol()
        adjust_h_count = True
        oechem.OESubsetMol(
            fragment, mol, frag_atom_pred, frag_bond_pred, adjust_h_count
        )
        fragments.append(fragment)

    return fragments


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]]:
    """Return the fragmentation function for the given type."""
    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())
