Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ jobs:
run: flake8 meze/ --count --select=E9,F63,F7,F82 --show-source --statistics

- name: Lint (style/complexity report only)
run: flake8 meze/ --count --exit-zero --max-complexity=10 --max-line-length=127 --statistics
run: flake8 meze/ --count --exit-zero --max-complexity=10 --max-line-length=79 --statistics

test:
runs-on: ubuntu-latest
Expand Down
24 changes: 15 additions & 9 deletions meze/ligand.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,15 +61,16 @@ def _validate_file(file):
return [file]
elif isinstance(file, list):
if len(file) > 2:
raise ValueError(
message = (
f"Too many values for 'file': {file}."
f"Expected a 'str' or a list of at most 2 input files."
"Expected a 'str' or a list of at most 2 input files."
)
log.error(message)
raise ValueError(message)
return file

raise TypeError(
f"Expected str or list[str], got {type(file)}"
)
message = f"Expected str or list[str], got {type(file)}"
log.error(message)
raise TypeError(message)

@staticmethod
def _check_files_exist(files):
Expand All @@ -84,10 +85,12 @@ def _validate_charge(charge):
try:
return float(charge)
except (TypeError, ValueError):
raise TypeError(
f"Ligand charge must be an integer or float "
message = (
"Ligand charge must be an integer or float "
f"(got {charge} of type {type(charge)})."
)
log.error(message)
raise TypeError(message)

@staticmethod
def _infer_ligand_name(file):
Expand Down Expand Up @@ -124,7 +127,10 @@ def parameterise(self,
directory = os.getcwd()

if len(self.file) > 1:
raise UserWarning(f"Expected one ligand file but got {self.file}")
warnings.warn(
f"Expected one ligand file but got {self.file}",
UserWarning
)
else:
file = self.file[0]

Expand Down
120 changes: 59 additions & 61 deletions meze/sofra.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,11 +34,11 @@
from MDAnalysis.core.groups import Residue as mdaResidue
import BioSimSpace as bss
from BioSimSpace._SireWrappers import System as bssSystem
from BioSimSpace.Types._time import Time as bssTime
from BioSimSpace.Types._temperature import Temperature as bssTemperature
from BioSimSpace.Types._pressure import Pressure as bssPressure
if TYPE_CHECKING:
from BioSimSpace.Protocol._protocol import Protocol as bssProtocol
from BioSimSpace.Types._time import Time as bssTime
from BioSimSpace.Types._temperature import Temperature as bssTemperature
from BioSimSpace.Types._pressure import Pressure as bssPressure
from .utils import (
_residue_restraint_mask,
_write_distance_restraints,
Expand Down Expand Up @@ -106,13 +106,11 @@ class MezeRecipe(BaseModel):
"g16", description="Gaussian version"
)
memory: float = Field(
12000, description="Memory for Gaussian calculations in MB"
12_000, description="Memory for Gaussian calculations in MB"
)

nprocshared: int = Field(
8, description="Number of processors for Gaussian calculations"
)

only_optimise_hydrogens: bool = Field(
True, description="Only optimise hydrogen atoms"
)
Expand Down Expand Up @@ -206,7 +204,7 @@ def __setitem__(self, key: str, value):

def to_json(self, file: str):
with open(file, "w") as ofile:
ofile.write(self.model_dump_json(indent=2))
ofile.write(self.model_dump_json(indent=2, fallback=str))


class ColdMezeRecipe(MezeRecipe):
Expand All @@ -232,10 +230,10 @@ class ColdMezeRecipe(MezeRecipe):
2, ge=1, le=2, description="Type of barostat, 1: Berendsen, 2: MC"
)

runtime: Union[float, bssTime] = Field(
runtime: Union[float, bss.Types.Time] = Field(
100.0, description="Simulation time in picoseconds"
)
dt: Union[float, bssTime] = Field(
dt: Union[float, bss.Types.Time] = Field(
0.001, description="Integrator timestep, in picoseconds"
)
start_temperature: float = Field(
Expand All @@ -255,9 +253,9 @@ class ColdMezeRecipe(MezeRecipe):
def validate_time(cls, value):
if isinstance(value, bss.Types.Time):
return value
if value < 0:
if value <= 0:
raise ValueError(
"dt, time must be greater than or equal to 0 picoseconds"
"dt, time must be greater than 0 picoseconds"
)
return bss.Types.Time(value, "picoseconds")

Expand All @@ -277,10 +275,10 @@ def validate_temperature_range(cls, value):
class HotMezeRecipe(MezeRecipe):
"""Meze workflow recipe for production runs
"""
runtime: Union[float, bssTime] = Field(
runtime: Union[float, bss.Types.Time] = Field(
100.0, description="Simulation time in nanoseconds"
)
dt: Union[float, bssTime] = Field(
dt: Union[float, bss.Types.Time] = Field(
0.002, description="Integrator timestep, in picoseconds"
)

Expand All @@ -289,9 +287,9 @@ class HotMezeRecipe(MezeRecipe):
def validate_time(cls, value):
if isinstance(value, bss.Types.Time):
return value
if value < 0:
if value <= 0:
raise ValueError(
"dt must be greater than or equal to 0 nanoseconds"
"dt must be greater than 0 nanoseconds"
)
return bss.Types.Time(value, "nanoseconds")

Expand All @@ -300,9 +298,9 @@ def validate_time(cls, value):
def validate_timestep(cls, value):
if isinstance(value, bss.Types.Time):
return value
if value < 0:
if value <= 0:
raise ValueError(
"dt must be greater than or equal to 0 picoseconds"
"dt must be greater than 0 picoseconds"
)
return bss.Types.Time(value, "picoseconds")

Expand All @@ -316,7 +314,7 @@ class AlchemicalMezeRecipe(MezeRecipe):
n_lambdas: int = Field(
16, ge=3, description="Number of lambda windows"
)
sampling_time: Union[float, bssTime] = Field(
sampling_time: Union[float, bss.Types.Time] = Field(
4.0, description="Runtime for each lambda window in ns."
)
restart_interval: int = Field(
Expand Down Expand Up @@ -355,21 +353,21 @@ class AlchemicalMezeRecipe(MezeRecipe):
@field_validator("sampling_time", mode="after")
@classmethod
def validate_sampling_time(cls, value):
if isinstance(value, bssTime):
if isinstance(value, bss.Types.Time):
return value
value = float(value)
if value <= 0:
raise ValueError("sampling_time must be greater than 0 ns")
return bssTime(value, "nanoseconds")
return bss.Types.Time(value, "nanoseconds")

@field_validator("dt", mode="after")
@classmethod
def validate_picosecond_times(cls, value):
if isinstance(value, bss.Types.Time):
return value
if value < 0:
if value <= 0:
raise ValueError(
"dt must be greater than or equal to 0 picoseconds"
"dt must be greater than 0 picoseconds"
)
return bss.Types.Time(value, "picoseconds")

Expand Down Expand Up @@ -2610,12 +2608,12 @@ def run(
barostat: Optional[int] = None,
n_sd_cycles: Optional[int] = None,
nb_cutoff: Optional[float] = None,
timestep: Optional[Union[float, bssTime]] = None,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = None,
start_temperature: Optional[Union[float, bssTemperature]] = 300,
end_temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = None,
timestep: Optional[Union[float, "bssTime"]] = None,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = None,
start_temperature: Optional[Union[float, "bssTemperature"]] = 300,
end_temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = None,
is_gpu: Optional[bool] = True,
engine_executable: Optional["str"] = None,
additional_positional_restraints: Optional[dict[str, Any]] = None,
Expand Down Expand Up @@ -2768,11 +2766,11 @@ def heat(
] = None,
restart: Optional[bool] = False,
restraint_weight: Optional[float] = None,
timestep: Optional[Union[float, bssTemperature]] = None,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = None,
start_temperature: Optional[Union[float, bssTemperature]] = 300,
end_temperature: Optional[Union[float, bssTemperature]] = 300,
timestep: Optional[Union[float, "bssTemperature"]] = None,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = None,
start_temperature: Optional[Union[float, "bssTemperature"]] = 300,
end_temperature: Optional[Union[float, "bssTemperature"]] = 300,
process_name: Optional[str] = "nvt",
is_gpu: Optional[bool] = True,
engine_executable: Optional[str] = None,
Expand Down Expand Up @@ -2811,10 +2809,10 @@ def pressurise(
] = None,
restart: Optional[bool] = False,
restraint_weight: Optional[float] = None,
timestep: Optional[Union[float, bssTemperature]] = None,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = 1.0,
timestep: Optional[Union[float, "bssTemperature"]] = None,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = 1.0,
process_name: Optional[str] = "npt",
is_gpu: Optional[bool] = True,
engine_executable: Optional[str] = None,
Expand Down Expand Up @@ -2930,10 +2928,10 @@ def run(
system: Optional[bssSystem] = None,
process_name: Optional[str] = "meze-run",
nb_cutoff: Optional[float] = None,
timestep: Optional[Union[float, bssTime]] = None,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = 1,
timestep: Optional[Union[float, "bssTime"]] = None,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = 1,
engine_executable: Optional[str] = None,
write_frequency: Optional[int] = 100000,
distance_write_frequency: Optional[int] = 10000,
Expand Down Expand Up @@ -3413,12 +3411,12 @@ def run(
barostat: Optional[int] = None,
n_sd_cycles: Optional[int] = None,
nb_cutoff: Optional[float] = None,
timestep: Optional[Union[float, bssTime]] = None,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = None,
start_temperature: Optional[Union[float, bssTemperature]] = 300,
end_temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = None,
timestep: Optional[Union[float, "bssTime"]] = None,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = None,
start_temperature: Optional[Union[float, "bssTemperature"]] = 300,
end_temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = None,
engine_executable: Optional[str] = None,
qm_theory: Optional[str] = "DFTB3",
metal_resids_for_distance_restraints: Optional[
Expand Down Expand Up @@ -3577,11 +3575,11 @@ def heat(
system: Optional[bssSystem] = None,
workdir: Optional[str] = None,
restart: Optional[bool] = False,
timestep: Optional[Union[float, bssTemperature]] = 0.001,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = None,
start_temperature: Optional[Union[float, bssTemperature]] = 300,
end_temperature: Optional[Union[float, bssTemperature]] = 300,
timestep: Optional[Union[float, "bssTemperature"]] = 0.001,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = None,
start_temperature: Optional[Union[float, "bssTemperature"]] = 300,
end_temperature: Optional[Union[float, "bssTemperature"]] = 300,
process_name: Optional[str] = "qm-nvt",
engine_executable: Optional[str] = None,
qm_theory: Optional[str] = "DFTB3",
Expand Down Expand Up @@ -3620,10 +3618,10 @@ def pressurise(
system: Optional[bssSystem] = None,
workdir: Optional[str] = None,
restart: Optional[bool] = False,
timestep: Optional[Union[float, bssTemperature]] = 0.001,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = 1.0,
timestep: Optional[Union[float, "bssTemperature"]] = 0.001,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = 1.0,
process_name: Optional[str] = "qm-npt",
engine_executable: Optional[str] = None,
qm_theory: Optional[str] = "DFTB3",
Expand Down Expand Up @@ -3718,10 +3716,10 @@ def run(
process_name: Optional[str] = "qm-meze-run",
ensemble: Optional[Literal["nvt", "npt"]] = "nvt",
nb_cutoff: Optional[float] = None,
timestep: Optional[Union[float, bssTime]] = 0.001,
runtime: Optional[Union[float, bssTime]] = None,
temperature: Optional[Union[float, bssTemperature]] = 300,
pressure: Optional[Union[float, bssPressure]] = None,
timestep: Optional[Union[float, "bssTime"]] = 0.001,
runtime: Optional[Union[float, "bssTime"]] = None,
temperature: Optional[Union[float, "bssTemperature"]] = 300,
pressure: Optional[Union[float, "bssPressure"]] = None,
engine_executable: Optional[str] = None,
write_frequency: Optional[int] = 500,
qm_theory: Optional[str] = "DFTB3",
Expand Down Expand Up @@ -4338,7 +4336,7 @@ def set_ligand_network(

log.info("Lomap finished succesfully. Parsing outputs.")
self.transformations, self.lomap_scores, network_file = (
lomap_directory, f"{self.group_name}_score_with_connection.txt"
self._parse_lomap_output(scores_file, lomap_directory)
)
self.save_network_file(network_file)

Expand Down
24 changes: 24 additions & 0 deletions tests/test_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
from unittest.mock import patch
import pytest
from meze.helpers import _check_ambertools


def test_check_ambertools_all_found_on_path():
with patch("meze.helpers.shutil.which", return_value="/usr/bin/tool"):
_check_ambertools()


def test_check_ambertools_found_via_amberhome():
with patch("meze.helpers.shutil.which", return_value=None), \
patch.dict("os.environ", {"AMBERHOME": "/opt/amber"}, clear=True), \
patch("meze.helpers.os.path.exists", return_value=True):
_check_ambertools()


def test_check_ambertools_missing_raises():
with patch("meze.helpers.shutil.which", return_value=None), \
patch.dict("os.environ", {}, clear=True):
with pytest.raises(
RuntimeError, match="AmberTools installation required"
):
_check_ambertools()
Loading
Loading