Coverage for tests/test_pre_transform.py: 100%

184 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-02 05:16 -0400

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 

35 

36from qg_test_utils import make_test_quantum_graph 

37 

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 

42 

43TESTDIR = os.path.abspath(os.path.dirname(__file__)) 

44_LOG = logging.getLogger(__name__) 

45 

46 

47class TestExecute(unittest.TestCase): 

48 """Test execution.""" 

49 

50 def setUp(self): 

51 self.file = tempfile.NamedTemporaryFile("w+") 

52 self.logger = logging.getLogger("lsst.ctrl.bps") 

53 

54 def tearDown(self): 

55 self.file.close() 

56 

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) 

69 

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) 

75 

76 

77class TestCreatingQuantumGraph(unittest.TestCase): 

78 """Test quantum graph creation.""" 

79 

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") 

90 

91 def tearDown(self): 

92 shutil.rmtree(self.tmpdir, ignore_errors=True) 

93 

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)) 

102 

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) 

109 

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)) 

118 

119 

120class TestUpdatingQuantumGraph(unittest.TestCase): 

121 """Test quantum graph update.""" 

122 

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") 

134 

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") 

139 

140 self.backup = Path(f"{self.src.parent}/{self.src.stem}_orig{self.src.suffix}") 

141 

142 def tearDown(self): 

143 shutil.rmtree(self.tmpdir, ignore_errors=True) 

144 

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") 

157 

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()) 

167 

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) 

174 

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") 

183 

184 

185class TestReadQuantumGraph(unittest.TestCase): 

186 """Test read_quantum_graph method.""" 

187 

188 def setUp(self): 

189 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR) 

190 self.logger = logging.getLogger("lsst.ctrl.bps") 

191 

192 def tearDown(self): 

193 shutil.rmtree(self.tmpdir, ignore_errors=True) 

194 

195 def testBadExtension(self): 

196 with self.assertRaisesRegex(ValueError, "Unrecognized extension for quantum graph file"): 

197 _ = pre_transform.read_quantum_graph("mygraph.badext") 

198 

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)) 

203 

204 qg = pre_transform.read_quantum_graph(filename) 

205 self.assertEqual(len(qg), 18) 

206 

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() 

220 

221 

222class TestAcquireQuantumGraph(unittest.TestCase): 

223 """Test acquire_quantum_graph method.""" 

224 

225 def setUp(self): 

226 self.tmpdir = tempfile.mkdtemp(dir=TESTDIR) 

227 self.logger = logging.getLogger("lsst.ctrl.bps") 

228 

229 def tearDown(self): 

230 shutil.rmtree(self.tmpdir, ignore_errors=True) 

231 

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() 

239 

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"}}) 

244 

245 results = pre_transform.acquire_quantum_graph(config, out_prefix=None) 

246 self.assertEqual(results, filename) 

247 mock_update.assert_called_once() 

248 

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}) 

253 

254 results = pre_transform.acquire_quantum_graph(config, out_prefix=None) 

255 self.assertEqual(results, filename) 

256 mock_update.assert_not_called() 

257 

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) 

265 

266 config = BpsConfig({"qgraphFile": str(filename), "finalJob": {"dummy_var": "dummy_val"}}) 

267 

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() 

272 

273 

274class TestClusterQuanta(unittest.TestCase): 

275 """Test cluster_quanta method. Other tests cover functions 

276 cluster_quanta calls so mocking them here. 

277 """ 

278 

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") 

293 

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") 

307 

308 

309if __name__ == "__main__": 

310 unittest.main()