Coverage for tests/test_aggregator.py: 99%

565 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-29 02:11 -0700

1# This file is part of pipe_base. 

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/>. 

27 

28from __future__ import annotations 

29 

30import dataclasses 

31import itertools 

32import json 

33import os 

34import re 

35import tempfile 

36import time 

37import types 

38import unittest.mock 

39import uuid 

40from collections.abc import Iterator 

41from contextlib import contextmanager 

42from typing import Any, cast 

43 

44import astropy.table 

45import click.testing 

46import numpy as np 

47import pydantic 

48from click.testing import CliRunner, Result 

49 

50import lsst.utils.tests 

51from lsst.daf.butler import Butler, ButlerLogRecords, QuantumBackedButler 

52from lsst.pex.config import Config 

53from lsst.pipe.base import ( 

54 AlgorithmError, 

55 QuantumAttemptStatus, 

56 QuantumSuccessCaveats, 

57 TaskMetadata, 

58) 

59from lsst.pipe.base import automatic_connection_constants as acc 

60from lsst.pipe.base.cli.cmd.commands import ( 

61 aggregate_graph as aggregate_graph_cli, 

62) 

63from lsst.pipe.base.cli.cmd.commands import ( 

64 provenance_report as provenance_report_cli, 

65) 

66from lsst.pipe.base.graph_walker import GraphWalker 

67from lsst.pipe.base.pipeline_graph import Edge 

68from lsst.pipe.base.quantum_graph import ( 

69 FORMAT_VERSION, 

70 PredictedDatasetInfo, 

71 PredictedQuantumGraph, 

72 PredictedQuantumInfo, 

73 ProvenanceDatasetInfo, 

74 ProvenanceQuantumGraph, 

75 ProvenanceQuantumInfo, 

76 ProvenanceQuantumReport, 

77 ProvenanceReport, 

78 ProvenanceTaskMetadataModel, 

79) 

80from lsst.pipe.base.quantum_graph.aggregator import AggregatorConfig, FatalWorkerError, aggregate_graph 

81from lsst.pipe.base.quantum_graph.aggregator._writer import Writer 

82from lsst.pipe.base.quantum_graph.ingest_graph import ingest_graph 

83from lsst.pipe.base.resource_usage import QuantumResourceUsage 

84from lsst.pipe.base.single_quantum_executor import SingleQuantumExecutor 

85from lsst.pipe.base.tests.mocks import ( 

86 DirectButlerRepo, 

87 DynamicConnectionConfig, 

88 DynamicTestPipelineTaskConfig, 

89) 

90from lsst.pipe.base.tests.util import patch_deterministic_uuid4 

91from lsst.utils.packages import Packages 

92 

93 

94@dataclasses.dataclass 

95class PrepInfo: 

96 """Struct of objects used in an aggregator test.""" 

97 

98 butler: Butler 

99 butler_path: str 

100 predicted: PredictedQuantumGraph 

101 predicted_path: str 

102 config: AggregatorConfig 

103 

104 

105class AggregatorTestCase(unittest.TestCase): 

106 """Unit tests for `lsst.pipe.base.quantum_graph.aggregator`.""" 

107 

108 @staticmethod 

109 @contextmanager 

110 def make_test_repo() -> Iterator[PrepInfo]: 

111 """Make a test data repository and predicted quantum graph. 

112 

113 Returns 

114 ------- 

115 prep_info : `PrepInfo` 

116 Objects used in aggregator tests. 

117 

118 Notes 

119 ----- 

120 The pipeline graph used by this task looks like this: 

121 

122 ■ calibrate: {detector, visit} 

123 ╭─┤ 

124 ■ │ consolidate: {visit} 

125 │ 

126 ■ resample: {patch, visit} 

127 │ 

128 ■ coadd: {band, patch} 

129 

130 The data can be visualized via:: 

131 

132 python -m lsst.daf.butler.tests.registry_data.spatial 

133 

134 One of the 'calibrate' quanta (visit=2, detector=2) is configured to 

135 fail with `lsst.pipe.base.AnnotatedPartialOutputsError`. This lets us 

136 test both success-with-caveats and failures, depending on how we 

137 configure the executor. This ``{visit: 2, detector: 2}`` data ID is 

138 the only one that overlaps ``{tract: 1, patch: 1}`` and 

139 ``{tract: 0, patch: 5}``, so it should chain to the 'resample' and 

140 'coadd' tasks, too. 

141 """ 

142 with patch_deterministic_uuid4(100): 

143 with DirectButlerRepo.make_temporary("base.yaml", "spatial.yaml") as (helper, root): 

144 calibrate_config = DynamicTestPipelineTaskConfig() 

145 calibrate_config.fail_exception = "lsst.pipe.base.AnnotatedPartialOutputsError" 

146 calibrate_config.fail_condition = "visit=2 AND detector=2" 

147 helper.add_task( 

148 "calibrate", 

149 config=calibrate_config, 

150 dimensions=["visit", "detector"], 

151 inputs={ 

152 "input_image": DynamicConnectionConfig( 

153 dataset_type_name="raw", 

154 dimensions=["visit", "detector"], 

155 ) 

156 }, 

157 prerequisite_inputs={ 

158 "refcat": DynamicConnectionConfig( 

159 dataset_type_name="references", 

160 dimensions=["htm7"], 

161 multiple=True, 

162 ) 

163 }, 

164 init_outputs={ 

165 "output_schema": DynamicConnectionConfig( 

166 dataset_type_name="source_schema", 

167 ) 

168 }, 

169 outputs={ 

170 "output_image": DynamicConnectionConfig( 

171 dataset_type_name="image", 

172 dimensions=["visit", "detector"], 

173 ), 

174 "output_table": DynamicConnectionConfig( 

175 dataset_type_name="source_detector", 

176 dimensions=["visit", "detector"], 

177 ), 

178 }, 

179 ) 

180 helper.add_task( 

181 "consolidate", 

182 dimensions=["visit"], 

183 init_inputs={ 

184 "input_schema": DynamicConnectionConfig( 

185 dataset_type_name="source_schema", 

186 ) 

187 }, 

188 inputs={ 

189 "input_table": DynamicConnectionConfig( 

190 dataset_type_name="source_detector", 

191 dimensions=["visit", "detector"], 

192 multiple=True, 

193 ) 

194 }, 

195 outputs={ 

196 "output_table": DynamicConnectionConfig( 

197 dataset_type_name="source", 

198 dimensions=["visit"], 

199 ) 

200 }, 

201 ) 

202 helper.add_task( 

203 "resample", 

204 dimensions=["patch", "visit"], 

205 inputs={ 

206 "input_image": DynamicConnectionConfig( 

207 dataset_type_name="image", 

208 dimensions=["visit", "detector"], 

209 multiple=True, 

210 ) 

211 }, 

212 outputs={ 

213 "output_image": DynamicConnectionConfig( 

214 dataset_type_name="warp", 

215 dimensions=["patch", "visit"], 

216 ) 

217 }, 

218 ) 

219 helper.add_task( 

220 "coadd", 

221 dimensions=["patch", "band"], 

222 inputs={ 

223 "input_image": DynamicConnectionConfig( 

224 dataset_type_name="warp", 

225 dimensions=["patch", "visit"], 

226 multiple=True, 

227 ) 

228 }, 

229 outputs={ 

230 "output_image": DynamicConnectionConfig( 

231 dataset_type_name="coadd", 

232 dimensions=["patch", "band"], 

233 ), 

234 }, 

235 ) 

236 pqgc = helper.make_quantum_graph_builder().finish(output="out_chain") 

237 # We use the butler root for various QG files just because it's 

238 # a convenient temporary directory. 

239 predicted_path = os.path.join(root, "predicted.qg") 

240 pqgc.write(predicted_path) 

