Coverage for tests/test_aggregator.py: 99%
565 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-14 02:17 -0700
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-14 02:17 -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/>.
28from __future__ import annotations
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
44import astropy.table
45import click.testing
46import numpy as np
47import pydantic
48from click.testing import CliRunner, Result
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
94@dataclasses.dataclass
95class PrepInfo:
96 """Struct of objects used in an aggregator test."""
98 butler: Butler
99 butler_path: str
100 predicted: PredictedQuantumGraph
101 predicted_path: str
102 config: AggregatorConfig
105class AggregatorTestCase(unittest.TestCase):
106 """Unit tests for `lsst.pipe.base.quantum_graph.aggregator`."""
108 @staticmethod
109 @contextmanager
110 def make_test_repo() -> Iterator[PrepInfo]:
111 """Make a test data repository and predicted quantum graph.
113 Returns
114 -------
115 prep_info : `PrepInfo`
116 Objects used in aggregator tests.
118 Notes
119 -----
120 The pipeline graph used by this task looks like this:
122 ■ calibrate: {detector, visit}
123 ╭─┤
124 ■ │ consolidate: {visit}
125 │
126 ■ resample: {patch, visit}
127 │
128 ■ coadd: {band, patch}
130 The data can be visualized via::
132 python -m lsst.daf.butler.tests.registry_data.spatial
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 )
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.
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.
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
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`.
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.
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
522 def _expect_all_exist(self, existence: list[bool], msg: str) -> None:
523 self.assertTrue(all(existence), msg=msg)
525 def _expect_none_exist(self, existence: list[bool], msg: str) -> None:
526 self.assertFalse(any(existence), msg=msg)
528 def _expect_one_missing(self, existence: list[bool], msg: str) -> None:
529 self.assertEqual(existence.count(False), 1, msg=msg)
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)
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)
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)
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.
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.
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
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.
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())
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.
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())
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)
707 def check_packages(self, butler: Butler) -> None:
708 """Check fetching package versions from the provenance graph.
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)
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.
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))
754 def check_quantum_table(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None:
755 """Check `ProvenanceQuantumGraph.make_quantum_table`.
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)
805 def check_exception_table(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None:
806 """Check `ProvenanceQuantumGraph.make_exception_table`.
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)])
822 def check_report(self, prov: ProvenanceQuantumGraph, expect_failure: bool) -> None:
823 """Check `ProvenanceQuantumGraph.make_report`.
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])
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])
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])
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])
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])
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)
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])
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)
1045 def test_aggregate_graph_cli_overrides(self) -> None:
1046 """Test that command-line options override config attributes as
1047 expected.
1048 """
1050 def mock_run(predicted_path: str, butler_path: str, config: AggregatorConfig) -> None:
1051 print(config.model_dump_json(indent=2))
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 )
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 )
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 )
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)
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.")
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))
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)
1198 ingest_graph(prep.butler_path, prep.config.output_path, transfer="move", batch_size=10)
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))
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))
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)
1242 def test_provenance_report_cli_overrides(self) -> None:
1243 """Test the provenance-report CLI command with a mocked
1244 implementation.
1245 """
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
1261 @pydantic.model_validator(mode="after")
1262 def _sort_states(self) -> MakeManyReportsArgs:
1263 self.states.sort(key=lambda e: e.value)
1264 return self
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
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
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))
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 )
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 )
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"]))
1384class CompressionDictTestCase(unittest.TestCase):
1385 """Test `Writer.make_compression_dictionary` guards against zstandard's
1386 undocumented requirements.
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 """
1395 class _Log:
1396 def __init__(self) -> None:
1397 self.messages: list[str] = []
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
1402 def _log(msg: str, *args: Any) -> None:
1403 self.messages.append(msg % args)
1405 return _log
1406 raise AttributeError(name)
1408 class _Config:
1409 zstd_dict_size = 256
1410 zstd_dict_n_inputs = 1
1411 zstd_dict_input_max_bytes = 1_048_576
1413 class _Scan:
1414 is_compressed = False
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
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
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 )
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"")
1445if __name__ == "__main__":
1446 lsst.utils.tests.init()
1447 unittest.main()