Skip to content
Open
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
60 changes: 46 additions & 14 deletions PQAnalysis/io/topology_file/topology_file_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from _io import TextIOWrapper as File # type: ignore

from beartype.typing import List
from beartype.typing import Iterable, List

from PQAnalysis import __package_name__
from PQAnalysis.io.base import BaseWriter
Expand Down Expand Up @@ -94,17 +94,57 @@ def write(self, bonded_topology: Topology | BondedTopology) -> None:
"Invalid bonded topology.", exception=TopologyFileError
)

self.open()

if bonded_topology.ordering_keys is not None:
keys = bonded_topology.ordering_keys
else:
keys = self.key_topology_map.keys()

for key in keys:
self.key_topology_map[key](bonded_topology, self.file)
self._check_types_given(bonded_topology, keys)

self.close()
self.open()

try:
for key in keys:
self.key_topology_map[key](bonded_topology, self.file)
finally:
self.close()

@classmethod
def _check_types_given(
cls, bonded_topology: BondedTopology, keys: Iterable[str]
) -> None:
"""
Checks that a type is defined for all bonded parameters to be written.

This check has to be performed before the output file is opened,
because opening it truncates an already existing file.

Parameters
----------
bonded_topology : BondedTopology
The bonded topology object to check.
keys : Iterable[str]
The keys of the topology sections that will be written.

Raises
------
TopologyFileError
If any bond, angle, dihedral or improper to be written
does not have a type defined.
"""

type_names = {
"bonds": "bond",
"angles": "angle",
"dihedrals": "dihedral",
"impropers": "improper",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

what about j coulping?

}

for key in keys:
if key in type_names:
cls._check_type_given(
getattr(bonded_topology, key), type_names[key]
)

@classmethod
def _write_bond_info(
Expand All @@ -127,8 +167,6 @@ def _write_bond_info(
If any bond in the bonded topology does not have a bond type defined.
"""

cls._check_type_given(bonded_topology.bonds, "bond")

if len(bonded_topology.bonds) != 0:
lines = cls._get_bond_lines(bonded_topology)
for line in lines:
Expand All @@ -155,8 +193,6 @@ def _write_angle_info(
If any angle in the bonded topology does not have an angle type defined.
"""

cls._check_type_given(bonded_topology.angles, "angle")

if len(bonded_topology.angles) != 0:
lines = cls._get_angle_lines(bonded_topology)
for line in lines:
Expand All @@ -183,8 +219,6 @@ def _write_dihedral_info(
If any dihedral in the bonded topology does not have a dihedral type defined.
"""

cls._check_type_given(bonded_topology.dihedrals, "dihedral")

if len(bonded_topology.dihedrals) != 0:
lines = cls._get_dihedral_lines(bonded_topology)
for line in lines:
Expand All @@ -205,8 +239,6 @@ def _write_improper_info(
The file object to write the improper information to.
"""

cls._check_type_given(bonded_topology.impropers, "improper")

if len(bonded_topology.impropers) != 0:
lines = cls._get_improper_lines(bonded_topology)
for line in lines:
Expand Down
18 changes: 18 additions & 0 deletions tests/io/test_topology_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,24 @@ def test__check_type_given(self):
"all dihedrals must have a dihedral type defined."
)

def test_write_does_not_truncate_on_invalid_topology(self, tmp_path):
"""
Test that a failing write leaves an already existing file untouched.
"""
filename = str(tmp_path / "topology.top")
content = "BONDS 1 1 0\n 1 2 1\nEND\n"

with open(filename, "w", encoding="utf-8") as file:
file.write(content)

topology = BondedTopology(bonds=[Bond(index1=1, index2=2)])

with pytest.raises(TopologyFileError):
TopologyFileWriter(filename, mode="o").write(topology)

with open(filename, "r", encoding="utf-8") as file:
assert file.read() == content

def test__get_bond_lines(self):
"""
Test the _get_bond_lines method.
Expand Down
Loading