241 config = AggregatorConfig( 

242 output_path=os.path.join(root, "provenance.qg"), 

243 # Set these small to see logic paths that otherwise only 

244 # affect large graphs. 

245 ingest_batch_size=10, 

246 zstd_dict_size=256, 

247 zstd_dict_n_inputs=16, 

248 ) 

249 yield PrepInfo( 

250 butler=helper.butler.clone(collections="out_chain"), 

251 butler_path=root, 

252 predicted=pqgc.assemble(), 

253 predicted_path=predicted_path, 

254 config=config, 

255 ) 

256 

257 def iter_graph_execution( 

258 self, 

259 repo: str, 

260 qg: PredictedQuantumGraph, 

261 raise_on_partial_outputs: bool, 

262 is_retry: bool = False, 

263 ) -> Iterator[uuid.UUID]: 

264 """Return an iterator that executes and yields quanta one by one. 

265 

266 Parameters 

267 ---------- 

268 repo : `str` 

269 Butler repository path or alias. 

270 qg : `lsst.pipe.base.quantum_graph.PredictedQuantumGraph` 

271 Predicted quantum graph. Must have datastore records attached, 

272 since execution uses a quantum-backed butler. 

273 raise_on_partial_outputs : `bool` 

274 Whether to raise on `lsst.pipe.base.AnnotatedPartialOutputsError` 

275 or treat it as a success with caveats. 

276 is_retry : `bool`, optional 

277 If `True`, this is a retry attempt and hence some outputs may 

278 already be present; skip successes and reprocess failures. 

279 

280 Returns 

281 ------- 

282 quanta : `~collections.abc.Iterator` [`uuid.UUID`] 

283 An iterator over all executed quantum IDs (not blocked ones). 

284 """ 

285 qbb = qg.make_init_qbb(repo) 

286 self.enterContext(qbb) 

287 qg.init_output_run(qbb) 

288 sqe = SingleQuantumExecutor( 

289 limited_butler_factory=lambda quantum: QuantumBackedButler.initialize( 

290 repo, 

291 quantum, 

292 qg.pipeline_graph.universe, 

293 ), 

294 assume_no_existing_outputs=not is_retry, 

295 skip_existing=is_retry, 

296 clobber_outputs=is_retry, 

297 raise_on_partial_outputs=raise_on_partial_outputs, 

298 ) 

299 qg.build_execution_quanta() 

300 xgraph = qg.quantum_only_xgraph 

301 walker = GraphWalker[uuid.UUID](xgraph.copy()) 

302 for ready in walker: 

303 for quantum_id in ready: 

304 info = xgraph.nodes[quantum_id] 

305 try: 

306 sqe.execute(info["pipeline_node"], info["quantum"], quantum_id) 

307 except AlgorithmError: 

308 walker.fail(quantum_id) 

309 else: 

310 walker.finish(quantum_id) 

311 yield quantum_id 

312 

313 def check_provenance_graph( 

314 self, 

315 pred: PredictedQuantumGraph, 

316 butler: Butler, 

317 expect_failure: bool, 

318 start_time: float, 

319 expect_failures_retried: bool = False, 

320 ) -> ProvenanceQuantumGraph: 

321 """Run a batter of tests on a provenance quantum graph produced by 

322 scanning the graph created by `make_test_repo`. 

323 

324 Parameters 

325 ---------- 

326 pred: `lsst.pipe.base.quantum_graph.PredictedQuantumGraph` 

327 Predicted quantum graph. 

328 prov_reader : \ 

329 `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraphReader` 

330 Reader for the provenance quantum graph. 

331 butler : `lsst.daf.butler.Butler` 

332 Client for the data repository. 

333 expect_failure : `bool` 

334 Whether to expect one quantum of 'calibrate' to fail (`True`) or 

335 succeed without writing anything (`False`). 

336 start_time : `float` 

337 A POSIX timestamp that strictly precedes the start time of any 

338 quantum's execution. 

339 expect_failures_retried : `bool`, optional 

340 If `True`, expect an initial attempt with failures prior to the 

341 most recent attempt. 

342 

343 Returns 

344 ------- 

345 prov : `ProvenanceQuantumGraph` 

346 The full provenance quantum graph. 

347 """ 

348 prov: ProvenanceQuantumGraph = butler.get("run_provenance") 

349 self.maxDiff = None 

350 self.assertEqual(prov.header.version, FORMAT_VERSION) 

351 self.assertEqual( 

352 list(butler.collections.get_info(prov.header.output).children), 

353 [prov.header.output_run] 

354 + list(butler.collections.query(prov.header.inputs, flatten_chains=True)), 

355 ) 

356 self.assertEqual(pred.quanta_by_task.keys(), prov.quanta_by_task.keys()) 

357 for task_label in pred.quanta_by_task: 

358 self.assertEqual(pred.quanta_by_task[task_label], prov.quanta_by_task[task_label]) 

359 self.assertEqual(pred.datasets_by_type.keys() - {"packages"}, prov.datasets_by_type.keys()) 

360 for dataset_type_name in prov.datasets_by_type: 

361 self.assertEqual( 

362 pred.datasets_by_type[dataset_type_name], prov.datasets_by_type[dataset_type_name] 

363 ) 

364 self.assertEqual(prov.init_quanta.keys(), pred.quanta_by_task.keys()) 

365 for quantum_id in pred: 

366 # Check consistency between the predicted and provenance quantum 

367 # node attributes. 

368 pred_qinfo: PredictedQuantumInfo = pred.bipartite_xgraph.nodes[quantum_id] 

369 prov_qinfo: ProvenanceQuantumInfo = prov.bipartite_xgraph.nodes[quantum_id] 

370 self.assertEqual(pred_qinfo["task_label"], prov_qinfo["task_label"]) 

371 self.assertEqual(pred_qinfo["data_id"], prov_qinfo["data_id"]) 

372 msg = f"{pred_qinfo['task_label']}@{pred_qinfo['data_id']}" 

373 # Check consistency between the predicted and provenance dataset 

374 # node attributes and edges. Also gather existence information for 

375 # use later. 

376 existence: dict[str, list[bool]] = {} 

377 pipeline_edges: list[Edge] 

378 for dataset_id, _, pipeline_edges in pred.bipartite_xgraph.in_edges( 

379 quantum_id, data="pipeline_edges" 

380 ): 

381 self.assertTrue(prov.bipartite_xgraph.has_predecessor(quantum_id, dataset_id)) 

382 for edge in pipeline_edges: 

383 existence.setdefault(edge.connection_name, []).append( 

384 self.check_dataset(dataset_id, pred, prov, butler) 

385 ) 

386 for _, dataset_id, pipeline_edges in pred.bipartite_xgraph.out_edges( 

387 quantum_id, data="pipeline_edges" 

388 ): 

389 self.assertTrue(prov.bipartite_xgraph.has_successor(quantum_id, dataset_id)) 

390 for edge in pipeline_edges: 

391 existence.setdefault(edge.connection_name, []).append( 

392 self.check_dataset(dataset_id, pred, prov, butler) 

393 ) 

394 # Check quantum status and dataset existence against the known 

395 # structure of the graph and where failures/caveats occur. 

396 match (pred_qinfo["task_label"], dict(pred_qinfo["data_id"].required)): 

397 case "calibrate", {"visit": 2, "detector": 2}: 

398 # This is the quantum that can directly raise. 

399 self._expect_all_exist(existence["input_image"], msg=msg) 

400 self._expect_all_exist(existence["refcat"], msg=msg) 

401 self._expect_none_exist(existence["output_image"], msg=msg) 

402 self._expect_none_exist(existence["output_table"], msg=msg) 

403 if expect_failure: 

404 self._expect_failure(prov_qinfo, existence, msg=msg) 

405 else: 

406 self._expect_successful( 

407 prov_qinfo, 

408 existence, 

409 caveats=( 

410 QuantumSuccessCaveats.PARTIAL_OUTPUTS_ERROR 

411 | QuantumSuccessCaveats.ALL_OUTPUTS_MISSING 

412 | QuantumSuccessCaveats.ANY_OUTPUTS_MISSING 

413 ), 

414 exception_type="lsst.pipe.base.tests.mocks.MockAlgorithmError", 

415 msg=msg, 

416 ) 

