Skip to content

Commit 92dd950

Browse files
committed
Add test for metadata
1 parent ce1d95d commit 92dd950

3 files changed

Lines changed: 143 additions & 2 deletions

File tree

‎tests/conftest.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ def pytest_collection_modifyitems(items):
3737
"tests.data.test_interface",
3838
"tests.data.test_params",
3939
"tests.data.test_source",
40+
"tests.data.test_writer",
4041
"tests.data.test_store",
4142
"tests.data.test_xarray",
4243
]

‎tests/data/test_writer.py‎

Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
"""Test common Writer features."""
2+
3+
import os
4+
5+
from neba.data import DataInterface, ParametersDict
6+
from neba.data.writer import MetadataGenerator, WriterAbstract, element
7+
8+
9+
class TestMetadata:
10+
ELTS_BASIC = [
11+
"written_with_interface",
12+
"creation_time",
13+
"creation_hostname",
14+
"creation_script",
15+
]
16+
ELTS_PARAMS = ["creation_params"]
17+
ELTS_GIT = ["creation_commit", "creation_diff"]
18+
19+
def get_interface(self, *args, **kwargs) -> DataInterface:
20+
class MyDataInterface(DataInterface):
21+
Parameters = ParametersDict
22+
23+
return MyDataInterface(*args, **kwargs)
24+
25+
def test_elements_selection(self):
26+
27+
di = self.get_interface()
28+
metadata = di.writer.get_metadata()
29+
30+
# git will fail on github CI for some reason
31+
for key in self.ELTS_BASIC + self.ELTS_PARAMS:
32+
assert key in metadata
33+
34+
# Test group skip
35+
assert (
36+
di.writer.metadata_generator(None, add_git_info=False).get_elements()
37+
== self.ELTS_BASIC + self.ELTS_PARAMS
38+
)
39+
40+
assert (
41+
di.writer.metadata_generator(None, add_params=False).get_elements()
42+
== self.ELTS_BASIC + self.ELTS_GIT
43+
)
44+
45+
metadata = di.writer.get_metadata(add_params=False, add_git_info=False)
46+
assert list(metadata.keys()) == self.ELTS_BASIC
47+
48+
# Test manually specifying elements (in different order)
49+
elts = ["creation_hostname", "written_with_interface", "creation_params"]
50+
metadata = di.writer.get_metadata(elements=elts)
51+
assert list(metadata.keys()) == elts
52+
53+
def test_generator_subclass(self):
54+
class MyGenerator(MetadataGenerator):
55+
@element
56+
def simple(self):
57+
return 0
58+
59+
@element(elements=["a", "b"])
60+
def multiple(self):
61+
return {"a": 0, "b": 1}
62+
63+
@element(elements=["b"])
64+
def last(self):
65+
return {"b": 5}
66+
67+
class MyDataInterface(DataInterface):
68+
Parameters = ParametersDict
69+
70+
class Writer(WriterAbstract):
71+
metadata_generator = MyGenerator
72+
73+
di = MyDataInterface()
74+
metadata = di.writer.get_metadata(add_git_info=False)
75+
76+
assert list(metadata.keys()) == self.ELTS_BASIC + self.ELTS_PARAMS + [
77+
"simple",
78+
"a",
79+
"b",
80+
]
81+
82+
assert metadata["simple"] == 0
83+
assert metadata["a"] == 0
84+
assert metadata["b"] == 5
85+
86+
def test_renaming(self):
87+
class MyGenerator(MetadataGenerator):
88+
@element
89+
def simple(self):
90+
return 0
91+
92+
@element(elements=["a", "b"])
93+
def multiple(self):
94+
return {"a": 0, "b": 1}
95+
96+
class MyDataInterface(DataInterface):
97+
Parameters = ParametersDict
98+
99+
class Writer(WriterAbstract):
100+
metadata_generator = MyGenerator
101+
102+
di = MyDataInterface()
103+
MyGenerator.simple.rename("simple_rename")
104+
MyGenerator.multiple.rename(a="a_rename")
105+
metadata = di.writer.get_metadata()
106+
assert metadata["simple_rename"] == 0
107+
assert metadata["a_rename"] == 0
108+
assert metadata["b"] == 1
109+
assert "a" not in metadata
110+
assert "simple" not in metadata
111+
112+
def test_parameters(self):
113+
params = dict(a=0, b=1)
114+
di = self.get_interface(params)
115+
metadata = di.writer.get_metadata(params_str=False)
116+
assert metadata["creation_params"] == params
117+
118+
metadata = di.writer.get_metadata()
119+
assert metadata["creation_params"] == '{"a": 0, "b": 1}'
120+
121+
def test_script_filename(self):
122+
di = self.get_interface()
123+
metadata = di.writer.get_metadata(
124+
creation_script="a", elements=["creation_script"]
125+
)
126+
assert metadata["creation_script"] == "a"
127+
128+
# Pytest messes with getting filename (similarly to IPython)
129+
# not sure how to test it properly
130+
131+
def test_git(self):
132+
di = self.get_interface()
133+
metadata = di.writer.get_metadata(creation_script=".")
134+
assert "creation_commit" in metadata
135+
136+
if (commit := os.environ.get("GITHUB_SHA")) is not None:
137+
assert metadata["creation_commit"] == commit

‎tests/data/test_xarray.py‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -161,8 +161,11 @@ def test_metadata(self, tmpdir):
161161

162162
written = xr.open_dataset(filename)
163163

164-
assert written.attrs["written_with_interface"] == "XarrayInterface"
165-
assert written.attrs["created_with_params"] == '{"a": 0}'
164+
assert (
165+
written.attrs["written_with_interface"]
166+
== "tests.data.test_xarray.XarrayInterface"
167+
)
168+
assert written.attrs["creation_params"] == '{"a": 0}'
166169

167170
def setup_multifile(self, tmpdir) -> tuple[xr.Dataset, list[xr.Dataset], list[str]]:
168171
data = np.arange(12).reshape(3, 4)

0 commit comments

Comments
 (0)