Skip to content

Commit aa79f5c

Browse files
Add tests for the get_transfer_builder utility
1 parent d493a85 commit aa79f5c

2 files changed

Lines changed: 227 additions & 0 deletions

File tree

tests/utils/conftest.py

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
# -*- coding: utf-8 -*-
2+
"""A collection of useful pytest fixtures.
3+
4+
* make_tempdir: returns a function to create temporary directories
5+
6+
* makeget_computer: returns a function to create computers (note: these need to be destroyed after the test
7+
by using 'clear_database' because as the python variables go out of scope, the tempdir used to create the
8+
computers is deleted).
9+
10+
"""
11+
import pytest
12+
13+
# pylint: disable=redefined-outer-name
14+
15+
16+
# Maybe this could go on aiida-core
17+
@pytest.fixture(scope='function')
18+
def make_tempdir(request):
19+
"""Provide a function to generate new temporary directories.
20+
21+
:return: The path to the directory
22+
:rtype: str
23+
"""
24+
25+
def _make_tempdir(parent_dir=None):
26+
import tempfile
27+
import shutil
28+
29+
try:
30+
dirpath = tempfile.mkdtemp(dir=parent_dir)
31+
except FileNotFoundError as exc:
32+
raise ValueError('The parent_dir provided was never created') from exc
33+
34+
# After the test function has completed, remove the directory again
35+
# Design Note: yields would be a simpler way to do this, but can't use it
36+
# inside a fixture factory without having problems with 'generator' object
37+
def cleanup():
38+
shutil.rmtree(dirpath)
39+
40+
request.addfinalizer(cleanup)
41+
42+
return dirpath
43+
44+
return _make_tempdir
45+
46+
47+
# Maybe this could go on aiida-core
48+
# It would be preferrable to do usefixtures but this just works for test, not other
49+
# fixtures (see https://github.com/pytest-dev/pytest/issues/3664).
50+
# We'll have to leave the pylint ignore until then
51+
#@pytest.mark.usefixtures('clear_database')
52+
# pylint: disable=unused-argument
53+
@pytest.fixture(scope='function')
54+
def makeget_computer(clear_database, make_tempdir):
55+
"""Provide a function to generate new computers.
56+
57+
:return: The computer node
58+
:rtype: :py:class:`aiida.orm.Computer`
59+
"""
60+
61+
def _makeget_computer(label='localhost-test', transport_type='local'):
62+
from aiida.orm import Computer
63+
from aiida.common.exceptions import NotExistent
64+
65+
try:
66+
computer = Computer.objects.get(label=label)
67+
68+
except NotExistent:
69+
computer = Computer(
70+
label=label,
71+
description=f'{label} computer set up by test manager',
72+
hostname=label,
73+
workdir=make_tempdir(),
74+
transport_type=transport_type,
75+
scheduler_type='direct'
76+
)
77+
computer.store()
78+
computer.set_minimum_job_poll_interval(0.)
79+
80+
if transport_type == 'local':
81+
computer.configure()
82+
else:
83+
raise NotImplementedError('Transport `{transport_type}` not implemented')
84+
85+
return computer
86+
87+
return _makeget_computer

tests/utils/test_transfer.py