417 if expect_failures_retried: 

418 self.assertEqual(len(prov_qinfo["attempts"]), 2) 

419 self.assertEqual( 

420 prov_qinfo["attempts"][0].exception.type_name, 

421 "lsst.pipe.base.tests.mocks.MockAlgorithmError", 

422 ) 

423 else: 

424 self.assertEqual(len(prov_qinfo["attempts"]), 1) 

425 case "consolidate", {"visit": 2}: 

426 # This quantum will succeed (with one predicted input 

427 # missing) or be blocked. 

428 self._expect_one_missing(existence["input_table"], msg=msg) 

429 if expect_failure: 

430 self._expect_blocked(prov_qinfo, existence, msg=msg) 

431 else: 

432 self._expect_successful(prov_qinfo, existence, msg=msg) 

433 self.assertEqual( 

434 len(prov_qinfo["attempts"]), expect_failures_retried or not expect_failure 

435 ) 

436 case ( 

437 "resample" | "coadd", 

438 {"tract": 1, "patch": 1} | {"tract": 0, "patch": 5}, 

439 ): 

440 # These quanta will be blocked by an upstream failure or do 

441 # chained caveats, since they won't have enough inputs to 

442 # run. 

443 if expect_failure: 

444 self._expect_blocked(prov_qinfo, existence, msg=msg) 

445 else: 

446 self._expect_successful( 

447 prov_qinfo, 

448 existence, 

449 caveats=( 

450 QuantumSuccessCaveats.ADJUST_QUANTUM_RAISED 

451 | QuantumSuccessCaveats.NO_WORK 

452 | QuantumSuccessCaveats.ALL_OUTPUTS_MISSING 

453 | QuantumSuccessCaveats.ANY_OUTPUTS_MISSING 

454 ), 

455 msg=msg, 

456 ) 

457 self.assertEqual( 

458 len(prov_qinfo["attempts"]), expect_failures_retried or not expect_failure 

459 ) 

460 case ( 

461 "resample", 

462 {"tract": 0, "patch": 4, "visit": 2} | {"tract": 1, "patch": 0, "visit": 2}, 

463 ): 

464 # This will succeed or be blocked, with one input missing 

465 # regardless. 

466 self._expect_one_missing(existence["input_image"], msg=msg) 

467 if expect_failure: 

468 self._expect_blocked(prov_qinfo, existence, msg=msg) 

469 else: 

470 self._expect_successful(prov_qinfo, existence, msg=msg) 

471 self.assertEqual( 

472 len(prov_qinfo["attempts"]), expect_failures_retried or not expect_failure 

473 ) 

474 case ( 

475 "coadd", 

476 {"tract": 0, "patch": 4, "band": "r"} | {"tract": 1, "patch": 0, "band": "r"}, 

477 ): 

478 # This will succeed with no inputs missing or be blocked 

479 # with one input missing. 

480 if expect_failure: 

481 self._expect_one_missing(existence["input_image"], msg=msg) 

482 self._expect_blocked(prov_qinfo, existence, msg=msg) 

483 else: 

484 self._expect_all_exist(existence["input_image"], msg=msg) 

485 self._expect_successful(prov_qinfo, existence, msg=msg) 

486 self.assertEqual( 

487 len(prov_qinfo["attempts"]), expect_failures_retried or not expect_failure 

488 ) 

489 case _: 

490 # All other quanta should succeed and have all inputs 

491 # present. 

492 for connection_name in prov_qinfo["pipeline_node"].inputs.keys(): 

493 self._expect_all_exist(existence[connection_name], msg=msg) 

494 self._expect_successful(prov_qinfo, existence, msg=msg) 

495 self.assertEqual(len(prov_qinfo["attempts"]), 1) 

496 self.check_metadata( 

497 quantum_id, 

498 prov, 

499 butler, 

500 expect_ingested=(prov_qinfo["status"] is QuantumAttemptStatus.SUCCESSFUL), 

501 ) 

502 self.check_log( 

503 quantum_id, 

504 prov, 

505 butler, 

506 expect_ingested=( 

507 prov_qinfo["status"] 

508 in ( 

509 QuantumAttemptStatus.SUCCESSFUL, 

510 QuantumAttemptStatus.FAILED, 

511 ) 

512 ), 

513 ) 

514 self.check_resource_usage_table(prov, expect_failure=expect_failure, start_time=start_time) 

515 self.check_packages(butler) 

516 self.check_configs(butler, prov) 

517 self.check_quantum_table(prov, expect_failure=expect_failure) 

518 self.check_exception_table(prov, expect_failure=expect_failure) 

519 self.check_report(prov, expect_failure=expect_failure) 

520 return prov 

521 

522 def _expect_all_exist(self, existence: list[bool], msg: str) -> None: 

523 self.assertTrue(all(existence), msg=msg) 

524 

525 def _expect_none_exist(self, existence: list[bool], msg: str) -> None: 

526 self.assertFalse(any(existence), msg=msg) 

527 

528 def _expect_one_missing(self, existence: list[bool], msg: str) -> None: 

529 self.assertEqual(existence.count(False), 1, msg=msg) 

530 

531 def _expect_successful( 

532 self, 

533 info: ProvenanceQuantumInfo, 

534 existence: dict[str, list[bool]], 

535 caveats: QuantumSuccessCaveats = QuantumSuccessCaveats.NO_CAVEATS, 

536 exception_type: str | None = None, 

537 *, 

538 msg: str, 

539 ) -> None: 

540 self.assertEqual(info["status"], QuantumAttemptStatus.SUCCESSFUL, msg=msg) 

541 self.assertEqual(info["caveats"], caveats, msg=msg) 

542 if exception_type is None: 

543 self.assertIsNone(info["exception"], msg=msg) 

544 else: 

545 assert info["exception"] is not None 

546 self.assertEqual(info["exception"].type_name, exception_type, msg=msg) 

547 self._expect_all_exist(existence[acc.LOG_OUTPUT_CONNECTION_NAME], msg=msg) 

548 self._expect_all_exist(existence[acc.METADATA_OUTPUT_CONNECTION_NAME], msg=msg) 

549 if not (caveats & QuantumSuccessCaveats.ANY_OUTPUTS_MISSING): 

550 for connection_name in info["pipeline_node"].outputs.keys(): 

551 self._expect_all_exist(existence[connection_name], msg=msg) 

552 if caveats & QuantumSuccessCaveats.ALL_OUTPUTS_MISSING: 

553 for connection_name in info["pipeline_node"].outputs.keys(): 

554 self._expect_none_exist(existence[connection_name], msg=msg) 

555 self.assertIsNotNone(info["resource_usage"], msg=msg) 

556 self.assertGreater(info["resource_usage"].total_time, 0, msg=msg) 

557 self.assertGreater(info["resource_usage"].memory, 0, msg=msg) 

558 

559 def _expect_failure( 

560 self, info: ProvenanceQuantumInfo, existence: dict[str, list[bool]], msg: str 

561 ) -> None: 

562 self.assertEqual(info["status"], QuantumAttemptStatus.FAILED, msg=msg) 

563 self.assertEqual(info["exception"].type_name, "lsst.pipe.base.tests.mocks.MockAlgorithmError") 

564 self._expect_all_exist(existence[acc.LOG_OUTPUT_CONNECTION_NAME], msg=msg) 

565 self._expect_none_exist(existence[acc.METADATA_OUTPUT_CONNECTION_NAME], msg=msg) 

566 for connection_name in info["pipeline_node"].outputs.keys(): 

567 self._expect_none_exist(existence[connection_name], msg=msg) 

568 

569 def _expect_blocked( 

570 self, 

571 info: ProvenanceQuantumInfo, 

572 existence: dict[str, list[bool]], 

573 msg: str, 

574 ) -> None: 

