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

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

27 

28"""Tests for non-public Butler._query functionality that is not specific to 

29any Butler or QueryDriver implementation. 

30""" 

31 

32from __future__ import annotations 

33 

34import unittest 

35from collections.abc import Iterable, Set 

36 

37import astropy.time 

38 

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 

46 

47 

48class ColumnSetTestCase(unittest.TestCase): 

49 """Tests for lsst.daf.butler.queries.ColumnSet.""" 

50 

51 def setUp(self) -> None: 

52 self.universe = DimensionUniverse() 

53 

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

95 

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) 

113 

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) 

142 

143 

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]] = [] 

152 

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) 

158 

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) 

164 

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) 

170 

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) 

176 

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) 

180 

181 

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

187 

188 def setUp(self) -> None: 

189 self.universe = DimensionUniverse() 

190 

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 

205 

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) 

221 

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

229 

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

292 

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) 

359 

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) 

412 

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) 

457 

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

504 

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. 

508 

509 

510class NaiveDisjointSetTestCase(unittest.TestCase): 

511 """Test the naive disjoint-set implementation that backs automatic overlap 

512 join creation. 

513 """ 

514 

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

524 

525 

526if __name__ == "__main__": 

527 unittest.main()