Coverage for tests/test_query_utilities.py: 100%
268 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-26 08:49 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-26 08:49 +0000
1# This file is part of daf_butler.
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/>.
28"""Tests for non-public Butler._query functionality that is not specific to
29any Butler or QueryDriver implementation.
30"""
32from __future__ import annotations
34import unittest
35from collections.abc import Iterable, Set
37import astropy.time
39from lsst.daf.butler import DimensionUniverse, InvalidQueryError, Timespan
40from lsst.daf.butler.dimensions import DimensionElement, DimensionGroup
41from lsst.daf.butler.queries import tree as qt
42from lsst.daf.butler.queries.expression_factory import ExpressionFactory
43from lsst.daf.butler.queries.overlaps import OverlapsVisitor, _NaiveDisjointSet
44from lsst.daf.butler.queries.visitors import PredicateVisitFlags
45from lsst.sphgeom import Mq3cPixelization, Region
48class ColumnSetTestCase(unittest.TestCase):
49 """Tests for lsst.daf.butler.queries.ColumnSet."""
51 def setUp(self) -> None:
52 self.universe = DimensionUniverse()
54 def test_basics(self) -> None:
55 columns = qt.ColumnSet(self.universe.conform(["detector"]))
56 self.assertNotEqual(columns, columns.dimensions.names) # intentionally not comparable to other sets
57 self.assertEqual(columns.dimensions, self.universe.conform(["detector"]))
58 self.assertFalse(columns.dataset_fields)
59 columns.dataset_fields["bias"].add("dataset_id")
60 self.assertEqual(dict(columns.dataset_fields), {"bias": {"dataset_id"}})
61 columns.dimension_fields["detector"].add("purpose")
62 self.assertEqual(columns.dimension_fields["detector"], {"purpose"})
63 self.assertTrue(columns)
64 self.assertEqual(
65 list(columns),
66 [(k, None) for k in columns.dimensions.data_coordinate_keys]
67 + [("detector", "purpose"), ("bias", "dataset_id")],
68 )
69 self.assertEqual(str(columns), "{instrument, detector, detector:purpose, bias:dataset_id}")
70 empty = qt.ColumnSet(self.universe.empty)
71 self.assertFalse(empty)
72 self.assertFalse(columns.issubset(empty))
73 self.assertTrue(columns.issuperset(empty))
74 self.assertTrue(columns.isdisjoint(empty))
75 copy = columns.copy()
76 self.assertEqual(columns, copy)
77 self.assertTrue(columns.issubset(copy))
78 self.assertTrue(columns.issuperset(copy))
79 self.assertFalse(columns.isdisjoint(copy))
80 copy.dataset_fields["bias"].add("timespan")
81 copy.dimension_fields["detector"].add("name")
82 copy.update_dimensions(self.universe.conform(["band"]))
83 self.assertEqual(copy.dataset_fields["bias"], {"dataset_id", "timespan"})
84 self.assertEqual(columns.dataset_fields["bias"], {"dataset_id"})
85 self.assertEqual(copy.dimension_fields["detector"], {"purpose", "name"})
86 self.assertEqual(columns.dimension_fields["detector"], {"purpose"})
87 self.assertTrue(columns.issubset(copy))
88 self.assertFalse(columns.issuperset(copy))
89 self.assertFalse(columns.isdisjoint(copy))
90 columns.update(copy)
91 self.assertEqual(columns, copy)
92 self.assertTrue(columns.is_timespan("visit", "timespan"))
93 self.assertFalse(columns.is_timespan("visit", None))
94 self.assertFalse(columns.is_timespan("detector", "purpose"))
96 def test_drop_dimension_keys(self):
97 columns = qt.ColumnSet(self.universe.conform(["physical_filter"]))
98 columns.drop_implied_dimension_keys()
99 self.assertEqual(list(columns), [("instrument", None), ("physical_filter", None)])
100 undropped = qt.ColumnSet(columns.dimensions)
101 self.assertTrue(columns.issubset(undropped))
102 self.assertFalse(columns.issuperset(undropped))
103 self.assertFalse(columns.isdisjoint(undropped))
104 band_only = qt.ColumnSet(self.universe.conform(["band"]))
105 self.assertFalse(columns.issubset(band_only))
106 self.assertFalse(columns.issuperset(band_only))
107 self.assertTrue(columns.isdisjoint(band_only))
108 copy = columns.copy()
109 copy.update(band_only)
110 self.assertEqual(copy, undropped)
111 columns.restore_dimension_keys()
112 self.assertEqual(columns, undropped)
114 def test_get_column_spec(self) -> None:
115 columns = qt.ColumnSet(self.universe.conform(["detector"]))
116 columns.dimension_fields["detector"].add("purpose")
117 columns.dataset_fields["bias"].update(["dataset_id", "run", "collection", "timespan", "ingest_date"])
118 self.assertEqual(columns.get_column_spec("instrument", None).name, "instrument")
119 self.assertEqual(columns.get_column_spec("instrument", None).type, "string")
120 self.assertEqual(columns.get_column_spec("instrument", None).nullable, False)
121 self.assertEqual(columns.get_column_spec("detector", None).name, "detector")
122 self.assertEqual(columns.get_column_spec("detector", None).type, "int")
123 self.assertEqual(columns.get_column_spec("detector", None).nullable, False)
124 self.assertEqual(columns.get_column_spec("detector", "purpose").name, "detector:purpose")
125 self.assertEqual(columns.get_column_spec("detector", "purpose").type, "string")
126 self.assertEqual(columns.get_column_spec("detector", "purpose").nullable, True)
127 self.assertEqual(columns.get_column_spec("bias", "dataset_id").name, "bias:dataset_id")
128 self.assertEqual(columns.get_column_spec("bias", "dataset_id").type, "uuid")
129 self.assertEqual(columns.get_column_spec("bias", "dataset_id").nullable, False)
130 self.assertEqual(columns.get_column_spec("bias", "run").name, "bias:run")
131 self.assertEqual(columns.get_column_spec("bias", "run").type, "string")
132 self.assertEqual(columns.get_column_spec("bias", "run").nullable, False)
133 self.assertEqual(columns.get_column_spec("bias", "collection").name, "bias:collection")
134 self.assertEqual(columns.get_column_spec("bias", "collection").type, "string")
135 self.assertEqual(columns.get_column_spec("bias", "collection").nullable, False)
136 self.assertEqual(columns.get_column_spec("bias", "timespan").name, "bias:timespan")
137 self.assertEqual(columns.get_column_spec("bias", "timespan").type, "timespan")
138 self.assertEqual(columns.get_column_spec("bias", "timespan").nullable, True)
139 self.assertEqual(columns.get_column_spec("bias", "ingest_date").name, "bias:ingest_date")
140 self.assertEqual(columns.get_column_spec("bias", "ingest_date").type, "datetime")
141 self.assertEqual(columns.get_column_spec("bias", "ingest_date").nullable, True)
144class _RecordingOverlapsVisitor(OverlapsVisitor):
145 def __init__(self, dimensions: DimensionGroup, calibration_dataset_types: Set[str] = frozenset()):
146 super().__init__(dimensions, calibration_dataset_types)
147 self.spatial_constraints: list[tuple[str, PredicateVisitFlags]] = []
148 self.spatial_joins: list[tuple[str, str, PredicateVisitFlags]] = []
149 self.temporal_dimension_joins: list[tuple[str, str, PredicateVisitFlags]] = []
150 self.validity_range_dimension_joins: list[tuple[str, str, PredicateVisitFlags]] = []
151 self.validity_range_joins: list[tuple[str, str, PredicateVisitFlags]] = []
153 def visit_spatial_constraint(
154 self, element: DimensionElement, region: Region, flags: PredicateVisitFlags
155 ) -> qt.Predicate | None:
156 self.spatial_constraints.append((element.name, flags))
157 return super().visit_spatial_constraint(element, region, flags)
159 def visit_spatial_join(
160 self, a: DimensionElement, b: DimensionElement, flags: PredicateVisitFlags
161 ) -> qt.Predicate | None:
162 self.spatial_joins.append((a.name, b.name, flags))
163 return super().visit_spatial_join(a, b, flags)
165 def visit_temporal_dimension_join(
166 self, a: DimensionElement, b: DimensionElement, flags: PredicateVisitFlags
167 ) -> qt.Predicate | None:
168 self.temporal_dimension_joins.append((a.name, b.name, flags))
169 return super().visit_temporal_dimension_join(a, b, flags)
171 def visit_validity_range_dimension_join(
172 self, a: str, b: DimensionElement, flags: PredicateVisitFlags
173 ) -> qt.Predicate | None:
174 self.validity_range_dimension_joins.append((a, b.name, flags))
175 return super().visit_validity_range_dimension_join(a, b, flags)
177 def visit_validity_range_join(self, a: str, b: str, flags: PredicateVisitFlags) -> qt.Predicate | None:
178 self.validity_range_joins.append((a, b, flags))
179 return super().visit_validity_range_join(a, b, flags)
182class OverlapsVisitorTestCase(unittest.TestCase):
183 """Tests for lsst.daf.butler.queries.overlaps.OverlapsVisitor, which is
184 responsible for validating and inferring spatial and temporal joins and
185 constraints.
186 """
188 def setUp(self) -> None:
189 self.universe = DimensionUniverse()
191 def run_visitor(
192 self,
193 dimensions: Iterable[str],
194 predicate: qt.Predicate,
195 expected: str | None = None,
196 join_operands: Iterable[DimensionGroup] = (),
197 calibration_dataset_types: Set[str] = frozenset(),
198 ) -> _RecordingOverlapsVisitor:
199 visitor = _RecordingOverlapsVisitor(self.universe.conform(dimensions), calibration_dataset_types)
200 if expected is None:
201 expected = str(predicate)
202 new_predicate = visitor.run(predicate, join_operands=join_operands)
203 self.assertEqual(str(new_predicate), expected)
204 return visitor
206 def test_trivial(self) -> None:
207 """Test the overlaps visitor when there is nothing spatial or temporal
208 in the query at all.
209 """
210 x = ExpressionFactory(self.universe)
211 # Trivial predicate.
212 visitor = self.run_visitor(["physical_filter"], qt.Predicate.from_bool(True))
213 self.assertFalse(visitor.spatial_joins)
214 self.assertFalse(visitor.spatial_constraints)
215 self.assertFalse(visitor.temporal_dimension_joins)
216 # Non-overlap predicate.
217 visitor = self.run_visitor(["physical_filter"], x.any(x.band == "r", x.band == "i"))
218 self.assertFalse(visitor.spatial_joins)
219 self.assertFalse(visitor.spatial_constraints)
220 self.assertFalse(visitor.temporal_dimension_joins)
222 def test_invalid_spatial_overlap_operands(self) -> None:
223 """Test defensive validation of spatial overlap operands."""
224 pixelization = Mq3cPixelization(10)
225 region = qt.make_column_literal(pixelization.quad(12058870))
226 visitor = _RecordingOverlapsVisitor(self.universe.empty)
227 with self.assertRaisesRegex(InvalidQueryError, "requires at least one dimension region column"):
228 visitor.visit_spatial_overlap(region, region, PredicateVisitFlags(0))
230 def test_one_spatial_family(self) -> None:
231 """Test the overlaps visitor when there is one spatial family."""
232 x = ExpressionFactory(self.universe)
233 pixelization = Mq3cPixelization(10)
234 region = pixelization.quad(12058870)
235 # Trivial predicate.
236 visitor = self.run_visitor(["visit"], qt.Predicate.from_bool(True))
237 self.assertFalse(visitor.spatial_joins)
238 self.assertFalse(visitor.spatial_constraints)
239 self.assertFalse(visitor.temporal_dimension_joins)
240 # Non-overlap predicate.
241 visitor = self.run_visitor(["visit"], x.any(x.band == "r", x.visit > 2))
242 self.assertFalse(visitor.spatial_joins)
243 self.assertFalse(visitor.spatial_constraints)
244 self.assertFalse(visitor.temporal_dimension_joins)
245 # Spatial constraint predicate, in various positions relative to other
246 # non-overlap predicates.
247 visitor = self.run_visitor(["visit"], x.visit.region.overlaps(region))
248 self.assertEqual(visitor.spatial_constraints, [(self.universe["visit"], PredicateVisitFlags(0))])
249 visitor = self.run_visitor(["visit"], x.all(x.visit.region.overlaps(region), x.band == "r"))
250 self.assertEqual(
251 visitor.spatial_constraints, [(self.universe["visit"], PredicateVisitFlags.HAS_AND_SIBLINGS)]
252 )
253 visitor = self.run_visitor(["visit"], x.any(x.visit.region.overlaps(region), x.band == "r"))
254 self.assertEqual(
255 visitor.spatial_constraints, [(self.universe["visit"], PredicateVisitFlags.HAS_OR_SIBLINGS)]
256 )
257 visitor = self.run_visitor(
258 ["visit"],
259 x.all(
260 x.any(x.literal(region).overlaps(x.visit.region), x.band == "r"),
261 x.visit.observation_reason == "science",
262 ),
263 )
264 self.assertEqual(
265 visitor.spatial_constraints,
266 [
267 (
268 self.universe["visit"],
269 PredicateVisitFlags.HAS_OR_SIBLINGS | PredicateVisitFlags.HAS_AND_SIBLINGS,
270 )
271 ],
272 )
273 visitor = self.run_visitor(
274 ["visit"],
275 x.any(
276 x.all(x.visit.region.overlaps(region), x.band == "r"),
277 x.visit.observation_reason == "science",
278 ),
279 )
280 self.assertEqual(
281 visitor.spatial_constraints,
282 [
283 (
284 self.universe["visit"],
285 PredicateVisitFlags.HAS_OR_SIBLINGS | PredicateVisitFlags.HAS_AND_SIBLINGS,
286 )
287 ],
288 )
289 # A spatial join between dimensions in the same family is an error.
290 with self.assertRaises(InvalidQueryError):
291 self.run_visitor(["patch", "tract"], x.patch.region.overlaps(x.tract.region))
293 def test_single_unambiguous_spatial_join(self) -> None:
294 """Test the overlaps visitor when there are two spatial families with
295 one dimension element in each, and hence exactly one join is needed.
296 """
297 x = ExpressionFactory(self.universe)
298 # Trivial predicate; an automatic join is added. Order of elements in
299 # automatic joins is lexicographical in order to be deterministic.
300 visitor = self.run_visitor(
301 ["visit", "tract"], qt.Predicate.from_bool(True), "tract.region OVERLAPS visit.region"
302 )
303 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags.HAS_AND_SIBLINGS)])
304 self.assertFalse(visitor.spatial_constraints)
305 self.assertFalse(visitor.temporal_dimension_joins)
306 # Non-overlap predicate; an automatic join is added.
307 visitor = self.run_visitor(
308 ["visit", "tract"],
309 x.all(x.band == "r", x.visit > 2),
310 "band == 'r' AND visit > 2 AND tract.region OVERLAPS visit.region",
311 )
312 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags.HAS_AND_SIBLINGS)])
313 self.assertFalse(visitor.spatial_constraints)
314 self.assertFalse(visitor.temporal_dimension_joins)
315 # The same overlap predicate that would be added automatically has been
316 # added manually.
317 visitor = self.run_visitor(
318 ["visit", "tract"],
319 x.tract.region.overlaps(x.visit.region),
320 "tract.region OVERLAPS visit.region",
321 )
322 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags(0))])
323 self.assertFalse(visitor.spatial_constraints)
324 self.assertFalse(visitor.temporal_dimension_joins)
325 # Add the join overlap predicate in an OR expression, which is unusual
326 # but enough to block the addition of an automatic join; we assume the
327 # user knows what they're doing.
328 visitor = self.run_visitor(
329 ["visit", "tract"],
330 x.any(x.visit > 2, x.tract.region.overlaps(x.visit.region)),
331 "visit > 2 OR tract.region OVERLAPS visit.region",
332 )
333 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags.HAS_OR_SIBLINGS)])
334 self.assertFalse(visitor.spatial_constraints)
335 self.assertFalse(visitor.temporal_dimension_joins)
336 # Add the join overlap predicate in a NOT expression, which is unusual
337 # but permitted in the same sense as OR expressions.
338 visitor = self.run_visitor(
339 ["visit", "tract"],
340 x.not_(x.tract.region.overlaps(x.visit.region)),
341 "NOT tract.region OVERLAPS visit.region",
342 )
343 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags.INVERTED)])
344 self.assertFalse(visitor.spatial_constraints)
345 self.assertFalse(visitor.temporal_dimension_joins)
346 # Add a "join operand" whose dimensions include both spatial families.
347 # This blocks an automatic join from being created, because we assume
348 # that join operand (e.g. a materialization or dataset search) already
349 # encodes some spatial join.
350 visitor = self.run_visitor(
351 ["visit", "tract"],
352 qt.Predicate.from_bool(True),
353 "True",
354 join_operands=[self.universe.conform(["tract", "visit"])],
355 )
356 self.assertFalse(visitor.spatial_joins)
357 self.assertFalse(visitor.spatial_constraints)
358 self.assertFalse(visitor.temporal_dimension_joins)
360 def test_single_flexible_spatial_join(self) -> None:
361 """Test the overlaps visitor when there are two spatial families and
362 one has multiple dimension elements.
363 """
364 x = ExpressionFactory(self.universe)
365 # Trivial predicate; an automatic join between the fine-grained
366 # elements is added. Order of elements in automatic joins is
367 # lexicographical in order to be deterministic.
368 visitor = self.run_visitor(
369 ["visit", "detector", "patch"],
370 qt.Predicate.from_bool(True),
371 "patch.region OVERLAPS visit_detector_region.region",
372 )
373 self.assertEqual(
374 visitor.spatial_joins, [("patch", "visit_detector_region", PredicateVisitFlags.HAS_AND_SIBLINGS)]
375 )
376 self.assertFalse(visitor.spatial_constraints)
377 self.assertFalse(visitor.temporal_dimension_joins)
378 # The same overlap predicate that would be added automatically has been
379 # added manually.
380 visitor = self.run_visitor(
381 ["visit", "detector", "patch"],
382 x.patch.region.overlaps(x.visit_detector_region.region),
383 "patch.region OVERLAPS visit_detector_region.region",
384 )
385 self.assertEqual(visitor.spatial_joins, [("patch", "visit_detector_region", PredicateVisitFlags(0))])
386 self.assertFalse(visitor.spatial_constraints)
387 self.assertFalse(visitor.temporal_dimension_joins)
388 # A coarse overlap join has been added; respect it and do not add an
389 # automatic one.
390 visitor = self.run_visitor(
391 ["visit", "detector", "patch"],
392 x.tract.region.overlaps(x.visit.region),
393 "tract.region OVERLAPS visit.region",
394 )
395 self.assertEqual(visitor.spatial_joins, [("tract", "visit", PredicateVisitFlags(0))])
396 self.assertFalse(visitor.spatial_constraints)
397 self.assertFalse(visitor.temporal_dimension_joins)
398 # Add a "join operand" whose dimensions include both spatial families
399 # with the most fine-grained dimensions in the query.
400 # This blocks an automatic join from being created, because we assume
401 # that join operand (e.g. a materialization or dataset search) already
402 # encodes some spatial join.
403 visitor = self.run_visitor(
404 ["visit", "detector", "patch"],
405 qt.Predicate.from_bool(True),
406 "True",
407 join_operands=[self.universe.conform(["patch", "visit_detector_region"])],
408 )
409 self.assertFalse(visitor.spatial_joins)
410 self.assertFalse(visitor.spatial_constraints)
411 self.assertFalse(visitor.temporal_dimension_joins)
413 def test_multiple_spatial_joins(self) -> None:
414 """Test the overlaps visitor when there are >2 spatial families."""
415 x = ExpressionFactory(self.universe)
416 # Trivial predicate. This is an error, because we cannot generate
417 # automatic spatial joins when there are more than two families
418 with self.assertRaises(InvalidQueryError):
419 self.run_visitor(["visit", "patch", "htm7"], qt.Predicate.from_bool(True))
420 # Predicate that joins one pair of families but orphans the the other;
421 # also an error.
422 with self.assertRaises(InvalidQueryError):
423 self.run_visitor(["visit", "patch", "htm7"], x.visit.region.overlaps(x.htm7.region))
424 # A sufficient overlap join predicate has been added; each family is
425 # connected to at least one other.
426 visitor = self.run_visitor(
427 ["visit", "patch", "htm7"],
428 x.all(x.tract.region.overlaps(x.visit.region), x.tract.region.overlaps(x.htm7.region)),
429 "tract.region OVERLAPS visit.region AND tract.region OVERLAPS htm7.region",
430 )
431 self.assertEqual(
432 visitor.spatial_joins,
433 [
434 ("tract", "visit", PredicateVisitFlags.HAS_AND_SIBLINGS),
435 ("tract", "htm7", PredicateVisitFlags.HAS_AND_SIBLINGS),
436 ],
437 )
438 self.assertFalse(visitor.spatial_constraints)
439 self.assertFalse(visitor.temporal_dimension_joins)
440 # Add a "join operand" whose dimensions includes two spatial families,
441 # with the most fine-grained dimensions in the query, and a predicate
442 # that joins the third in.
443 visitor = self.run_visitor(
444 ["visit", "patch", "htm7"],
445 x.tract.region.overlaps(x.htm7.region),
446 "tract.region OVERLAPS htm7.region",
447 join_operands=[self.universe.conform(["visit", "patch"])],
448 )
449 self.assertEqual(
450 visitor.spatial_joins,
451 [
452 ("tract", "htm7", PredicateVisitFlags(0)),
453 ],
454 )
455 self.assertFalse(visitor.spatial_constraints)
456 self.assertFalse(visitor.temporal_dimension_joins)
458 def test_one_temporal_family(self) -> None:
459 """Test the overlaps visitor when there is one temporal family."""
460 x = ExpressionFactory(self.universe)
461 begin = astropy.time.Time("2020-01-01T00:00:00", format="isot", scale="tai")
462 end = astropy.time.Time("2020-01-01T00:01:00", format="isot", scale="tai")
463 timespan = Timespan(begin, end)
464 # Trivial predicate.
465 visitor = self.run_visitor(["exposure"], qt.Predicate.from_bool(True))
466 self.assertFalse(visitor.spatial_joins)
467 self.assertFalse(visitor.spatial_constraints)
468 self.assertFalse(visitor.temporal_dimension_joins)
469 # Non-overlap predicate.
470 visitor = self.run_visitor(["exposure"], x.any(x.band == "r", x.exposure > 2))
471 self.assertFalse(visitor.spatial_joins)
472 self.assertFalse(visitor.spatial_constraints)
473 self.assertFalse(visitor.temporal_dimension_joins)
474 # Temporal constraint predicate.
475 visitor = self.run_visitor(["exposure"], x.exposure.timespan.overlaps(timespan))
476 self.assertFalse(visitor.spatial_joins)
477 self.assertFalse(visitor.spatial_constraints)
478 self.assertFalse(visitor.temporal_dimension_joins)
479 # A temporal join between dimensions in the same family is an error.
480 with self.assertRaises(InvalidQueryError):
481 self.run_visitor(["exposure", "visit"], x.exposure.timespan.overlaps(x.visit.timespan))
482 # Overlap join with a calibration dataset's validity ranges.
483 visitor = self.run_visitor(
484 ["exposure"], x.exposure.timespan.overlaps(x["bias"].timespan), calibration_dataset_types={"bias"}
485 )
486 self.assertFalse(visitor.spatial_joins)
487 self.assertFalse(visitor.spatial_constraints)
488 self.assertFalse(visitor.temporal_dimension_joins)
489 self.assertEqual(
490 visitor.validity_range_dimension_joins, [("bias", "exposure", PredicateVisitFlags(0))]
491 )
492 self.assertFalse(visitor.validity_range_joins)
493 # Overlap join between two calibration dataset validity ranges.
494 # (It's not clear this kind of query is ever useful in practice, but
495 # there's a good consistency argument for what it ought to do).
496 visitor = self.run_visitor(
497 [], x["flat"].timespan.overlaps(x["bias"].timespan), calibration_dataset_types={"bias", "flat"}
498 )
499 self.assertFalse(visitor.spatial_joins)
500 self.assertFalse(visitor.spatial_constraints)
501 self.assertFalse(visitor.temporal_dimension_joins)
502 self.assertFalse(visitor.validity_range_dimension_joins)
503 self.assertEqual(visitor.validity_range_joins, [("flat", "bias", PredicateVisitFlags(0))])
505 # There are no tests for temporal dimension joins, because the default
506 # dimension universe only has one spatial family, and the untested logic
507 # trivially duplicates the spatial-join logic.
510class NaiveDisjointSetTestCase(unittest.TestCase):
511 """Test the naive disjoint-set implementation that backs automatic overlap
512 join creation.
513 """
515 def test_naive_disjoint_set(self) -> None:
516 s = _NaiveDisjointSet(range(8))
517 self.assertCountEqual(s.subsets(), [{n} for n in range(8)])
518 s.merge(3, 4)
519 self.assertCountEqual(s.subsets(), [{0}, {1}, {2}, {3, 4}, {5}, {6}, {7}])
520 s.merge(2, 1)
521 self.assertCountEqual(s.subsets(), [{0}, {1, 2}, {3, 4}, {5}, {6}, {7}])
522 s.merge(1, 3)
523 self.assertCountEqual(s.subsets(), [{0}, {1, 2, 3, 4}, {5}, {6}, {7}])
526if __name__ == "__main__":
527 unittest.main()