575 self.assertEqual(info["status"], QuantumAttemptStatus.BLOCKED, msg=msg) 

576 self.assertEqual(info["attempts"], []) 

577 self._expect_none_exist(existence[acc.LOG_OUTPUT_CONNECTION_NAME], msg=msg) 

578 self._expect_none_exist(existence[acc.METADATA_OUTPUT_CONNECTION_NAME], msg=msg) 

579 for connection_name in info["pipeline_node"].outputs.keys(): 

580 self._expect_none_exist(existence[connection_name], msg=msg) 

581 

582 def check_dataset( 

583 self, 

584 dataset_id: uuid.UUID, 

585 pred: PredictedQuantumGraph, 

586 prov: ProvenanceQuantumGraph, 

587 butler: Butler, 

588 ) -> bool: 

589 """Check a provenance dataset for consistency with its predicted 

590 counterpart. 

591 

592 Parameters 

593 ---------- 

594 dataset_id : `uuid.UUID` 

595 Unique ID for the dataset. 

596 pred: `lsst.pipe.base.quantum_graph.PredictedQuantumGraph` 

597 Predicted quantum graph. 

598 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

599 Provenance quantum graph. 

600 butler : `lsst.daf.butler.Butler` 

601 Client for the data repository. 

602 

603 Returns 

604 ------- 

605 exists : `bool` 

606 Whether the dataset was marked as existing in the provenance 

607 quantum graph. 

608 """ 

609 pred_info: PredictedDatasetInfo = pred.bipartite_xgraph.nodes[dataset_id] 

610 prov_info: ProvenanceDatasetInfo = prov.bipartite_xgraph.nodes[dataset_id] 

611 self.assertEqual(pred_info["dataset_type_name"], prov_info["dataset_type_name"]) 

612 self.assertEqual(pred_info["data_id"], prov_info["data_id"]) 

613 self.assertEqual(pred_info["run"], prov_info["run"]) 

614 exists = prov_info["produced"] 

615 dataset_type_name = prov_info["dataset_type_name"] 

616 # We can remove this guard when we ingest QG-backed metadata and logs. 

617 if not dataset_type_name.endswith("_metadata") and not dataset_type_name.endswith("_log"): 

618 self.assertEqual( 

619 butler.get_dataset(dataset_id) is not None, 

620 exists, 

621 msg=( 

622 f"Ingest/existence inconsistency for {dataset_type_name}" 

623 f"@{prov_info['data_id']}/{dataset_id}]" 

624 ), 

625 ) 

626 return exists 

627 

628 def check_metadata( 

629 self, quantum_id: uuid.UUID, prov: ProvenanceQuantumGraph, butler: Butler, expect_ingested: bool 

630 ) -> None: 

631 """Check reading a metadata dataset from the butler, and check that the 

632 original metadata file has been deleted. 

633 

634 Parameters 

635 ---------- 

636 quantum_id : `uuid.UUID` 

637 Unique ID for the quantum this metadata belongs to. 

638 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

639 Provenance quantum graph. 

640 butler : `lsst.daf.butler.Butler` 

641 Client for the data repository. 

642 expect_ingested : `bool` 

643 Whether the metadata dataset should have been ingested. 

644 """ 

645 dataset_id = prov.bipartite_xgraph.nodes[quantum_id]["metadata_id"] 

646 ref = butler.get_dataset(dataset_id) 

647 if not expect_ingested: 

648 self.assertIsNone(ref) 

649 return 

650 assert ref is not None 

651 metadata = butler.get(ref) 

652 self.assertIsInstance(metadata, TaskMetadata) 

653 graph_path = butler.getURI("run_provenance") 

654 self.assertEqual(butler.getURI(ref), graph_path) 

655 # We now delete the metadata dataset, in order let us get the original 

656 # location from the butler and check that there's nothing there. Note 

657 # that this doesn't actually delete the file because the butler knows 

658 # it's shared with other datasets. 

659 butler.pruneDatasets([ref], disassociate=True, unstore=True, purge=True) 

660 original_path = butler.getURI(ref, predict=True) 

661 self.assertTrue(graph_path.exists()) 

662 self.assertNotEqual(graph_path, original_path) 

663 self.assertFalse(original_path.exists()) 

664 

665 def check_log( 

666 self, quantum_id: uuid.UUID, prov: ProvenanceQuantumGraph, butler: Butler, expect_ingested: bool 

667 ) -> None: 

668 """Check reading a log dataset from the butler, and check that the 

669 original log file has been deleted. 

670 

671 Parameters 

672 ---------- 

673 quantum_id : `uuid.UUID` 

674 Unique ID for the quantum this log belongs to. 

675 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

676 Provenance quantum graph. 

677 butler : `lsst.daf.butler.Butler` 

678 Client for the data repository. 

679 expect_ingested : `bool` 

680 Whether the metadata dataset should have been ingested. 

681 """ 

682 dataset_id = prov.bipartite_xgraph.nodes[quantum_id]["log_id"] 

683 ref = butler.get_dataset(dataset_id) 

684 if not expect_ingested: 

685 self.assertIsNone(ref) 

686 return 

687 assert ref is not None 

688 log = butler.get(ref) 

689 self.assertIsInstance(log, ButlerLogRecords) 

690 graph_path = butler.getURI("run_provenance") 

691 self.assertEqual(butler.getURI(ref), graph_path) 

692 # We now delete the log dataset, in order let us get the original 

693 # location from the butler and check that there's nothing there. Note 

694 # that this doesn't actually delete the file because the butler knows 

695 # it's shared with other datasets. 

696 butler.pruneDatasets([ref], disassociate=True, unstore=True, purge=True) 

697 original_path = butler.getURI(ref, predict=True) 

698 self.assertTrue(graph_path.exists()) 

699 self.assertNotEqual(graph_path, original_path) 

700 self.assertFalse(original_path.exists()) 

701 

702 def check_configs(self, butler: Butler, prov: ProvenanceQuantumGraph) -> None: 

703 for task_node in prov.pipeline_graph.tasks.values(): 

704 config = butler.get(task_node.init.config_output.dataset_type_name) 

705 self.assertIsInstance(config, Config) 

706 

707 def check_packages(self, butler: Butler) -> None: 

708 """Check fetching package versions from the provenance graph. 

709 

710 Parameters 

711 ---------- 

712 butler : `lsst.daf.butler.Butler` 

713 Client for the data repository. 

714 """ 

715 packages = butler.get("run_provenance.packages") 

716 self.assertIsInstance(packages, Packages) 

717 self.assertIn("pipe_base", packages) 

718 

719 def check_resource_usage_table( 

720 self, prov: ProvenanceQuantumGraph, expect_failure: bool, start_time: float 

721 ) -> None: 

722 """Check building a resource usage table from the provenance graph. 

723 

724 Parameters 

725 ---------- 

726 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

727 Reader for the provenance quantum graph. 

728 expect_failure : `bool` 

729 Whether to expect one quantum of 'calibrate' to fail (`True`) or 

730 succeed without writing anything (`False`). 

731 start_time : `float` 

732 A POSIX timestamp that strictly precedes the start time of any 

733 quantum's execution. 

734 """ 

735 tbl = prov.make_task_resource_usage_table("calibrate", include_data_ids=True) 

736 self.assertEqual(len(tbl), prov.header.n_task_quanta["calibrate"]) 

737 self.assertCountEqual( 

738 tbl.colnames, 

739 ["quantum_id"] 

740 + list(prov.pipeline_graph.tasks["calibrate"].dimensions.names) 

741 + list(QuantumResourceUsage.model_fields), 

742 ) 

743 # Check that quantum start times are bounded by the before-execution 

744 # start_time and now. This makes sure we didn't get any timezone 

745 # shenanigans. 

746 end_time = time.time() 

747 for quantum_start_time in tbl["start"]: 

