diff --git a/src/sire/restraints/__init__.py b/src/sire/restraints/__init__.py index 07371e06a..6afff5aea 100644 --- a/src/sire/restraints/__init__.py +++ b/src/sire/restraints/__init__.py @@ -6,6 +6,7 @@ "distance", "boresch", "boresch_search", + "check_boresch_search", "inverse_bond", "inverse_distance", "rmsd", @@ -26,4 +27,4 @@ rmsd, ) from ._standard_state_correction import get_standard_state_correction -from ._boresch_search import boresch_search +from ._boresch_search import boresch_search, check_boresch_search diff --git a/src/sire/restraints/_boresch_search.py b/src/sire/restraints/_boresch_search.py index 2e59b2d67..adf195f8b 100644 --- a/src/sire/restraints/_boresch_search.py +++ b/src/sire/restraints/_boresch_search.py @@ -2,7 +2,7 @@ Automatic Boresch restraint generation for ABFE simulations. """ -__all__ = ["boresch_search"] +__all__ = ["boresch_search", "check_boresch_search"] from collections import deque as _deque @@ -51,6 +51,75 @@ def _build_triplets(connectivity, anchor_idx, mol, is_lig, ghost_elem, h_elem): return triplets +def _find_ligand(system): + """ + Return the perturbable molecule and its non-ghost, non-H atom indices + (lambda=0 state), validating that there is enough to build a restraint. + """ + from ..legacy import Mol as _SireMol + + pert_mols = system.molecules("property is_perturbable") + if pert_mols.num_molecules() != 1: + raise ValueError( + "System must contain exactly one perturbable molecule for Boresch " + f"restraint generation; found {pert_mols.num_molecules()}." + ) + pert_mol = pert_mols.molecule(0) + + ghost_elem = _SireMol.Element(0) + h_elem = _SireMol.Element("H") + + heavy_idxs = [] + for atom in pert_mol.atoms(): + elem0 = atom.property("element0") + if elem0 != ghost_elem and elem0 != h_elem: + heavy_idxs.append(atom.index()) + + if len(heavy_idxs) < 3: + raise ValueError( + f"Ligand has only {len(heavy_idxs)} non-ghost, non-hydrogen atom(s); " + "need at least 3 for Boresch restraints." + ) + + return pert_mol, heavy_idxs + + +def _find_rxrx_ligand_atoms(pert_mol, heavy_idxs): + """ + Return the RXRX candidate anchor atoms (non-terminal heavy atoms) and the + ligand N/O atoms available as hydrogen-bond partners. + """ + from ..legacy import Mol as _SireMol + + ghost_elem = _SireMol.Element(0) + h_elem = _SireMol.Element("H") + hbond_elems = (_SireMol.Element("N"), _SireMol.Element("O")) + connectivity = pert_mol.connectivity() + + candidate_set = { + idx + for idx in heavy_idxs + if len(_nonH_bonded(connectivity, idx, pert_mol, True, ghost_elem, h_elem)) >= 2 + } + if not candidate_set: + raise ValueError( + "No non-terminal ligand heavy atoms (bonded to >= 2 heavy atoms) " + "were found for RXRX restraint search." + ) + + hbond_idxs = [ + idx + for idx in heavy_idxs + if pert_mol.atom(idx).property("element0") in hbond_elems + ] + if not hbond_idxs: + raise ValueError( + "Ligand has no N/O heavy atoms to use as hydrogen-bond partners." + ) + + return candidate_set, hbond_idxs + + def _circular_mean_std(values): """ Circular mean/std (radians) for periodic dihedral samples. @@ -470,8 +539,6 @@ def _boresch_search_rxrx( ghost_elem = _SireMol.Element(0) h_elem = _SireMol.Element("H") - n_elem = _SireMol.Element("N") - o_elem = _SireMol.Element("O") # ------------------------------------------------------------------------- # 1. Locate the perturbable molecule (ligand) and its candidate atoms. @@ -480,27 +547,10 @@ def _boresch_search_rxrx( # pool; the candidate anchor set is restricted to non-terminal atoms # (bonded to >= 2 non-H heavy atoms), matching the RXRX paper's Figure 2B. # ------------------------------------------------------------------------- - pert_mols = system.molecules("property is_perturbable") - if pert_mols.num_molecules() != 1: - raise ValueError( - "System must contain exactly one perturbable molecule for Boresch " - f"restraint generation; found {pert_mols.num_molecules()}." - ) - pert_mol = pert_mols.molecule(0) + pert_mol, lig_heavy_idxs = _find_ligand(system) pert_mol_num = pert_mol.number() lig_connectivity = pert_mol.connectivity() - - lig_heavy_idxs = [] - for atom in pert_mol.atoms(): - elem0 = atom.property("element0") - if elem0 != ghost_elem and elem0 != h_elem: - lig_heavy_idxs.append(atom.index()) - - if len(lig_heavy_idxs) < 3: - raise ValueError( - f"Ligand has only {len(lig_heavy_idxs)} non-ghost, non-hydrogen atom(s); " - "need at least 3 for Boresch restraints." - ) + candidate_set, lig_hbond_idxs = _find_rxrx_ligand_atoms(pert_mol, lig_heavy_idxs) # Lambda=0 atomic masses, used for the mass-weighted centre of mass below # (topology doesn't change between frames, so this is computed once). @@ -508,20 +558,6 @@ def _boresch_search_rxrx( [float(pert_mol.atom(idx).property("mass0").value()) for idx in lig_heavy_idxs] ) - candidate_set = set() - for idx in lig_heavy_idxs: - n_heavy = len( - _nonH_bonded(lig_connectivity, idx, pert_mol, True, ghost_elem, h_elem) - ) - if n_heavy >= 2: - candidate_set.add(idx) - - if not candidate_set: - raise ValueError( - "No non-terminal ligand heavy atoms (bonded to >= 2 heavy atoms) " - "were found for RXRX restraint search." - ) - def _attached_H(connectivity, at_idx, mol, is_lig): bonded = connectivity.connections_to(at_idx) result = [] @@ -531,15 +567,6 @@ def _attached_H(connectivity, at_idx, mol, is_lig): result.append(b_idx) return result - lig_hbond_idxs = [ - idx - for idx in lig_heavy_idxs - if pert_mol.atom(idx).property("element0") in (n_elem, o_elem) - ] - if not lig_hbond_idxs: - raise ValueError( - "Ligand has no N/O heavy atoms to use as hydrogen-bond partners." - ) lig_donor_H = { idx: _attached_H(lig_connectivity, idx, pert_mol, True) for idx in lig_hbond_idxs @@ -1105,31 +1132,12 @@ def _boresch_search_aldeghi( # ------------------------------------------------------------------------- # 1. Locate the perturbable molecule (ligand). # ------------------------------------------------------------------------- - pert_mols = system.molecules("property is_perturbable") - if pert_mols.num_molecules() != 1: - raise ValueError( - "System must contain exactly one perturbable molecule for Boresch " - f"restraint generation; found {pert_mols.num_molecules()}." - ) - pert_mol = pert_mols.molecule(0) + pert_mol, lig_atom_idxs = _find_ligand(system) pert_mol_num = pert_mol.number() ghost_elem = _SireMol.Element(0) h_elem = _SireMol.Element("H") - # Collect non-ghost, non-H ligand AtomIdx values (lambda=0 state). - lig_atom_idxs = [] - for atom in pert_mol.atoms(): - elem0 = atom.property("element0") - if elem0 != ghost_elem and elem0 != h_elem: - lig_atom_idxs.append(atom.index()) - - if len(lig_atom_idxs) < 3: - raise ValueError( - f"Ligand has only {len(lig_atom_idxs)} non-ghost, non-hydrogen atom(s); " - "need at least 3 for Boresch restraints." - ) - lig_connectivity = pert_mol.connectivity() # ------------------------------------------------------------------------- @@ -1587,3 +1595,46 @@ def boresch_search( raise ValueError( f"Unknown 'protocol'={protocol!r}; must be 'rxrx' or 'aldeghi'." ) + + +def check_boresch_search(system, protocol="rxrx"): + """ + Check that the ligand in a system is suitable for ``boresch_search`` + with the given protocol, without requiring any trajectory frames. This + allows topology-level failures, such as a ligand with no N/O atoms to + act as hydrogen-bond partners for the "rxrx" protocol, to be caught + before the restraint-search trajectory is generated. + + Parameters + ---------- + + system : sire.system.System + A Sire system containing exactly one perturbable molecule. + + protocol : str + The restraint search protocol to check: ``"rxrx"`` (default) or + ``"aldeghi"``. + + Raises + ------ + + ValueError + If ``boresch_search`` is guaranteed to fail for this system and + protocol. The message matches the one ``boresch_search`` would raise. + """ + from ..system import System as _System + + if not isinstance(system, _System): + raise TypeError( + f"'system' must be of type 'sire.system.System', got {type(system)}" + ) + + if protocol not in ("rxrx", "aldeghi"): + raise ValueError( + f"Unknown 'protocol'={protocol!r}; must be 'rxrx' or 'aldeghi'." + ) + + pert_mol, heavy_idxs = _find_ligand(system) + + if protocol == "rxrx": + _find_rxrx_ligand_atoms(pert_mol, heavy_idxs) diff --git a/tests/restraints/test_boresch_search.py b/tests/restraints/test_boresch_search.py index d1cc7282a..2a5d976a4 100644 --- a/tests/restraints/test_boresch_search.py +++ b/tests/restraints/test_boresch_search.py @@ -241,3 +241,50 @@ def test_tight_cutoff_raises(self, abfe_system): with pytest.raises(ValueError, match="cutoff"): boresch_search(abfe_system, protocol="aldeghi", cutoff="0.1 A") + + +@pytest.fixture(scope="module") +def no_hbond_system(abfe_system): + """The test system with the ligand's N/O atoms changed to carbon.""" + mols = abfe_system.clone() + cursor = mols.molecules("property is_perturbable")[0].cursor() + for atom in cursor.atoms(): + if atom["element0"].symbol() in ("N", "O"): + atom["element0"] = sr.mol.Element("C") + mols.update(cursor.commit()) + return mols + + +class TestCheckBoreschSearch: + @pytest.mark.parametrize("protocol", PROTOCOLS) + def test_valid_system_passes(self, abfe_system, protocol): + from sire.restraints import check_boresch_search + + check_boresch_search(abfe_system, protocol=protocol) + + def test_no_frames_required(self, abfe_system): + from sire.restraints import check_boresch_search + + mols = abfe_system.clone() + mols.delete_all_frames() + check_boresch_search(mols) + + def test_rxrx_no_hbond_partners_raises(self, no_hbond_system): + from sire.restraints import boresch_search, check_boresch_search + + with pytest.raises(ValueError, match="N/O"): + check_boresch_search(no_hbond_system, protocol="rxrx") + + with pytest.raises(ValueError, match="N/O"): + boresch_search(no_hbond_system, protocol="rxrx") + + def test_aldeghi_no_hbond_partners_passes(self, no_hbond_system): + from sire.restraints import check_boresch_search + + check_boresch_search(no_hbond_system, protocol="aldeghi") + + def test_invalid_protocol_raises(self, abfe_system): + from sire.restraints import check_boresch_search + + with pytest.raises(ValueError, match="protocol"): + check_boresch_search(abfe_system, protocol="nonsense")