Coverage for tests/test_pre_transform.py: 100%
184 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-08-29 09:18 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-08-29 09:18 +0000
1# This file is part of ctrl_bps.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (https://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This software is dual licensed under the GNU General Public License and also
10# under a 3-clause BSD license. Recipients may choose which of these licenses
11# to use; please see the files gpl-3.0.txt and/or bsd_license.txt,
12# respectively. If you choose the GPL option then the following text applies
13# (but note that there is still no warranty even if you opt for BSD instead):
14#
15# This program is free software: you can redistribute it and/or modify
16# it under the terms of the GNU General Public License as published by
17# the Free Software Foundation, either version 3 of the License, or
18# (at your option) any later version.
19#
20# This program is distributed in the hope that it will be useful,
21# but WITHOUT ANY WARRANTY; without even the implied warranty of
22# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
23# GNU General Public License for more details.
24#
25# You should have received a copy of the GNU General Public License
26# along with this program. If not, see <https://www.gnu.org/licenses/>.
27import errno
28import logging
29import os
30import shutil
31import sys
32import tempfile
33import unittest
34from pathlib import Path
36from qg_test_utils import make_test_quantum_graph
38from lsst.ctrl.bps import BpsConfig, BpsSubprocessError, ClusteredQuantumGraph, pre_transform
39from lsst.pipe.base import QuantumGraph
40from lsst.pipe.base.quantum_graph import PredictedQuantumGraph
41from lsst.pipe.base.tests.mocks import InMemoryRepo
43TESTDIR = os.path.abspath(os.path.dirname(__file__))
44_LOG = logging.getLogger(__name__)
47class TestExecute(unittest.TestCase):
48 """Test execution."""
50 def setUp(self):
51 self.file = tempfile.NamedTemporaryFile("w+")
52 self.logger = logging.getLogger("lsst.ctrl.bps")
54 def tearDown(self):
55 self.file.close()
57 def testSuccessfulExecution(self):
58 """Test exit status if command succeeded."""
59 content = "Successful execution"
60 command = f"{sys.executable} -c 'print(\"{content}\")'"
61 with self.assertLogs(logger=self.logger, level="INFO") as cm:
62 status = pre_transform.execute(command, self.file.name)
63 self.assertIn(content, cm.output[0])
64 self.file.seek(0)
65 file_contents = self.file.read()
66 self.assertIn(command, file_contents)
67 self.assertIn(content, file_contents)
68 self.assertEqual(status, 0)
70 def testFailingExecution(self):
71 """Test exit status if command failed."""
72 status = pre_transform.execute("false", self.file.name)
73 self.assertIn("false", self.file.read())
74 self.assertNotEqual(status, 0)
77class TestCreatingQuantumGraph(unittest.TestCase):
78 """Test quantum graph creation."""
80 def setUp(self):
81 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR)
82 self.settings = {
83 "createQuantumGraph": "touch {qgraphFile}",
84 "submitPath": self.tmpdir,
85 "whenSaveJobQgraph": "NEVER",
86 "uniqProcName": "my_test",
87 "qgraphFileTemplate": "{uniqProcName}.qg",
88 }
89 self.logger = logging.getLogger("lsst.ctrl.bps")
91 def tearDown(self):
92 shutil.rmtree(self.tmpdir, ignore_errors=True)
94 def testSuccess(self):
95 """Test if a new quantum graph was created successfully."""
96 config = BpsConfig(self.settings, search_order=[])
97 with self.assertLogs(logger=self.logger, level="INFO") as cm:
98 qgraph_filename = pre_transform.create_quantum_graph(config, self.tmpdir)
99 _, command = config.search("createQuantumGraph", opt={"curvals": {"qgraphFile": qgraph_filename}})
100 self.assertIn(command, cm.output[0])
101 self.assertTrue(os.path.exists(qgraph_filename))
103 def testCommandMissing(self):
104 """Test if error is caught when the command is missing."""
105 del self.settings["createQuantumGraph"]
106 config = BpsConfig(self.settings, search_order=[])
107 with self.assertRaisesRegex(KeyError, "command.*not found"):
108 pre_transform.create_quantum_graph(config, self.tmpdir)
110 def testFailure(self):
111 """Test if error is caught when the quantum graph creation fails."""
112 self.settings["createQuantumGraph"] = "bash -c 'exit 2'"
113 config = BpsConfig(self.settings, search_order=[])
114 with self.assertRaises(BpsSubprocessError) as cm:
115 pre_transform.create_quantum_graph(config, self.tmpdir)
116 self.assertEqual(cm.exception.errno, errno.ENOENT)
117 self.assertIn("non-zero exit code", str(cm.exception))
120class TestUpdatingQuantumGraph(unittest.TestCase):
121 """Test quantum graph update."""
123 def setUp(self):
124 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR)
125 self.settings = {
126 "updateQuantumGraph": "bash -c 'echo foo > {qgraphFile}'",
127 "submitPath": self.tmpdir,
128 "whenSaveJobQgraph": "NEVER",
129 "uniqProcName": "my_test",
130 "qgraphFileTemplate": "{uniqProcName}.qg",
131 "inputQgraphFile": f"{self.tmpdir}/src.qg",
132 }
133 self.logger = logging.getLogger("lsst.ctrl.bps")
135 # Create a file in the temporary directory that will serve as
136 # the file with a quantum graph that needs updating.
137 self.src = Path(self.settings["inputQgraphFile"])
138 self.src.write_text("foo\n")
140 self.backup = Path(f"{self.src.parent}/{self.src.stem}_orig{self.src.suffix}")
142 def tearDown(self):
143 shutil.rmtree(self.tmpdir, ignore_errors=True)
145 def testSuccess(self):
146 """Test if the quantum graph was updated."""
147 config = BpsConfig(self.settings, search_order=[])
148 with self.assertLogs(logger=self.logger, level="INFO") as cm:
149 pre_transform.update_quantum_graph(config, str(self.src), self.tmpdir)
150 _, command = config.search("updateQuantumGraph", opt={"curvals": {"qgraphFile": str(self.src)}})
151 self.assertIn("backing up", cm.output[0].lower())
152 self.assertIn("completed", cm.output[1].lower())
153 self.assertIn(command, cm.output[2])
154 self.assertTrue(self.src.read_text(), "bar\n")
155 self.assertTrue(self.backup.is_file())
156 self.assertTrue(self.backup.read_text(), "foo\n")
158 def testSuccessInPlace(self):
159 """Test if a quantum graph was updated inplace."""
160 config = BpsConfig(self.settings, search_order=[])
161 with self.assertLogs(logger=self.logger, level="INFO") as cm:
162 pre_transform.update_quantum_graph(config, str(self.src), self.tmpdir, inplace=True)
163 _, command = config.search("updateQuantumGraph", opt={"curvals": {"qgraphFile": str(self.src)}})
164 self.assertIn(command, cm.output[0])
165 self.assertTrue(self.src.read_text(), "bar\n")
166 self.assertFalse(self.backup.is_file())
168 def testCommandMissing(self):
169 """Test if error is caught when the command is missing."""
170 del self.settings["updateQuantumGraph"]
171 config = BpsConfig(self.settings, search_order=[])
172 with self.assertRaisesRegex(KeyError, "command.*not found"):
173 pre_transform.update_quantum_graph(config, str(self.src), self.tmpdir)
175 def testFailure(self):
176 """Test if error is caught when the command fails."""
177 self.settings["updateQuantumGraph"] = "bash -c 'exit 2'"
178 config = BpsConfig(self.settings, search_order=[])
179 with self.assertRaises(BpsSubprocessError) as cm:
180 pre_transform.update_quantum_graph(config, str(self.src), self.tmpdir)
181 self.assertEqual(cm.exception.errno, errno.ENOENT)
182 self.assertRegex(str(cm.exception), "non-zero exit code")
185class TestReadQuantumGraph(unittest.TestCase):
186 """Test read_quantum_graph method."""
188 def setUp(self):
189 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR)
190 self.logger = logging.getLogger("lsst.ctrl.bps")
192 def tearDown(self):
193 shutil.rmtree(self.tmpdir, ignore_errors=True)
195 def testBadExtension(self):
196 with self.assertRaisesRegex(ValueError, "Unrecognized extension for quantum graph file"):
197 _ = pre_transform.read_quantum_graph("mygraph.badext")
199 def testReadQG(self):
200 filename = f"{self.tmpdir}/testQG.qg"
201 self.qg = make_test_quantum_graph("test_read", save_filename=filename)
202 self.assertTrue(os.path.exists(filename))
204 qg = pre_transform.read_quantum_graph(filename)
205 self.assertEqual(len(qg), 18)
207 @unittest.mock.patch.object(QuantumGraph, "loadUri")
208 @unittest.mock.patch.object(PredictedQuantumGraph, "from_old_quantum_graph")
209 def testReadOldQG(self, mock_from, mock_load):
210 # Instead of maintaining old format in ctrl_bps, just make sure bps
211 # function calls methods to read old format and convert old format.
212 # This test will fail if from_old_quantum_graph or QuantumGraph are
213 # removed. Those calls should also be removed from read_quantum_graph.
214 mock_from.return_value = "quantum_graph"
215 mock_load.return_value = "quantum_graph"
216 filename = f"{self.tmpdir}/testQG.qgraph"
217 _ = pre_transform.read_quantum_graph(filename)
218 mock_load.assert_called_once()
219 mock_from.assert_called_once()
222class TestAcquireQuantumGraph(unittest.TestCase):
223 """Test acquire_quantum_graph method."""
225 def setUp(self):
226 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR)
227 self.logger = logging.getLogger("lsst.ctrl.bps")
229 def tearDown(self):
230 shutil.rmtree(self.tmpdir, ignore_errors=True)
232 @unittest.mock.patch("lsst.ctrl.bps.pre_transform.create_quantum_graph")
233 def testCreateGraph(self, mock_create):
234 qgraph_filename = f"{self.tmpdir}/created.qg"
235 mock_create.return_value = qgraph_filename
236 results = pre_transform.acquire_quantum_graph(BpsConfig({}), out_prefix=self.tmpdir)
237 self.assertEqual(results, qgraph_filename)
238 mock_create.assert_called_once()
240 @unittest.mock.patch("lsst.ctrl.bps.pre_transform.update_quantum_graph")
241 def testExistingGraphNoCopy(self, mock_update):
242 filename = "original.qg"
243 config = BpsConfig({"qgraphFile": filename, "finalJob": {"dummy_var": "dummy_val"}})
245 results = pre_transform.acquire_quantum_graph(config, out_prefix=None)
246 self.assertEqual(results, filename)
247 mock_update.assert_called_once()
249 @unittest.mock.patch("lsst.ctrl.bps.pre_transform.update_quantum_graph")
250 def testExistingGraphNoCopyNoUpdate(self, mock_update):
251 filename = "original.qg"
252 config = BpsConfig({"qgraphFile": filename})
254 results = pre_transform.acquire_quantum_graph(config, out_prefix=None)
255 self.assertEqual(results, filename)
256 mock_update.assert_not_called()
258 @unittest.mock.patch("lsst.ctrl.bps.pre_transform.update_quantum_graph")
259 def testExistingGraphCopy(self, mock_update):
260 filename = Path(self.tmpdir) / "original.qg"
261 with open(filename, "w") as fh:
262 fh.write("test file")
263 path = Path(self.tmpdir) / "run_dir"
264 path.mkdir(parents=True, exist_ok=True)
266 config = BpsConfig({"qgraphFile": str(filename), "finalJob": {"dummy_var": "dummy_val"}})
268 results = pre_transform.acquire_quantum_graph(config, out_prefix=path)
269 self.assertEqual(results, str(path / filename.name))
270 self.assertTrue(Path(results).exists())
271 mock_update.assert_called_once()
274class TestClusterQuanta(unittest.TestCase):
275 """Test cluster_quanta method. Other tests cover functions
276 cluster_quanta calls so mocking them here.
277 """
279 @unittest.mock.patch.object(ClusteredQuantumGraph, "validate")
280 def testValidate(self, mock_validate):
281 """Test that actually calls validate per config."""
282 mock_validate.side_effect = RuntimeError("Fake error")
283 settings = {
284 "clusterAlgorithm": "lsst.ctrl.bps.quantum_clustering_funcs.single_quantum_clustering",
285 "uniqProcName": "my_test",
286 "validateClusteredQgraph": True,
287 }
288 config = BpsConfig(settings, search_order=[])
289 with InMemoryRepo() as repo:
290 qgraph = repo.make_quantum_graph()
291 with self.assertRaisesRegex(RuntimeError, "Fake error"):
292 _ = pre_transform.cluster_quanta(config, qgraph, "a_name")
294 @unittest.mock.patch.object(ClusteredQuantumGraph, "validate")
295 def testNoValidate(self, mock_validate):
296 """Test that doesn't call validate per config."""
297 mock_validate.side_effect = RuntimeError("Fake error")
298 settings = {
299 "clusterAlgorithm": "lsst.ctrl.bps.quantum_clustering_funcs.single_quantum_clustering",
300 "uniqProcName": "my_test",
301 "validateClusteredQgraph": False,
302 }
303 config = BpsConfig(settings, search_order=[])
304 with InMemoryRepo() as repo:
305 qgraph = repo.make_quantum_graph()
306 _ = pre_transform.cluster_quanta(config, qgraph, "a_name")
309if __name__ == "__main__":
310 unittest.main()