748 self.assertGreater(quantum_start_time, start_time) 

749 self.assertLess(quantum_start_time, end_time) 

750 self.assertTrue(np.all(tbl["init_time"] >= 0.0)) 

751 self.assertTrue(np.all(tbl["prep_time"] > 0.0)) 

752 self.assertTrue(np.all(tbl["run_time"] >= 0.0)) 

753 

754 def check_quantum_table(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None: 

755 """Check `ProvenanceQuantumGraph.make_quantum_table`. 

756 

757 Parameters 

758 ---------- 

759 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

760 Reader for the provenance quantum graph. 

761 expect_failure : `bool` 

762 Whether to expect one quantum of 'calibrate' to fail (`True`) or 

763 succeed without writing anything (`False`). 

764 """ 

765 t = prov.make_quantum_table() 

766 self.assertEqual(list(t["Task"]), ["calibrate", "consolidate", "resample", "coadd"]) 

767 self.assertEqual(t["TOTAL"][0], 8) 

768 self.assertEqual(t["EXPECTED"][0], 8) 

769 self.assertEqual(t["Blocked"][0], 0) 

770 self.assertEqual(t["TOTAL"][1], 2) 

771 self.assertEqual(t["EXPECTED"][1], 2) 

772 self.assertEqual(t["TOTAL"][2], 10) 

773 self.assertEqual(t["EXPECTED"][2], 10) 

774 if expect_failure: 

775 # calibrate 

776 self.assertEqual(t["Successful"][0], 7) 

777 self.assertEqual(t["Caveats"][0], "") 

778 self.assertEqual(t["Failed"][0], 1) 

779 # consolidate 

780 self.assertEqual(t["Successful"][1], 1) 

781 self.assertEqual(t["Caveats"][1], "") 

782 self.assertEqual(t["Failed"][1], 0) 

783 self.assertEqual(t["Blocked"][1], 1) 

784 # resample 

785 self.assertEqual(t["Successful"][2], 6) 

786 self.assertEqual(t["Caveats"][2], "") 

787 self.assertEqual(t["Failed"][2], 0) 

788 self.assertEqual(t["Blocked"][2], 4) 

789 else: 

790 # calibrate 

791 self.assertEqual(t["Successful"][0], 8) 

792 self.assertEqual(t["Caveats"][0], "*P(1)") 

793 self.assertEqual(t["Failed"][0], 0) 

794 # consolidate 

795 self.assertEqual(t["Successful"][1], 2) 

796 self.assertEqual(t["Caveats"][1], "") 

797 self.assertEqual(t["Failed"][1], 0) 

798 self.assertEqual(t["Blocked"][1], 0) 

799 # resample 

800 self.assertEqual(t["Successful"][2], 10) 

801 self.assertEqual(t["Caveats"][2], "*A(2)") 

802 self.assertEqual(t["Failed"][2], 0) 

803 self.assertEqual(t["Blocked"][2], 0) 

804 

805 def check_exception_table(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None: 

806 """Check `ProvenanceQuantumGraph.make_exception_table`. 

807 

808 Parameters 

809 ---------- 

810 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

811 Reader for the provenance quantum graph. 

812 expect_failure : `bool` 

813 Whether to expect one quantum of 'calibrate' to fail (`True`) or 

814 succeed without writing anything (`False`). 

815 """ 

816 t = prov.make_exception_table() 

817 self.assertEqual(list(t["Task"]), ["calibrate"]) 

818 self.assertEqual(list(t["Exception"]), ["lsst.pipe.base.tests.mocks.MockAlgorithmError"]) 

819 self.assertEqual(list(t["Successes"]), [int(not expect_failure)]) 

820 self.assertEqual(list(t["Failures"]), [int(expect_failure)]) 

821 

822 def check_report(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None: 

823 """Check `ProvenanceQuantumGraph.make_report`. 

824 

825 Parameters 

826 ---------- 

827 prov : `lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph` 

828 Reader for the provenance quantum graph. 

829 expect_failure : `bool` 

830 Whether to expect one quantum of 'calibrate' to fail (`True`) or 

831 succeed without writing anything (`False`). 

832 """ 

833 with tempfile.TemporaryDirectory(ignore_cleanup_errors=True) as data_id_table_dir: 

834 report = prov.make_status_report( 

835 also=QuantumAttemptStatus.SUCCESSFUL, data_id_table_dir=data_id_table_dir 

836 ) 

837 task_label = "calibrate" 

838 status_name = "FAILED" if expect_failure else "SUCCESSFUL" 

839 exc_type = "lsst.pipe.base.tests.mocks.MockAlgorithmError" 

840 self.assertEqual(report.root.keys(), {task_label}) 

841 self.assertEqual(report.root[task_label].keys(), {status_name}) 

842 self.assertEqual(report.root[task_label][status_name].keys(), {exc_type}) 

843 self.assertEqual(len(report.root[task_label][status_name][exc_type]), 1) 

844 qr = report.root[task_label][status_name][exc_type][0] 

845 self.assertIsInstance(qr, ProvenanceQuantumReport) 

846 self.assertEqual( 

847 qr.data_id, 

848 { 

849 "instrument": "Cam1", 

850 "visit": 2, 

851 "detector": 2, 

852 "band": "r", 

853 "day_obs": 20210909, 

854 "physical_filter": "Cam1-R1", 

855 }, 

856 ) 

857 tbl = astropy.table.Table.read( 

858 os.path.join(data_id_table_dir, task_label, status_name, f"{exc_type}.ecsv") 

859 ) 

860 self.assertCountEqual(tbl.colnames, ["instrument", "visit", "detector"]) 

861 self.assertEqual(len(tbl), 1) 

862 self.assertEqual(list(tbl["instrument"]), ["Cam1"]) 

863 self.assertEqual(list(tbl["detector"]), [2]) 

864 self.assertEqual(list(tbl["visit"]), [2]) 

865 

866 def test_all_successful(self) -> None: 

867 """Test running a full graph with no failures, and then scanning the 

868 results with incomplete=False. 

869 """ 

870 with self.make_test_repo() as prep: 

871 prep.config.incomplete = False 

872 start_time = time.time() 

873 attempted_quanta = list( 

874 self.iter_graph_execution(prep.butler_path, prep.predicted, raise_on_partial_outputs=False) 

875 ) 

876 self.assertCountEqual(attempted_quanta, prep.predicted.quantum_only_xgraph.nodes.keys()) 

877 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

878 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

879 prov = self.check_provenance_graph( 

880 prep.predicted, 

881 prep.butler, 

882 expect_failure=False, 

883 start_time=start_time, 

884 ) 

885 self.check_no_original_dirs(prep.butler_path, prov.header.output_run) 

886 for i, quantum_id in enumerate(attempted_quanta): 

887 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

888 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, attempted_quanta[:i]) 

889 

890 def test_all_successful_two_phase(self) -> None: 

891 """Test running some of a graph with no failures, scanning with 

892 incomplete=True, then finishing the graph and scanning again. 

893 """ 

894 with self.make_test_repo() as prep: 

895 start_time = time.time() 

896 execution_iter = self.iter_graph_execution( 

897 prep.butler_path, prep.predicted, raise_on_partial_outputs=False 

898 ) 

899 attempted_quanta = list(itertools.islice(execution_iter, 9)) 

900 self.assertEqual(len(attempted_quanta), 9) 

901 # Run the scanner while telling it execution is incomplete, so it 

902 # just abandons incomplete quanta and doesn't write the provenance 

903 # QG. 

904 prep.config.incomplete = True 

905 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

906 self.assertFalse(os.path.exists(cast(str, prep.config.output_path))) 

907 # Finish executing the quanta. 

908 attempted_quanta.extend(execution_iter) 

909 # Scan again, and write the provenance QG. 

910 prep.config.incomplete = False 

911 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

912 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

913 prov = self.check_provenance_graph( 

914 prep.predicted, 

915 prep.butler, 

916 expect_failure=False, 

917 start_time=start_time, 

918 ) 

919 for i, quantum_id in enumerate(attempted_quanta): 

920 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

921 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, attempted_quanta[:i]) 

922 

923 def test_some_failed(self) -> None: 

924 """Test running a full graph with some failures, and then scanning the 

925 results with incomplete=False. 

926 """ 

927 with self.make_test_repo() as prep: 

928 prep.config.incomplete = False 

929 start_time = time.time() 

930 attempted_quanta = list( 

931 self.iter_graph_execution(prep.butler_path, prep.predicted, raise_on_partial_outputs=True) 

932 ) 

933 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

934 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

935 prov = self.check_provenance_graph( 

936 prep.predicted, 

937 prep.butler, 

938 expect_failure=True, 

939 start_time=start_time, 

940 ) 

941 self.check_no_original_dirs(prep.butler_path, prov.header.output_run) 

942 for i, quantum_id in enumerate(attempted_quanta): 

943 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

944 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, attempted_quanta[:i]) 