Lines changed: 140 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,140 @@
1+
# -*- coding: utf-8 -*-
2+
"""Unit tests for the :py:mod:`~aiida_quantumespresso.utils.transfer` module."""
3+
import pytest
4+
from packaging import version
5+
import aiida
6+
7+
from aiida.orm import FolderData, RemoteData
8+
from aiida.engine import ProcessBuilder
9+
from aiida_quantumespresso.utils.transfer import get_transfer_builder
10+
11+
pytestmark = pytest.mark.skipif(
12+
version.parse(aiida.get_version()) >= version.parse('1.6'),
13+
reason='Transfer was released in AiiDA core v1.6',
14+
)
15+
16+
17+
@pytest.fixture(name='make_data_source')
18+
def fixture_make_data_source(make_tempdir):
19+
"""Provide a function to generate data sources specific for this test."""
20+
21+
def _make_data_source(computer=None, populate_level=0):
22+
import os
23+
import io
24+
25+
if computer is None:
26+
node = FolderData()
27+
28+
filepaths = []
29+
if populate_level > 0:
30+
filepaths.append('data-file-schema.xml')
31+
filepaths.append('charge-density.dat')
32+
if populate_level > 1:
33+
filepaths.append('paw.txt')
34+
35+
for filepath in filepaths:
36+
node.put_object_from_filelike(io.StringIO(), path=filepath)
37+
38+
else:
39+
remote_path = make_tempdir(parent_dir=computer.get_workdir())
40+
node = RemoteData(computer=computer, remote_path=remote_path)
41+
42+
filepaths = []
43+
if populate_level > 0:
44+
filepaths.append('out/aiida.save/data-file-schema.xml')
45+
filepaths.append('out/aiida.save/charge-density.dat')
46+
if populate_level > 1:
47+
filepaths.append('out/aiida.save/paw.txt')
48+
49+
for filepath in filepaths:
50+
fullpath = os.path.join(remote_path, filepath)
51+
os.makedirs(os.path.dirname(fullpath), exist_ok=True)
52+
open(fullpath, 'w').close()
53+
54+
return node.store()
55+
56+
return _make_data_source
57+
58+
59+
@pytest.mark.parametrize('track', [True, False])
60+
@pytest.mark.parametrize('populate_level', [1, 2])
61+
@pytest.mark.parametrize('source_is_remote, provide_computer', [(True, True), (True, False), (False, True)])
62+
def test_transfer_builder_results(
63+
makeget_computer, make_data_source, source_is_remote, provide_computer, populate_level, track
64+
):
65+
"""Test all viable input sets for the get_transfer_builder utility function."""
66+
67+
if provide_computer:
68+
builder_computer = makeget_computer('computer1')
69+
check_computer = builder_computer
70+
else:
71+
builder_computer = None
72+
73+
if source_is_remote:
74+
folder_computer = makeget_computer('computer2')
75+
data_source = make_data_source(computer=folder_computer, populate_level=populate_level)
76+
retrieve_expected = True
77+
expected_listname = 'symlink_files'
78+
check_computer = folder_computer
79+
else:
80+
data_source = make_data_source(computer=None, populate_level=populate_level)
81+
retrieve_expected = False
82+
expected_listname = 'local_files'
83+
84+
if source_is_remote and provide_computer:
85+
with pytest.warns(UserWarning) as warnings:
86+
builder = get_transfer_builder(data_source, computer=builder_computer, track=track)
87+
assert len(warnings) == 1
88+
assert 'ignore' in str(warnings[0].message)
89+
assert f'{builder_computer}' in str(warnings[0].message)
90+
assert f'{folder_computer}' in str(warnings[0].message)
91+
# Makes sure the information is there, but there is no generic way of checking the correctedness
92+
# of the warning (i.e. which computer is the one selected) without checking for the exact text
93+
else:
94+
builder = get_transfer_builder(data_source, computer=builder_computer, track=track)
95+
96+
assert isinstance(builder, ProcessBuilder)
97+
assert builder.metadata['computer'].pk == check_computer.pk
98+
assert builder.source_nodes['source_node'] == data_source
99+
assert builder.instructions.get_dict()['retrieve_files'] == retrieve_expected
100+
assert expected_listname in builder.instructions.get_dict()
101+
102+
if source_is_remote:
103+
# paw.txt are always loaded when remote copying
104+
expected_setlist = {
105+
('source_node', 'out/aiida.save/data-file-schema.xml', 'data-file-schema.xml'),
106+
('source_node', 'out/aiida.save/charge-density.dat', 'charge-density.dat'),
107+
('source_node', 'out/aiida.save/paw.txt', 'paw.txt'),
108+
}
109+
110+
else:
111+
expected_setlist = {
112+
('source_node', 'data-file-schema.xml', 'out/aiida.save/data-file-schema.xml'),
113+
('source_node', 'charge-density.dat', 'out/aiida.save/charge-density.dat'),
114+
}
115+
if populate_level == 2:
116+
expected_setlist.add(('source_node', 'paw.txt', 'out/aiida.save/paw.txt'))
117+
118+
# Note: I need to do the following transformation manually because sometimes what I get from:
119+
#
120+
# > builder.instructions.get_dict()[expected_listname]
121+
#
122+
# is a list of (unhashable) lists, and other times it returns a list of tupples
123+
# This needs to be checked in aiida-core before changing the following to a simpler syntax
124+
obtained_setlist = set()
125+
for element in builder.instructions.get_dict()[expected_listname]:
126+
obtained_setlist.add(tuple(element))
127+
assert obtained_setlist == expected_setlist
128+
129+
130+
@pytest.mark.parametrize('track', [True, False])
131+
@pytest.mark.parametrize('populate_level', [1, 2])
132+
def test_transfer_builder_raise(make_data_source, populate_level, track):
133+
"""Test all raises for the get_transfer_builder utility function."""
134+
135+
data_source = make_data_source(computer=None, populate_level=populate_level)
136+
137+
with pytest.raises(ValueError) as execinfo:
138+
_ = get_transfer_builder(data_source, computer=None, track=track)
139+
140+
assert 'computer' in str(execinfo.value)

0 commit comments

Comments
 (0)