945 

946 def test_some_failed_two_phase(self) -> None: 

947 """Test running a full graph with some failures, then scanning the 

948 results with incomplete=True, then scanning again with 

949 incomplete=False. 

950 """ 

951 with self.make_test_repo() as prep: 

952 start_time = time.time() 

953 attempted_quanta = list( 

954 self.iter_graph_execution(prep.butler_path, prep.predicted, raise_on_partial_outputs=True) 

955 ) 

956 prep.config.incomplete = True 

957 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

958 prep.config.incomplete = False 

959 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

960 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

961 prov = self.check_provenance_graph( 

962 prep.predicted, 

963 prep.butler, 

964 expect_failure=True, 

965 start_time=start_time, 

966 ) 

967 for i, quantum_id in enumerate(attempted_quanta): 

968 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

969 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, attempted_quanta[:i]) 

970 

971 def test_retry(self) -> None: 

972 """Test running a full graph with some failures, rerunning the quanta 

973 that failed or were blocked in the first attempt, and then scanning 

974 for provenance. 

975 """ 

976 with self.make_test_repo() as prep: 

977 start_time = time.time() 

978 attempted_quanta_1 = list( 

979 self.iter_graph_execution(prep.butler_path, prep.predicted, raise_on_partial_outputs=True) 

980 ) 

981 attempted_quanta_2 = list( 

982 self.iter_graph_execution( 

983 prep.butler_path, prep.predicted, raise_on_partial_outputs=False, is_retry=True 

984 ) 

985 ) 

986 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

987 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

988 prov = self.check_provenance_graph( 

989 prep.predicted, 

990 prep.butler, 

991 expect_failure=False, 

992 start_time=start_time, 

993 expect_failures_retried=True, 

994 ) 

995 for i, quantum_id in enumerate(attempted_quanta_1): 

996 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

997 self.assertEqual(qinfo["attempts"][0].previous_process_quanta, attempted_quanta_1[:i]) 

998 expected: list[uuid.UUID] = [] 

999 for quantum_id in attempted_quanta_2: 

1000 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

1001 if ( 

1002 quantum_id in attempted_quanta_1 

1003 and qinfo["attempts"][0].status is QuantumAttemptStatus.SUCCESSFUL 

1004 ): 

1005 # These weren't actually attempted twice, since they 

1006 # were already successful in the first round. 

1007 self.assertEqual(len(qinfo["attempts"]), 1) 

1008 else: 

1009 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, expected) 

1010 expected.append(quantum_id) 

1011 

1012 def test_promise_ingest_graph(self) -> None: 

1013 """Test running with promise_ingest_graph=True.""" 

1014 with self.make_test_repo() as prep: 

1015 prep.config.incomplete = False 

1016 prep.config.promise_ingest_graph = True 

1017 start_time = time.time() 

1018 attempted_quanta = list( 

1019 self.iter_graph_execution(prep.butler_path, prep.predicted, raise_on_partial_outputs=True) 

1020 ) 

1021 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

1022 self.assertFalse(prep.butler.query_datasets("calibrate_metadata", explain=False)) 

1023 self.assertFalse(prep.butler.query_datasets("consolidate_log", explain=False)) 

1024 self.assertFalse(prep.butler.query_datasets("resample_config", explain=False)) 

1025 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

1026 prov = self.check_provenance_graph( 

1027 prep.predicted, 

1028 prep.butler, 

1029 expect_failure=True, 

1030 start_time=start_time, 

1031 ) 

1032 self.check_no_original_dirs(prep.butler_path, prov.header.output_run) 

1033 for i, quantum_id in enumerate(attempted_quanta): 

1034 qinfo: ProvenanceQuantumInfo = prov.quantum_only_xgraph.nodes[quantum_id] 

1035 self.assertEqual(qinfo["attempts"][-1].previous_process_quanta, attempted_quanta[:i]) 

1036 

1037 def test_worker_failures(self) -> None: 

1038 """Test that if failures occur on (multiple) workers we shut down 

1039 gracefully instead of hanging. 

1040 """ 

1041 with self.make_test_repo() as prep: 

1042 with self.assertRaises(FatalWorkerError): 

1043 aggregate_graph(prep.predicted_path, "nonexistent", prep.config) 

1044 

1045 def test_aggregate_graph_cli_overrides(self) -> None: 

1046 """Test that command-line options override config attributes as 

1047 expected. 

1048 """ 

1049 

1050 def mock_run(predicted_path: str, butler_path: str, config: AggregatorConfig) -> None: 

1051 print(config.model_dump_json(indent=2)) 

1052 

1053 def check(result: Result, **kwargs: Any) -> None: 

1054 self.assertEqual(result.exit_code, 0, msg=result.output) 

1055 self.assertEqual( 

1056 result.output.strip(), AggregatorConfig(**kwargs).model_dump_json(indent=2).strip() 

1057 ) 

1058 

1059 self.maxDiff = None 

1060 runner = CliRunner() 

1061 with unittest.mock.patch("lsst.pipe.base.quantum_graph.aggregator.aggregate_graph", mock_run): 

1062 check(runner.invoke(aggregate_graph_cli, ("pg", "repo"))) 

1063 check(runner.invoke(aggregate_graph_cli, ("pg", "repo", "--output", "out")), output_path="out") 

1064 check(runner.invoke(aggregate_graph_cli, ("pg", "repo", "-j", "4")), n_processes=4) 

1065 check( 

1066 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--incomplete")), 

1067 incomplete=True, 

1068 ) 

1069 check(runner.invoke(aggregate_graph_cli, ("pg", "repo", "--dry-run")), dry_run=True) 

1070 check( 

1071 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--interactive-status")), 

1072 interactive_status=True, 

1073 ) 

1074 check( 

1075 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--log-status-interval", "120")), 

1076 log_status_interval=120, 

1077 ) 

1078 check( 

1079 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--no-register-dataset-types")), 

1080 register_dataset_types=False, 

1081 ) 

1082 check( 

1083 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--no-update-output-chain")), 

1084 update_output_chain=False, 

1085 ) 

1086 check( 

1087 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--worker-log-dir", "wlogs")), 

1088 worker_log_dir="wlogs", 

1089 ) 

1090 check( 

1091 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--worker-log-level", "DEBUG")), 

1092 worker_log_level="DEBUG", 

1093 ) 

1094 check( 

1095 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--zstd-level", "11")), 

1096 zstd_level=11, 

1097 ) 

1098 check( 

1099 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--zstd-dict-size", "143")), 

1100 zstd_dict_size=143, 

1101 ) 

1102 check( 

1103 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--zstd-dict-n-inputs", "2")), 

1104 zstd_dict_n_inputs=2, 

1105 ) 

1106 check( 

1107 runner.invoke(aggregate_graph_cli, ("pg", "repo", "--mock-storage-classes")), 

1108 mock_storage_classes=True, 

1109 ) 

1110 

1111 def check_provenance_report(self, result: click.testing.Result, root: str) -> None: 

1112 self.maxDiff = None 

1113 self.assertEqual(result.exit_code, 0, msg=result.output) 

1114 self.assertRegex( 

1115 result.output.replace("\n", ""), 

1116 re.compile( 

1117 r"\s*Task\s+Caveats\s+Failed\s+Blocked\s+Successful\s+TOTAL EXPECTED" 

1118 r"[-\s]*" 

1119 r"\s*calibrate \s* 1 \s* 0 \s* 7 \s* 8 \s* 8" 

1120 r"\s*consolidate \s* 0 \s* 1 \s* 1 \s* 2 \s* 2" 

1121 r"\s*resample \s* 0 \s* 4 \s* 6 \s* 10 \s* 10" 

1122 r"\s*coadd \s* 0 \s* 4 \s* 6 \s* 10 \s* 10" 

1123 r".*(Caveats[-\s]*)?.*" 

1124 r"\s*Task\s+Exception\s+Successes\s+Failures" 

1125 r"[-\s]*" 

1126 r"calibrate\s*lsst.pipe.base.tests.mocks.MockAlgorithmError \s* 0 \s* 1" 

1127 ), 

1128 ) 

1129 with open(os.path.join(root, "report.json")) as report_file: 

1130 report = ProvenanceReport.model_validate_json(report_file.read()) 

1131 self.assertEqual(report.root.keys(), {"calibrate"}) 

1132 self.assertEqual(report.root["calibrate"].keys(), {"FAILED"}) 

1133 self.assertEqual( 

1134 report.root["calibrate"]["FAILED"].keys(), {"lsst.pipe.base.tests.mocks.MockAlgorithmError"} 

1135 ) 

1136 self.assertEqual( 

1137 len(report.root["calibrate"]["FAILED"]["lsst.pipe.base.tests.mocks.MockAlgorithmError"]), 1 

1138 ) 

1139 self.assertEqual( 

1140 list(os.walk(os.path.join(root, "data_ids"))), 

1141 [ 

1142 (os.path.join(root, "data_ids"), ["calibrate"], []), 

1143 (os.path.join(root, "data_ids", "calibrate"), ["FAILED"], []), 

1144 ( 

1145 os.path.join(root, "data_ids", "calibrate", "FAILED"), 

1146 [], 

1147 ["lsst.pipe.base.tests.mocks.MockAlgorithmError.ecsv"], 

1148 ), 

1149 ], 

1150 ) 

1151 

1152 def check_provenance_report_json(self, result: click.testing.Result, root: str) -> None: 

1153 """Check that the json output is valid and has the correct number of 

1154 rows. 

1155 """ 

1156 loaded_json = json.loads(result.output) 

1157 self.assertEqual(len(loaded_json["tasks"]), 4) 

1158 self.assertEqual(len(loaded_json["exceptions"]), 1) 

1159 

1160 def check_no_original_dirs(self, butler_path: str, output_run: str) -> None: 

1161 """Check that there are no config/log/metadata directories in 

1162 the butler's directory for this output run. 

1163 """ 

1164 root = os.path.join(butler_path, output_run) 

1165 for subdir in os.listdir(root): 

1166 if subdir.endswith("_config") or subdir.endswith("_metadata") or subdir.endswith("_log"): 

1167 raise AssertionError(f"Directory {os.path.join(root, subdir)} still exists.") 

1168 

1169 def test_provenance_report_content(self) -> None: 

1170 """Test the provenance-report CLI command.""" 

1171 with self.make_test_repo() as prep: 

1172 prep.config.incomplete = False 

1173 prep.config.promise_ingest_graph = True 

1174 for _ in self.iter_graph_execution( 

1175 prep.butler_path, prep.predicted, raise_on_partial_outputs=True 

1176 ): 

1177 pass 

1178 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

1179 self.assertFalse(prep.butler.query_datasets("calibrate_metadata", explain=False)) 

1180 self.assertFalse(prep.butler.query_datasets("consolidate_log", explain=False)) 

1181 self.assertFalse(prep.butler.query_datasets("resample_config", explain=False)) 

1182 

1183 # First test on a provenance graph file that has not been ingested. 

1184 runner = CliRunner() 

1185 report_root = os.path.join(prep.butler_path, "uningested") 

1186 result = runner.invoke( 

1187 provenance_report_cli, 

1188 ( 

1189 cast(str, prep.config.output_path), 

1190 "--status-report", 

1191 os.path.join(report_root, "report.json"), 

1192 "--data-id-table-dir", 

1193 os.path.join(report_root, "data_ids"), 

1194 ), 

1195 ) 

1196 self.check_provenance_report(result, report_root) 

1197 

1198 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10) 

1199 

1200 runner = CliRunner() 

1201 report_root = os.path.join(prep.butler_path, "ingested") 

1202 result = runner.invoke( 

1203 provenance_report_cli, 

1204 ( 

1205 prep.butler_path, 

1206 *prep.butler.collections.defaults, 

1207 "--status-report", 

1208 os.path.join(report_root, "report.json"), 

1209 "--data-id-table-dir", 

1210 os.path.join(report_root, "data_ids"), 

1211 ), 

1212 ) 

1213 self.check_provenance_report(result, os.path.join(report_root)) 

1214 

1215 def test_provenance_report_cli_json_output(self) -> None: 

1216 """Test the json output option with a mocked implementation.""" 

1217 with self.make_test_repo() as prep: 

1218 prep.config.incomplete = False 

1219 prep.config.promise_ingest_graph = True 

1220 for _ in self.iter_graph_execution( 

1221 prep.butler_path, prep.predicted, raise_on_partial_outputs=True 

1222 ): 

1223 pass 

1224 aggregate_graph(prep.predicted_path, prep.butler_path, prep.config) 

1225 self.assertFalse(prep.butler.query_datasets("calibrate_metadata", explain=False)) 

1226 self.assertFalse(prep.butler.query_datasets("consolidate_log", explain=False)) 

1227 self.assertFalse(prep.butler.query_datasets("resample_config", explain=False)) 

1228 

1229 # First test on a provenance graph file that has not been ingested. 

1230 runner = CliRunner() 

1231 report_root = os.path.join(prep.butler_path, "uningested") 

1232 result = runner.invoke( 

1233 provenance_report_cli, 

1234 ( 

1235 cast(str, prep.config.output_path), 

1236 "--format", 

1237 "json", 

1238 ), 

1239 ) 

1240 self.check_provenance_report_json(result, report_root) 

1241 

1242 def test_provenance_report_cli_overrides(self) -> None: 

1243 """Test the provenance-report CLI command with a mocked 

1244 implementation. 

1245 """ 

1246 

1247 class MakeManyReportsArgs(pydantic.BaseModel): 

1248 status_report_file: str | None = None 

1249 print_quantum_table: bool = True 

1250 print_exception_table: bool = True 

1251 states: list[QuantumAttemptStatus] = pydantic.Field( 

1252 default_factory=lambda: [ 

1253 QuantumAttemptStatus.FAILED, 

1254 QuantumAttemptStatus.ABORTED, 

1255 QuantumAttemptStatus.ABORTED_SUCCESS, 

1256 ] 

1257 ) 

1258 with_caveats: QuantumSuccessCaveats | None = None 

1259 data_id_table_dir: str | None = None 

1260 

1261 @pydantic.model_validator(mode="after") 

1262 def _sort_states(self) -> MakeManyReportsArgs: 

1263 self.states.sort(key=lambda e: e.value) 

1264 return self 

1265 

1266 class MockProvenanceQuantumGraph(pydantic.BaseModel): 

1267 repo_or_filename: str 

1268 collection: str | None 

1269 quanta: list[uuid.UUID] | None = None 

1270 datasets: list[uuid.UUID] | None = None 

1271 writeable: bool = False 

1272 make_many_reports_args: MakeManyReportsArgs | None = None 

1273 

1274 @classmethod 

1275 @contextmanager 

1276 def from_args( 

1277 cls, repo_or_filename: str, /, **kwargs: Any 

1278 ) -> Iterator[tuple[MockProvenanceQuantumGraph, None]]: 

1279 yield cls(repo_or_filename=repo_or_filename, **kwargs), None 

1280 

1281 def make_many_reports(self, **kwargs: Any) -> None: 

1282 self.make_many_reports_args = MakeManyReportsArgs(**kwargs) 

1283 print(self.model_dump_json(indent=2)) 

1284 

1285 def check( 

1286 result: Result, repo_or_filename: str, collection: str | None = None, **kwargs: Any 

1287 ) -> None: 

1288 self.assertEqual(result.exit_code, 0, msg=result.output) 

1289 self.assertEqual( 

1290 result.output.strip(), 

1291 MockProvenanceQuantumGraph( 

1292 repo_or_filename=repo_or_filename, 

1293 collection=collection, 

1294 datasets=[], 

1295 writeable=False, 

1296 make_many_reports_args=MakeManyReportsArgs(**kwargs), 

1297 ).model_dump_json(indent=2), 

1298 ) 

1299 

1300 self.maxDiff = None 

1301 runner = CliRunner() 

1302 with unittest.mock.patch( 

1303 "lsst.pipe.base.quantum_graph.ProvenanceQuantumGraph", MockProvenanceQuantumGraph 

1304 ): 

1305 check( 

1306 runner.invoke(provenance_report_cli, ("repo", "collection1")), 

1307 repo_or_filename="repo", 

1308 collection="collection1", 

1309 ) 

1310 check( 

1311 runner.invoke(provenance_report_cli, ("filename1",)), 

1312 repo_or_filename="filename1", 

1313 collection=None, 

1314 ) 

1315 check( 

1316 runner.invoke(provenance_report_cli, ("repo", "collection1", "--no-quantum-table")), 

1317 repo_or_filename="repo", 

1318 collection="collection1", 

1319 print_quantum_table=False, 

1320 ) 

1321 check( 

1322 runner.invoke(provenance_report_cli, ("repo", "collection1", "--no-exception-table")), 

1323 repo_or_filename="repo", 

1324 collection="collection1", 

1325 print_exception_table=False, 

1326 ) 

1327 check( 

1328 runner.invoke(provenance_report_cli, ("repo", "collection1", "--status-report", "filename2")), 

1329 repo_or_filename="repo", 

1330 collection="collection1", 

1331 status_report_file="filename2", 

1332 ) 

1333 check( 

1334 runner.invoke(provenance_report_cli, ("repo", "collection1", "--state", "SUCCESSFUL")), 

1335 repo_or_filename="repo", 

1336 collection="collection1", 

1337 states=[ 

1338 QuantumAttemptStatus.SUCCESSFUL, 

1339 QuantumAttemptStatus.FAILED, 

1340 QuantumAttemptStatus.ABORTED, 

1341 QuantumAttemptStatus.ABORTED_SUCCESS, 

1342 ], 

1343 ) 

1344 check( 

1345 runner.invoke(provenance_report_cli, ("repo", "collection1", "--no-state", "FAILED")), 

1346 repo_or_filename="repo", 

1347 collection="collection1", 

1348 states=[ 

1349 QuantumAttemptStatus.ABORTED, 

1350 QuantumAttemptStatus.ABORTED_SUCCESS, 

1351 ], 

1352 ) 

1353 check( 

1354 runner.invoke( 

1355 provenance_report_cli, ("repo", "collection1", "--caveat", "PARTIAL_OUTPUTS_ERROR") 

1356 ), 

1357 repo_or_filename="repo", 

1358 collection="collection1", 

1359 with_caveats=QuantumSuccessCaveats.PARTIAL_OUTPUTS_ERROR, 

1360 states=[ 

1361 QuantumAttemptStatus.SUCCESSFUL, 

1362 QuantumAttemptStatus.FAILED, 

1363 QuantumAttemptStatus.ABORTED, 

1364 QuantumAttemptStatus.ABORTED_SUCCESS, 

1365 ], 

1366 ) 

1367 check( 

1368 runner.invoke(provenance_report_cli, ("repo", "collection1", "--data-id-table-dir", "dir1")), 

1369 repo_or_filename="repo", 

1370 collection="collection1", 

1371 data_id_table_dir="dir1", 

1372 ) 

1373 

1374 def test_bad_metadata_readable(self) -> None: 

1375 """Test that consolidated metadata accidentally written with floats 

1376 transformed to JSON null are now readable. 

1377 """ 

1378 with open(os.path.join(os.path.dirname(__file__), "data", "DM-54057.json")) as stream: 

1379 data = stream.read() 

1380 prov_md = ProvenanceTaskMetadataModel.model_validate_json(data) 

1381 self.assertTrue(np.isnan(prov_md.attempts[0]["calibrateImage:psf_measure_psf"]["spatialFitChi2"])) 

1382 

1383 

1384class CompressionDictTestCase(unittest.TestCase): 

1385 """Test `Writer.make_compression_dictionary` guards against zstandard's 

1386 undocumented requirements. 

1387 

1388 zstandard's ``ZDICT_optimizeTrainFromBuffer_fastCover`` refuses to train 

1389 from fewer than 5 training samples (with its default 0.75 train/test split, 

1390 that means at least 7 samples overall) and raises 

1391 ``cannot train dict: Src size is incorrect`` otherwise. We set the 

1392 threshold at 10 to be safe. 

1393 """ 

1394 

1395 class _Log: 

1396 def __init__(self) -> None: 

1397 self.messages: list[str] = [] 

1398 

1399 def __getattr__(self, name: str) -> Any: 

1400 if name in ("info", "warning", "warn"): 1400 ↛ 1406line 1400 didn't jump to line 1406 because the condition on line 1400 was always true

1401 

1402 def _log(msg: str, *args: Any) -> None: 

1403 self.messages.append(msg % args) 

1404 

1405 return _log 

1406 raise AttributeError(name) 

1407 

1408 class _Config: 

1409 zstd_dict_size = 256 

1410 zstd_dict_n_inputs = 1 

1411 zstd_dict_input_max_bytes = 1_048_576 

1412 

1413 class _Scan: 

1414 is_compressed = False 

1415 

1416 def __init__(self, i: int, scale: int = 64) -> None: 

1417 self.quantum = b"q%d" % i * scale 

1418 self.metadata = b"m%d" % i * scale 

1419 self.logs = b"l%d" % i * scale 

1420 

1421 def make_writer(self, n_scans: int) -> Writer: 

1422 writer = object.__new__(Writer) 

1423 writer.comms = types.SimpleNamespace(config=self._Config(), log=self._Log()) 

1424 writer.predicted = types.SimpleNamespace(quantum_datasets={}) 

1425 writer.pending_compression_training = [self._Scan(i) for i in range(n_scans)] 

1426 return writer 

1427 

1428 def test_fallback_to_no_dictionary_with_too_few_samples(self) -> None: 

1429 # 3 scans -> 9 samples, fewer than the 10 required by our code. 

1430 writer = self.make_writer(3) 

1431 cdict = writer.make_compression_dictionary() 

1432 self.assertEqual(cdict.as_bytes(), b"") 

1433 self.assertTrue( 

1434 any("Only 9 < 10 samples" in message for message in writer.comms.log.messages), 

1435 "expected a warning explaining the fallback.", 

1436 ) 

1437 

1438 def test_trains_dictionary_with_enough_samples(self) -> None: 

1439 # 4 scans -> 12 samples, enough for zstandard to train. 

1440 writer = self.make_writer(4) 

1441 cdict = writer.make_compression_dictionary() 

1442 self.assertNotEqual(cdict.as_bytes(), b"") 

1443 

1444 

1445if __name__ == "__main__": 

1446 lsst.utils.tests.init() 

1447 unittest.main()