Coverage for python/lsst/dax/apdb/tests/_apdb.py: 99%
714 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-29 09:48 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-29 09:48 +0000
1# This file is part of dax_apdb.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (http://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 program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <http://www.gnu.org/licenses/>.
22from __future__ import annotations
24__all__ = ["ApdbSchemaUpdateTest", "ApdbTest", "update_schema_yaml"]
26import contextlib
27import itertools
28import logging.config
29import os
30import tempfile
31from abc import ABC, abstractmethod
32from collections.abc import Iterator
33from tempfile import TemporaryDirectory
34from typing import TYPE_CHECKING, Any
36import astropy.time
37import felis.datamodel
38import numpy
39import pandas
40import yaml
42from lsst.sphgeom import Angle, Circle, LonLat, Region, UnitVector3d
44from .. import (
45 Apdb,
46 ApdbConfig,
47 ApdbReassignDiaSourceToSSObjectRecord,
48 ApdbReplica,
49 ApdbTableData,
50 ApdbTables,
51 ApdbUpdateRecord,
52 ApdbWithdrawDiaSourceRecord,
53 DiaForcedSourceId,
54 DiaObjectId,
55 DiaSourceId,
56 IncompatibleVersionError,
57 ReplicaChunk,
58 VersionTuple,
59)
60from .data_factory import (
61 makeForcedSourceCatalog,
62 makeObjectCatalog,
63 makeSourceCatalog,
64 makeTimestamp,
65 makeTimestampColumn,
66)
67from .utils import TestCaseMixin
69if TYPE_CHECKING:
70 from ..pixelization import Pixelization
73# Optionally configure logging from a config file.
74if log_config := os.environ.get("DAX_APDB_TEST_LOG_CONFIG"): 74 ↛ 75line 74 didn't jump to line 75 because the condition on line 74 was never true
75 logging.config.fileConfig(log_config)
78def _make_region(xyz: tuple[float, float, float] | LonLat = (1.0, 1.0, -1.0)) -> Region:
79 """Make a region to use in tests"""
80 if isinstance(xyz, LonLat):
81 pointing_v = UnitVector3d(xyz)
82 else:
83 pointing_v = UnitVector3d(*xyz)
84 fov = 0.0013 # radians
85 region = Circle(pointing_v, Angle(fov / 2))
86 return region
89@contextlib.contextmanager
90def update_schema_yaml(
91 schema_file: str,
92 drop_metadata: bool = False,
93 version: str | None = None,
94) -> Iterator[str]:
95 """Update schema definition and return name of the new schema file.
97 Parameters
98 ----------
99 schema_file : `str`
100 Path for the existing YAML file with APDB schema.
101 drop_metadata : `bool`
102 If `True` then remove metadata table from the list of tables.
103 version : `str` or `None`
104 If non-empty string then set schema version to this string, if empty
105 string then remove schema version from config, if `None` - don't change
106 the version in config.
108 Yields
109 ------
110 Path for the updated configuration file.
111 """
112 with open(schema_file) as yaml_stream:
113 schemas_list = list(yaml.load_all(yaml_stream, Loader=yaml.SafeLoader))
114 # Edit YAML contents.
115 for schema in schemas_list:
116 # Optionally drop metadata table.
117 if drop_metadata:
118 schema["tables"] = [table for table in schema["tables"] if table["name"] != "metadata"]
119 if version is not None:
120 if version == "":
121 del schema["version"]
122 else:
123 schema["version"] = version
125 with TemporaryDirectory(ignore_cleanup_errors=True) as tmpdir:
126 output_path = os.path.join(tmpdir, "schema.yaml")
127 with open(output_path, "w") as yaml_stream:
128 yaml.dump_all(schemas_list, stream=yaml_stream)
129 yield output_path
132class ApdbTest(TestCaseMixin, ABC):
133 """Base class for Apdb tests that can be specialized for concrete
134 implementation.
136 This can only be used as a mixin class for a unittest.TestCase and it
137 calls various assert methods.
138 """
140 visit_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
142 processing_time = astropy.time.Time("2021-01-01T12:00:00", format="isot", scale="tai")
144 fsrc_requires_id_list = False
145 """Should be set to True if getDiaForcedSources requires object IDs"""
147 enable_replica: bool = False
148 """Set to true when support for replication is configured"""
150 use_mjd: bool = True
151 """If True then timestamp columns are MJD TAI."""
153 extra_chunk_columns = 1
154 """Number of additional columns in chunk tables."""
156 meta_row_count = 3
157 """Initial row count in metadata table."""
159 # number of columns as defined in tests/config/schema.yaml
160 table_column_count = {
161 ApdbTables.DiaObject: 8,
162 ApdbTables.DiaObjectLast: 6,
163 ApdbTables.DiaSource: 13,
164 ApdbTables.DiaForcedSource: 9,
165 ApdbTables.SSObject: 3,
166 }
168 @abstractmethod
169 def make_instance(self, **kwargs: Any) -> ApdbConfig:
170 """Make database instance and return configuration for it."""
171 raise NotImplementedError()
173 @abstractmethod
174 def getDiaObjects_table(self) -> ApdbTables:
175 """Return type of table returned from getDiaObjects method."""
176 raise NotImplementedError()
178 @abstractmethod
179 def pixelization(self, config: ApdbConfig) -> Pixelization:
180 """Return pixelization used by implementation."""
181 raise NotImplementedError()
183 def assert_catalog(self, catalog: Any, rows: int, table: ApdbTables) -> None:
184 """Validate catalog type and size
186 Parameters
187 ----------
188 catalog : `object`
189 Expected type of this is ``pandas.DataFrame``.
190 rows : `int`
191 Expected number of rows in a catalog.
192 table : `ApdbTables`
193 APDB table type.
194 """
195 self.assertIsInstance(catalog, pandas.DataFrame)
196 self.assertEqual(catalog.shape[0], rows)
197 self.assertEqual(catalog.shape[1], self.table_column_count[table])
199 def assert_table_data(self, catalog: Any, rows: int, table: ApdbTables) -> None:
200 """Validate catalog type and size
202 Parameters
203 ----------
204 catalog : `object`
205 Expected type of this is `ApdbTableData`.
206 rows : `int`
207 Expected number of rows in a catalog.
208 table : `ApdbTables`
209 APDB table type.
210 extra_columns : `int`
211 Count of additional columns expected in ``catalog``.
212 """
213 self.assertIsInstance(catalog, ApdbTableData)
214 n_rows = sum(1 for row in catalog.rows())
215 self.assertEqual(n_rows, rows)
216 # One extra column for replica chunk id
217 self.assertEqual(
218 len(catalog.column_names()), self.table_column_count[table] + self.extra_chunk_columns
219 )
221 def assert_column_types(self, catalog: Any, types: dict[str, felis.datamodel.DataType]) -> None:
222 column_defs = dict(catalog.column_defs())
223 for column, datatype in types.items():
224 self.assertEqual(column_defs[column], datatype)
226 def make_region(self, xyz: tuple[float, float, float] | LonLat = (1.0, 1.0, -1.0)) -> Region:
227 """Make a region to use in tests"""
228 return _make_region(xyz)
230 def test_makeSchema(self) -> None:
231 """Test for making APDB schema."""
232 config = self.make_instance()
233 apdb = Apdb.from_config(config)
235 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObject))
236 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObjectLast))
237 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaSource))
238 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaForcedSource))
239 self.assertIsNotNone(apdb.tableDef(ApdbTables.metadata))
240 self.assertIsNotNone(apdb.tableDef(ApdbTables.SSObject))
241 self.assertIsNotNone(apdb.tableDef(ApdbTables.SSSource))
242 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObject_To_Object_Match))
244 # Test from_uri factory method with the same config.
245 with tempfile.NamedTemporaryFile() as tmpfile:
246 config.save(tmpfile.name)
247 apdb = Apdb.from_uri(tmpfile.name)
249 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObject))
250 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObjectLast))
251 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaSource))
252 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaForcedSource))
253 self.assertIsNotNone(apdb.tableDef(ApdbTables.metadata))
254 self.assertIsNotNone(apdb.tableDef(ApdbTables.SSObject))
255 self.assertIsNotNone(apdb.tableDef(ApdbTables.SSSource))
256 self.assertIsNotNone(apdb.tableDef(ApdbTables.DiaObject_To_Object_Match))
258 def test_empty_gets(self) -> None:
259 """Test for getting data from empty database.
261 All get() methods should return empty results, only useful for
262 checking that code is not broken.
263 """
264 # use non-zero months for Forced/Source fetching
265 config = self.make_instance()
266 apdb = Apdb.from_config(config)
268 region = self.make_region()
269 visit_time = self.visit_time
271 res: pandas.DataFrame | None
273 # get objects by region
274 res = apdb.getDiaObjects(region)
275 self.assert_catalog(res, 0, self.getDiaObjects_table())
277 # get sources by region
278 res = apdb.getDiaSources(region, None, visit_time)
279 self.assert_catalog(res, 0, ApdbTables.DiaSource)
281 res = apdb.getDiaSources(region, [], visit_time)
282 self.assert_catalog(res, 0, ApdbTables.DiaSource)
284 # get sources by object ID, non-empty object list
285 res = apdb.getDiaSources(region, [1, 2, 3], visit_time)
286 self.assert_catalog(res, 0, ApdbTables.DiaSource)
288 # get forced sources by object ID, empty object list
289 res = apdb.getDiaForcedSources(region, [], visit_time)
290 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
292 # get sources by object ID, non-empty object list
293 res = apdb.getDiaForcedSources(region, [1, 2, 3], visit_time)
294 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
296 # data_factory's ccdVisitId generation corresponds to (1, 1)
297 res = apdb.containsVisitDetector(visit=1, detector=1)
298 self.assertFalse(res)
300 # get sources by region
301 if self.fsrc_requires_id_list: 301 ↛ 305line 301 didn't jump to line 305 because the condition on line 301 was always true
302 with self.assertRaises(NotImplementedError):
303 apdb.getDiaForcedSources(region, None, visit_time)
304 else:
305 res = apdb.getDiaForcedSources(region, None, visit_time)
306 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
308 def test_empty_gets_0months(self) -> None:
309 """Test for getting data from empty database.
311 All get() methods should return empty DataFrame or None.
312 """
313 # set read_sources_months to 0 so that Forced/Sources are None
314 config = self.make_instance(read_sources_months=0, read_forced_sources_months=0)
315 apdb = Apdb.from_config(config)
317 region = self.make_region()
318 visit_time = self.visit_time
320 res: pandas.DataFrame | None
322 # get objects by region
323 res = apdb.getDiaObjects(region)
324 self.assert_catalog(res, 0, self.getDiaObjects_table())
326 # get sources by region
327 res = apdb.getDiaSources(region, None, visit_time)
328 self.assertIs(res, None)
330 # get sources by object ID, empty object list
331 res = apdb.getDiaSources(region, [], visit_time)
332 self.assertIs(res, None)
334 # get forced sources by object ID, empty object list
335 res = apdb.getDiaForcedSources(region, [], visit_time)
336 self.assertIs(res, None)
338 # Database is empty, no images exist.
339 res = apdb.containsVisitDetector(visit=1, detector=1)
340 self.assertFalse(res)
342 def test_storeObjects(self) -> None:
343 """Store and retrieve DiaObjects."""
344 # don't care about sources.
345 config = self.make_instance()
346 apdb = Apdb.from_config(config)
348 region = self.make_region()
349 visit_time = self.visit_time
351 # make catalog with Objects
352 catalog = makeObjectCatalog(region, 100)
354 # store catalog
355 apdb.store(visit_time, catalog)
357 # read it back and check sizes
358 res = apdb.getDiaObjects(region)
359 self.assert_catalog(res, len(catalog), self.getDiaObjects_table())
361 # TODO: test apdb.contains with generic implementation from DM-41671
363 def test_storeObjects_empty(self) -> None:
364 """Test calling storeObject when there are no objects: see DM-43270."""
365 config = self.make_instance()
366 apdb = Apdb.from_config(config)
367 region = self.make_region()
368 visit_time = self.visit_time
369 # make catalog with no Objects
370 catalog = makeObjectCatalog(region, 0)
372 with self.assertLogs("lsst.dax.apdb", level="DEBUG") as cm:
373 apdb.store(visit_time, catalog)
374 self.assertIn("No objects", "\n".join(cm.output))
376 def test_storeMovingObject(self) -> None:
377 """Store and retrieve DiaObject which changes its position."""
378 # don't care about sources.
379 config = self.make_instance()
380 apdb = Apdb.from_config(config)
381 pixelization = self.pixelization(config)
383 lon_deg, lat_deg = 0.0, 0.0
384 lonlat1 = LonLat.fromDegrees(lon_deg - 1.0, lat_deg)
385 lonlat2 = LonLat.fromDegrees(lon_deg + 1.0, lat_deg)
386 uv1 = UnitVector3d(lonlat1)
387 uv2 = UnitVector3d(lonlat2)
389 # Check that they fall into different pixels.
390 self.assertNotEqual(pixelization.pixel(uv1), pixelization.pixel(uv2))
392 # Store one object at two different positions.
393 visit_time1 = self.visit_time
394 catalog1 = makeObjectCatalog(lonlat1, 1)
395 apdb.store(visit_time1, catalog1)
397 visit_time2 = visit_time1 + astropy.time.TimeDelta(120.0, format="sec")
398 catalog1 = makeObjectCatalog(lonlat2, 1)
399 apdb.store(visit_time2, catalog1)
401 # Make region covering both points.
402 region = Circle(UnitVector3d(LonLat.fromDegrees(lon_deg, lat_deg)), Angle.fromDegrees(1.1))
403 self.assertTrue(region.contains(uv1))
404 self.assertTrue(region.contains(uv2))
406 # Read it back, must return the latest one.
407 res = apdb.getDiaObjects(region)
408 self.assert_catalog(res, 1, self.getDiaObjects_table())
410 def test_storeSources(self) -> None:
411 """Store and retrieve DiaSources."""
412 config = self.make_instance()
413 apdb = Apdb.from_config(config)
415 region = self.make_region()
416 visit_time = self.visit_time
418 # have to store Objects first
419 objects = makeObjectCatalog(region, 100)
420 oids = list(objects["diaObjectId"])
421 sources = makeSourceCatalog(objects, visit_time, use_mjd=self.use_mjd)
423 # save the objects and sources
424 apdb.store(visit_time, objects, sources)
426 # read it back, no ID filtering
427 res = apdb.getDiaSources(region, None, visit_time)
428 self.assert_catalog(res, len(sources), ApdbTables.DiaSource)
430 # read it back and filter by ID
431 res = apdb.getDiaSources(region, oids, visit_time)
432 self.assert_catalog(res, len(sources), ApdbTables.DiaSource)
434 # read it back to get schema
435 res = apdb.getDiaSources(region, [], visit_time)
436 self.assert_catalog(res, 0, ApdbTables.DiaSource)
438 # test if a visit is present
439 # data_factory's ccdVisitId generation corresponds to (1, 1)
440 res = apdb.containsVisitDetector(visit=1, detector=1)
441 self.assertTrue(res)
442 # non-existent image
443 res = apdb.containsVisitDetector(visit=2, detector=42)
444 self.assertFalse(res)
446 def test_storeForcedSources(self) -> None:
447 """Store and retrieve DiaForcedSources."""
448 config = self.make_instance()
449 apdb = Apdb.from_config(config)
451 region = self.make_region()
452 visit_time = self.visit_time
454 # have to store Objects first
455 objects = makeObjectCatalog(region, 100)
456 oids = list(objects["diaObjectId"])
457 catalog = makeForcedSourceCatalog(objects, visit_time, use_mjd=self.use_mjd)
459 apdb.store(visit_time, objects, forced_sources=catalog)
461 # read it back and check sizes
462 res = apdb.getDiaForcedSources(region, oids, visit_time)
463 self.assert_catalog(res, len(catalog), ApdbTables.DiaForcedSource)
465 # read it back to get schema
466 res = apdb.getDiaForcedSources(region, [], visit_time)
467 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
469 # data_factory's ccdVisitId generation corresponds to (1, 1)
470 res = apdb.containsVisitDetector(visit=1, detector=1)
471 self.assertTrue(res)
472 # non-existent image
473 res = apdb.containsVisitDetector(visit=2, detector=42)
474 self.assertFalse(res)
476 def test_null_integer_type(self) -> None:
477 """Test that integer column with NULLs correct type on select."""
478 config = self.make_instance()
479 apdb = Apdb.from_config(config)
481 region = self.make_region()
482 visit_time = self.visit_time
484 # have to store Objects first
485 objects = makeObjectCatalog(region, 100)
486 sources = makeSourceCatalog(objects, visit_time, use_mjd=self.use_mjd)
487 # Reset some diaObjectIds to NULL.
488 sources.loc[0:10, "diaObjectId"] = None
490 # save the objects and sources
491 apdb.store(visit_time, objects, sources)
493 # read it back, no ID filtering
494 res = apdb.getDiaSources(region, None, visit_time)
495 self.assert_catalog(res, len(sources), ApdbTables.DiaSource)
496 assert res is not None, "Expecting catalog, not None"
497 self.assertEqual(res.dtypes["diaObjectId"], pandas.Int64Dtype())
499 def test_timestamps(self) -> None:
500 """Check that timestamp return type is as expected."""
501 config = self.make_instance()
502 apdb = Apdb.from_config(config)
504 region = self.make_region()
505 visit_time = self.visit_time
507 # Cassandra has a millisecond precision, so subtract 1ms to allow for
508 # truncated returned values.
509 time_before = makeTimestamp(self.processing_time, self.use_mjd, -1)
510 objects = makeObjectCatalog(region, 100)
511 oids = list(objects["diaObjectId"])
512 catalog = makeForcedSourceCatalog(
513 objects, visit_time, processing_time=self.processing_time, use_mjd=self.use_mjd
514 )
515 time_after = makeTimestamp(self.processing_time, self.use_mjd)
517 apdb.store(visit_time, objects, forced_sources=catalog)
519 # read it back and check sizes
520 res = apdb.getDiaForcedSources(region, oids, visit_time)
521 assert res is not None
522 self.assert_catalog(res, len(catalog), ApdbTables.DiaForcedSource)
524 time_processed_column = makeTimestampColumn("time_processed", self.use_mjd)
525 self.assertIn(time_processed_column, res.dtypes)
526 dtype = res.dtypes[time_processed_column]
527 timestamp_type_names = (
528 ("float64",) if self.use_mjd else ("datetime64[ms]", "datetime64[us]", "datetime64[ns]")
529 )
530 self.assertIn(dtype.name, timestamp_type_names)
531 # Verify that returned time is sensible.
532 self.assertTrue(all(time_before <= dt <= time_after for dt in res[time_processed_column]))
534 def test_getDiaObjectsForDedup(self) -> None:
535 """Test getDiaObjectsForDedup() method."""
536 config = self.make_instance()
537 apdb = Apdb.from_config(config)
539 region1 = self.make_region((1.0, 1.0, -1.0))
540 region2 = self.make_region((-1.0, 1.0, -1.0))
541 region3 = self.make_region((-1.0, -1.0, -1.0))
542 nobj = 100
543 objects1 = makeObjectCatalog(region1, nobj)
544 objects2 = makeObjectCatalog(region2, nobj, start_id=nobj * 2)
545 objects3 = makeObjectCatalog(region3, nobj, start_id=nobj * 4)
547 visits = [
548 (astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai"), objects1),
549 (astropy.time.Time("2021-01-01T00:10:00", format="isot", scale="tai"), objects2),
550 (astropy.time.Time("2021-01-01T00:20:00", format="isot", scale="tai"), objects3),
551 ]
553 for visit_time, objects in visits:
554 apdb.store(visit_time, objects)
556 catalog = apdb.getDiaObjectsForDedup()
557 self.assertEqual(len(catalog), 300)
559 catalog = apdb.getDiaObjectsForDedup(visits[0][0])
560 self.assertEqual(len(catalog), 300)
562 catalog = apdb.getDiaObjectsForDedup(visits[1][0])
563 self.assertEqual(len(catalog), 200)
565 catalog = apdb.getDiaObjectsForDedup(visits[2][0])
566 self.assertEqual(len(catalog), 100)
568 time = astropy.time.Time("2021-01-01T00:30:00", format="isot", scale="tai")
569 catalog = apdb.getDiaObjectsForDedup(time)
570 self.assertEqual(len(catalog), 0)
572 def test_getDiaSourcesForDiaObjects(self) -> None:
573 """Test getDiaSourcesForDiaObjects() method."""
574 config = self.make_instance()
575 apdb = Apdb.from_config(config)
576 # Monkey-patch APDB instance to set current time.
577 apdb._current_time = lambda: self.processing_time # type: ignore[method-assign]
579 region1 = self.make_region((1.0, 1.0, -1.0))
580 region2 = self.make_region((-1.0, 1.0, -1.0))
581 region3 = self.make_region((-1.0, -1.0, -1.0))
582 nobj = 100
583 objects1 = makeObjectCatalog(region1, nobj)
584 objects2 = makeObjectCatalog(region2, nobj, start_id=nobj * 2)
585 objects3 = makeObjectCatalog(region3, nobj, start_id=nobj * 4)
587 visits = [
588 (astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai"), objects1),
589 (astropy.time.Time("2021-01-01T00:10:00", format="isot", scale="tai"), objects2),
590 (astropy.time.Time("2021-01-01T00:20:00", format="isot", scale="tai"), objects3),
591 ]
593 start_id = 1_000_000
594 for visit_time, objects in visits:
595 sources = makeSourceCatalog(objects, visit_time, start_id=start_id, use_mjd=self.use_mjd)
596 apdb.store(visit_time, objects, sources)
597 start_id += 1_000_000
599 # Take a small number of objects from different regions.
600 object_ids = [
601 DiaObjectId.from_named_tuple(next(objects1.itertuples())),
602 DiaObjectId.from_named_tuple(next(objects2.itertuples())),
603 DiaObjectId.from_named_tuple(next(objects3.itertuples())),
604 ]
606 catalog = apdb.getDiaSourcesForDiaObjects(object_ids, visits[0][0])
607 self.assertEqual(len(catalog), 3)
608 self.assertEqual(set(catalog["diaObjectId"]), {1, 200, 400})
609 self.assertEqual(set(catalog["diaSourceId"]), {1_000_000, 2_000_000, 3_000_000})
611 catalog = apdb.getDiaSourcesForDiaObjects(object_ids, visits[2][0])
612 self.assertEqual(len(catalog), 1)
613 self.assertEqual(set(catalog["diaObjectId"]), {400})
614 self.assertEqual(set(catalog["diaSourceId"]), {3_000_000})
616 def test_reassignDiaSourcesToDiaObjects(self) -> None:
617 """Test reassignDiaSourcesToDiaObjects() method."""
618 config = self.make_instance()
619 apdb = Apdb.from_config(config)
620 apdb._current_time = lambda: self.processing_time # type: ignore[method-assign]
621 apdb_replica = ApdbReplica.from_config(config)
623 visit_time = self.visit_time
624 lonlat1 = LonLat.fromDegrees(0.0, 0.0)
625 lonlat2 = LonLat.fromDegrees(180.0, 0.0)
626 # regons around lonlat1/2
627 region1 = self.make_region(lonlat1)
628 region2 = self.make_region(lonlat2)
630 # Store 3 objects and sources at the same position in each region.
631 objects = makeObjectCatalog(lonlat1, 3, start_id=100)
632 sources = makeSourceCatalog(objects, visit_time, start_id=1000, use_mjd=self.use_mjd)
633 apdb.store(visit_time, objects, sources)
635 objects = makeObjectCatalog(lonlat2, 3, start_id=200)
636 sources = makeSourceCatalog(objects, visit_time, start_id=2000, use_mjd=self.use_mjd)
637 apdb.store(visit_time, objects, sources)
639 # check that everything as we think it is.
640 objects = apdb.getDiaObjects(region1)
641 self.assertEqual(set(objects["diaObjectId"]), {100, 101, 102})
642 self.assertEqual(list(objects["nDiaSources"]), [1, 1, 1])
643 sources = apdb.getDiaSources(region1, [100, 101, 102], visit_time)
644 assert sources is not None
645 self.assertEqual(set(sources["diaSourceId"]), {1000, 1001, 1002})
646 self.assertEqual(set(sources["diaObjectId"]), {100, 101, 102})
648 dia_source_ids = [DiaSourceId.from_named_tuple(row) for row in sources.itertuples()]
650 # Reassign sources in region1 and increment/decrement nDiaSources.
651 reassign = {
652 dia_source_id: 100
653 for dia_source_id in dia_source_ids
654 if dia_source_id.diaSourceId in (1001, 1002)
655 }
656 apdb.reassignDiaSourcesToDiaObjects(reassign)
658 objects = apdb.getDiaObjects(region1)
659 self.assertEqual(set(objects["nDiaSources"]), {0, 3})
660 sources = apdb.getDiaSources(region1, [100], visit_time)
661 assert sources is not None
662 self.assertEqual(set(sources["diaSourceId"]), {1000, 1001, 1002})
663 self.assertEqual(set(sources["diaObjectId"]), {100})
665 sources = apdb.getDiaSources(region2, [201, 202], visit_time)
666 assert sources is not None
667 self.assertEqual(set(sources["diaSourceId"]), {2001, 2002})
668 dia_source_ids = [DiaSourceId.from_named_tuple(row) for row in sources.itertuples()]
670 # Reassign but do not increment/decrement nDiaSources.
671 reassign = {
672 dia_source_id: 200
673 for dia_source_id in dia_source_ids
674 if dia_source_id.diaSourceId in (2001, 2002)
675 }
676 apdb.reassignDiaSourcesToDiaObjects(
677 reassign, increment_nDiaSources=False, decrement_nDiaSources=False
678 )
680 objects = apdb.getDiaObjects(region2)
681 self.assertEqual(set(objects["nDiaSources"]), {1})
682 sources = apdb.getDiaSources(region2, [200], visit_time)
683 assert sources is not None
684 self.assertEqual(set(sources["diaSourceId"]), {2000, 2001, 2002})
685 self.assertEqual(set(sources["diaObjectId"]), {200})
687 replica_chunks = apdb_replica.getReplicaChunks()
688 if not self.enable_replica:
689 self.assertIsNone(replica_chunks)
690 else:
691 assert replica_chunks is not None
693 # There could be one or two chunks.
694 self.assertTrue(1 <= len(replica_chunks) <= 2)
696 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks])
697 # Two reassignments for region1, three increments/decrements for
698 # that region, plus two reassignments for region2 without
699 # increments/decrements.
700 self.assertEqual(len(update_records), 2 + 3 + 2)
702 def test_setValidityEnd(self) -> None:
703 """Store DiaObjects and truncate validity for some."""
704 # don't care about sources.
705 config = self.make_instance()
706 apdb = Apdb.from_config(config)
707 apdb._current_time = lambda: self.processing_time # type: ignore[method-assign]
708 apdb_replica = ApdbReplica.from_config(config)
710 region = self.make_region()
711 visit_time = self.visit_time
713 # make catalog with Objects
714 catalog = makeObjectCatalog(region, 100)
716 # store catalog
717 apdb.store(visit_time, catalog)
719 # read it back and check sizes
720 res = apdb.getDiaObjects(region)
721 self.assert_catalog(res, 100, self.getDiaObjects_table())
723 # Select first 10 objects.
724 object_ids = [DiaObjectId.from_named_tuple(row) for row in catalog.iloc[:10].itertuples()]
725 count = apdb.setValidityEnd(object_ids, self.processing_time)
726 self.assertEqual(count, 10)
728 res = apdb.getDiaObjects(region)
729 self.assert_catalog(res, 90, self.getDiaObjects_table())
731 replica_chunks = apdb_replica.getReplicaChunks()
732 if not self.enable_replica:
733 self.assertIsNone(replica_chunks)
734 else:
735 # Check that there are 10 update records in replica tables.
736 assert replica_chunks is not None
738 # There could be one or two chunks.
739 self.assertTrue(1 <= len(replica_chunks) <= 2)
741 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks])
742 self.assertEqual(len(update_records), 10)
744 # Check that empty list works.
745 count = apdb.setValidityEnd(object_ids, self.processing_time)
746 self.assertEqual(count, 0)
748 # Try with non-existing object.
749 object_ids = [DiaObjectId.from_named_tuple(row) for row in catalog.iloc[10:12].itertuples()]
750 object_ids += [DiaObjectId(diaObjectId=1_000_000, ra=0.0, dec=0.0)]
751 with self.assertRaises(LookupError):
752 apdb.setValidityEnd(object_ids, self.processing_time, raise_on_missing_id=True)
754 count = apdb.setValidityEnd(object_ids, self.processing_time)
755 self.assertEqual(count, 2)
757 def test_resetDedup(self) -> None:
758 """Test resetDedup method."""
759 # don't care about sources.
760 config = self.make_instance()
761 apdb = Apdb.from_config(config)
763 region = self.make_region()
765 # make catalog with Objects
766 objects = makeObjectCatalog(region, 100)
768 visit_time1 = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
769 dedup_time1 = astropy.time.Time("2021-01-01T12:00:00", format="isot", scale="tai")
770 visit_time2 = astropy.time.Time("2021-01-02T00:00:00", format="isot", scale="tai")
771 dedup_time2 = astropy.time.Time("2021-01-02T12:00:00", format="isot", scale="tai")
773 # store catalog
774 apdb.store(visit_time1, objects)
776 catalog = apdb.getDiaObjectsForDedup()
777 self.assertEqual(len(catalog), 100)
779 catalog = apdb.getDiaObjectsForDedup(visit_time1)
780 self.assertEqual(len(catalog), 100)
782 apdb.resetDedup(dedup_time1)
784 catalog = apdb.getDiaObjectsForDedup(visit_time1)
785 self.assertEqual(len(catalog), self._count_after_reset_dedup(100))
787 apdb.store(visit_time2, objects)
789 catalog = apdb.getDiaObjectsForDedup()
790 self.assertEqual(len(catalog), 100)
792 catalog = apdb.getDiaObjectsForDedup(dedup_time1)
793 self.assertEqual(len(catalog), 100)
795 apdb.resetDedup(dedup_time2)
797 catalog = apdb.getDiaObjectsForDedup(dedup_time1)
798 self.assertEqual(len(catalog), self._count_after_reset_dedup(100))
800 catalog = apdb.getDiaObjectsForDedup()
801 self.assertEqual(len(catalog), 0)
803 def _count_after_reset_dedup(self, count_before: int) -> int:
804 """Return the number of rows that will be returned by
805 getDiaObjectsForDedup() after resetDedup() was called. For SQL backend
806 deduplication data comes from a regular table, and it is not removed
807 by resetDedup().
808 """
809 raise NotImplementedError()
811 def test_withdraw_sources(self) -> None:
812 """Test withdrawDiaSources() method."""
813 config = self.make_instance()
814 apdb = Apdb.from_config(config)
815 apdb._current_time = lambda: self.processing_time # type: ignore[method-assign]
816 apdb_replica = ApdbReplica.from_config(config)
818 lonlat1 = LonLat.fromDegrees(0.0, 0.0)
819 region1 = self.make_region(lonlat1)
820 lonlat2 = LonLat.fromDegrees(45.0, 0.0)
821 region2 = self.make_region(lonlat2)
823 # Store 3 objects and sources at the same position in each region.
824 # The code originally updated nDiaSources so there are objects with
825 # nDiaSources > 1, but we dropped that option.
826 visit_time1 = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
827 objects1 = makeObjectCatalog(lonlat1, 3, start_id=100)
828 sources1 = makeSourceCatalog(objects1, visit_time1, start_id=1000, use_mjd=self.use_mjd)
829 apdb.store(visit_time1, objects1, sources1)
831 visit_time2 = astropy.time.Time("2021-01-01T00:01:00", format="isot", scale="tai")
832 objects2 = makeObjectCatalog(lonlat2, 3, start_id=200)
833 sources2 = makeSourceCatalog(objects2, visit_time2, start_id=2000, use_mjd=self.use_mjd)
834 apdb.store(visit_time2, objects2, sources2)
836 # Fetch everything and verify.
837 objects1 = apdb.getDiaObjects(region1)
838 objects2 = apdb.getDiaObjects(region2)
839 self.assertEqual({row.diaObjectId for row in objects1.itertuples()}, {100, 101, 102})
840 self.assertEqual({row.diaObjectId for row in objects2.itertuples()}, {200, 201, 202})
842 sources1 = apdb.getDiaSources(region1, None, visit_time2)
843 sources2 = apdb.getDiaSources(region2, None, visit_time2)
844 assert sources1 is not None and sources2 is not None
845 source_ids = [
846 DiaSourceId.from_named_tuple(row)
847 for row in itertools.chain(sources1.itertuples(), sources2.itertuples())
848 ]
849 sources_by_id = {source_id.diaSourceId: source_id for source_id in source_ids}
851 if self.use_mjd:
852 self.assertTrue(all(pandas.isna(sources1["timeWithdrawnMjdTai"])))
853 self.assertTrue(all(pandas.isna(sources2["timeWithdrawnMjdTai"])))
854 else:
855 self.assertTrue(all(pandas.isnull(sources1["time_withdrawn"])))
856 self.assertTrue(all(pandas.isnull(sources2["time_withdrawn"])))
858 # Withdraw a bunch of sources.
859 withdraw_time1 = astropy.time.Time("2021-01-01T10:00:00", format="isot", scale="tai")
860 withdraw_time2 = astropy.time.Time("2021-01-02T10:00:00", format="isot", scale="tai")
861 apdb.withdrawDiaSources([sources_by_id[i] for i in (1000, 2001)], timeWithdrawn=withdraw_time1)
862 # Withdraw diaSourceId=1000 second time, should not change it.
863 apdb.withdrawDiaSources([sources_by_id[i] for i in (1000, 2000, 1002)], timeWithdrawn=withdraw_time2)
865 sources1 = apdb.getDiaSources(region1, None, visit_time2)
866 sources2 = apdb.getDiaSources(region2, None, visit_time2)
867 assert sources1 is not None and sources2 is not None
868 sources1.set_index("diaSourceId", inplace=True)
869 sources2.set_index("diaSourceId", inplace=True)
870 if self.use_mjd:
871 self.assertEqual(sources1.loc[1000, "timeWithdrawnMjdTai"], withdraw_time1.mjd)
872 self.assertTrue(numpy.isnan(sources1.loc[1001, "timeWithdrawnMjdTai"]))
873 self.assertEqual(sources1.loc[1002, "timeWithdrawnMjdTai"], withdraw_time2.mjd)
874 self.assertEqual(sources2.loc[2000, "timeWithdrawnMjdTai"], withdraw_time2.mjd)
875 self.assertEqual(sources2.loc[2001, "timeWithdrawnMjdTai"], withdraw_time1.mjd)
876 self.assertTrue(numpy.isnan(sources2.loc[2002, "timeWithdrawnMjdTai"]))
877 else:
878 # Exact type and values depend on backend, I don't want to
879 # overcomplicate it, just check for NaT.
880 self.assertEqual(
881 dict(numpy.isnat(sources1["time_withdrawn"])), {1000: False, 1001: True, 1002: False}
882 )
883 self.assertEqual(
884 dict(numpy.isnat(sources2["time_withdrawn"])), {2000: False, 2001: False, 2002: True}
885 )
887 # Check replication update tables.
888 replica_chunks = apdb_replica.getReplicaChunks()
889 if not self.enable_replica:
890 self.assertIsNone(replica_chunks)
891 else:
892 # Check that there are 4 update records in replica tables.
893 assert replica_chunks is not None
895 # There could be one or two chunks.
896 self.assertTrue(1 <= len(replica_chunks) <= 2)
898 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks])
899 self.assertEqual(len(update_records), 4)
901 def test_withdraw_forced_sources(self) -> None:
902 """Test withdrawDiaForcedSources() method."""
903 config = self.make_instance()
904 apdb = Apdb.from_config(config)
905 apdb._current_time = lambda: self.processing_time # type: ignore[method-assign]
906 apdb_replica = ApdbReplica.from_config(config)
908 lonlat = LonLat.fromDegrees(0.0, 0.0)
909 region = self.make_region(lonlat)
911 # Store 3 objects and sources at the same position in each region.
912 objects = makeObjectCatalog(lonlat, 3, start_id=100)
913 fsources = makeForcedSourceCatalog(objects, self.visit_time, use_mjd=self.use_mjd)
914 apdb.store(self.visit_time, objects, None, fsources)
916 fsources = apdb.getDiaForcedSources(region, [100, 101, 102], self.visit_time)
917 assert fsources is not None
918 source_ids = [DiaForcedSourceId.from_named_tuple(row) for row in fsources.itertuples()]
919 sources_by_id = {source_id.diaObjectId: source_id for source_id in source_ids}
921 if self.use_mjd:
922 self.assertTrue(all(pandas.isna(fsources["timeWithdrawnMjdTai"])))
923 else:
924 self.assertTrue(all(pandas.isnull(fsources["time_withdrawn"])))
926 # Withdraw sources.
927 withdraw_time1 = astropy.time.Time("2021-01-01T10:00:00", format="isot", scale="tai")
928 apdb.withdrawDiaForcedSources([sources_by_id[i] for i in (100, 102)], timeWithdrawn=withdraw_time1)
929 withdraw_time2 = astropy.time.Time("2021-01-02T10:00:00", format="isot", scale="tai")
930 # DiaSourceId=102 withdrawn second time, it has no effect.
931 apdb.withdrawDiaForcedSources([sources_by_id[i] for i in (101, 102)], timeWithdrawn=withdraw_time2)
933 fsources = apdb.getDiaForcedSources(region, [100, 101, 102], self.visit_time)
934 assert fsources is not None
935 fsources.set_index("diaObjectId", inplace=True)
936 if self.use_mjd:
937 self.assertEqual(fsources.loc[100, "timeWithdrawnMjdTai"], withdraw_time1.mjd)
938 self.assertEqual(fsources.loc[101, "timeWithdrawnMjdTai"], withdraw_time2.mjd)
939 self.assertEqual(fsources.loc[102, "timeWithdrawnMjdTai"], withdraw_time1.mjd)
940 else:
941 # Exact type and values depend on backend, I don't want to
942 # overcomplicate it, just check that all of them are not NULL.
943 self.assertFalse(any(pandas.isnull(fsources["time_withdrawn"])))
945 # Check replication update tables.
946 replica_chunks = apdb_replica.getReplicaChunks()
947 if not self.enable_replica:
948 self.assertIsNone(replica_chunks)
949 else:
950 # Check that there are 3 update records in replica tables.
951 assert replica_chunks is not None
953 # There could be one or two chunks.
954 self.assertTrue(1 <= len(replica_chunks) <= 2)
956 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks])
957 self.assertEqual(len(update_records), 3)
959 def test_getChunks(self) -> None:
960 """Store and retrieve replica chunks."""
961 # don't care about sources.
962 config = self.make_instance()
963 apdb = Apdb.from_config(config)
964 apdb_replica = ApdbReplica.from_config(config)
965 visit_time = self.visit_time
967 region1 = self.make_region((1.0, 1.0, -1.0))
968 region2 = self.make_region((-1.0, -1.0, -1.0))
969 nobj = 100
970 objects1 = makeObjectCatalog(region1, nobj)
971 objects2 = makeObjectCatalog(region2, nobj, start_id=nobj * 2)
973 # With the default 10 minutes replica chunk window we should have 4
974 # records.
975 visits = [
976 (astropy.time.Time("2021-01-01T00:01:00", format="isot", scale="tai"), objects1),
977 (astropy.time.Time("2021-01-01T00:02:00", format="isot", scale="tai"), objects2),
978 (astropy.time.Time("2021-01-01T00:11:00", format="isot", scale="tai"), objects1),
979 (astropy.time.Time("2021-01-01T00:12:00", format="isot", scale="tai"), objects2),
980 (astropy.time.Time("2021-01-01T00:45:00", format="isot", scale="tai"), objects1),
981 (astropy.time.Time("2021-01-01T00:46:00", format="isot", scale="tai"), objects2),
982 (astropy.time.Time("2021-03-01T00:01:00", format="isot", scale="tai"), objects1),
983 (astropy.time.Time("2021-03-01T00:02:00", format="isot", scale="tai"), objects2),
984 ]
986 start_id = 0
987 for visit_time, objects in visits:
988 sources = makeSourceCatalog(objects, visit_time, start_id=start_id, use_mjd=self.use_mjd)
989 fsources = makeForcedSourceCatalog(objects, visit_time, visit=start_id, use_mjd=self.use_mjd)
990 apdb.store(visit_time, objects, sources, fsources)
991 start_id += nobj
993 replica_chunks = apdb_replica.getReplicaChunks()
994 if not self.enable_replica:
995 self.assertIsNone(replica_chunks)
997 with self.assertRaisesRegex(ValueError, "APDB is not configured for replication"):
998 apdb_replica.getTableDataChunks(ApdbTables.DiaObject, [])
1000 else:
1001 assert replica_chunks is not None
1002 self.assertEqual(len(replica_chunks), 4)
1004 with self.assertRaisesRegex(ValueError, "does not support replica chunks"):
1005 apdb_replica.getTableDataChunks(ApdbTables.SSObject, [])
1007 def _check_chunks(replica_chunks: list[ReplicaChunk], n_records: int | None = None) -> None:
1008 if n_records is None:
1009 n_records = len(replica_chunks) * nobj
1010 res = apdb_replica.getTableDataChunks(
1011 ApdbTables.DiaObject, (chunk.id for chunk in replica_chunks)
1012 )
1013 self.assert_table_data(res, n_records, ApdbTables.DiaObject)
1014 validityStartColumn = "validityStartMjdTai" if self.use_mjd else "validityStart"
1015 validityStartType = (
1016 felis.datamodel.DataType.double if self.use_mjd else felis.datamodel.DataType.timestamp
1017 )
1018 self.assert_column_types(
1019 res,
1020 {
1021 "apdb_replica_chunk": felis.datamodel.DataType.long,
1022 "diaObjectId": felis.datamodel.DataType.long,
1023 validityStartColumn: validityStartType,
1024 "ra": felis.datamodel.DataType.double,
1025 "dec": felis.datamodel.DataType.double,
1026 "parallax": felis.datamodel.DataType.float,
1027 "nDiaSources": felis.datamodel.DataType.int,
1028 },
1029 )
1031 res = apdb_replica.getTableDataChunks(
1032 ApdbTables.DiaSource, (chunk.id for chunk in replica_chunks)
1033 )
1034 self.assert_table_data(res, n_records, ApdbTables.DiaSource)
1035 self.assert_column_types(
1036 res,
1037 {
1038 "apdb_replica_chunk": felis.datamodel.DataType.long,
1039 "diaSourceId": felis.datamodel.DataType.long,
1040 "visit": felis.datamodel.DataType.long,
1041 "detector": felis.datamodel.DataType.short,
1042 },
1043 )
1045 res = apdb_replica.getTableDataChunks(
1046 ApdbTables.DiaForcedSource, (chunk.id for chunk in replica_chunks)
1047 )
1048 self.assert_table_data(res, n_records, ApdbTables.DiaForcedSource)
1049 self.assert_column_types(
1050 res,
1051 {
1052 "apdb_replica_chunk": felis.datamodel.DataType.long,
1053 "diaObjectId": felis.datamodel.DataType.long,
1054 "visit": felis.datamodel.DataType.long,
1055 "detector": felis.datamodel.DataType.short,
1056 },
1057 )
1059 # read it back and check sizes
1060 _check_chunks(replica_chunks, 800)
1061 _check_chunks(replica_chunks[1:], 600)
1062 _check_chunks(replica_chunks[1:-1], 400)
1063 _check_chunks(replica_chunks[2:3], 200)
1064 _check_chunks([])
1066 # try to remove some of those
1067 deleted_chunks = replica_chunks[:1]
1068 apdb_replica.deleteReplicaChunks(chunk.id for chunk in deleted_chunks)
1070 # All queries on deleted ids should return empty set.
1071 _check_chunks(deleted_chunks, 0)
1073 replica_chunks = apdb_replica.getReplicaChunks()
1074 assert replica_chunks is not None
1075 self.assertEqual(len(replica_chunks), 3)
1077 _check_chunks(replica_chunks, 600)
1079 def test_reassignObjects(self) -> None:
1080 """Reassign DiaObjects."""
1081 # don't care about sources.
1082 config = self.make_instance()
1083 apdb = Apdb.from_config(config)
1085 region = self.make_region()
1086 visit_time = self.visit_time
1087 objects = makeObjectCatalog(region, 100)
1088 oids = list(objects["diaObjectId"])
1089 sources = makeSourceCatalog(objects, visit_time, use_mjd=self.use_mjd)
1090 apdb.store(visit_time, objects, sources)
1092 # read it back and filter by ID
1093 res = apdb.getDiaSources(region, oids, visit_time)
1094 self.assert_catalog(res, len(sources), ApdbTables.DiaSource)
1096 apdb.reassignDiaSources({1: 1, 2: 2, 5: 5})
1097 res = apdb.getDiaSources(region, oids, visit_time)
1098 self.assert_catalog(res, len(sources) - 3, ApdbTables.DiaSource)
1100 with self.assertRaisesRegex(ValueError, r"do not exist.*\D1000"):
1101 apdb.reassignDiaSources(
1102 {
1103 1000: 1,
1104 7: 3,
1105 }
1106 )
1107 self.assert_catalog(res, len(sources) - 3, ApdbTables.DiaSource)
1109 def test_storeUpdateRecord(self) -> None:
1110 """Test _storeUpdateRecord() method."""
1111 config = self.make_instance()
1112 apdb = Apdb.from_config(config)
1114 # Times are totally arbitrary.
1115 update_time_ns1 = 2_000_000_000_000_000_000
1116 update_time_ns2 = 2_000_000_001_000_000_000
1117 records = [
1118 ApdbReassignDiaSourceToSSObjectRecord(
1119 update_time_ns=update_time_ns1,
1120 update_order=0,
1121 diaSourceId=1,
1122 ssObjectId=1,
1123 ssObjectReassocTimeMjdTai=60000.0,
1124 ra=45.0,
1125 dec=-45.0,
1126 midpointMjdTai=60000.0,
1127 ),
1128 ApdbWithdrawDiaSourceRecord(
1129 update_time_ns=update_time_ns1,
1130 update_order=1,
1131 diaSourceId=123456,
1132 timeWithdrawnMjdTai=61000.0,
1133 ra=45.0,
1134 dec=-45.0,
1135 midpointMjdTai=60000.0,
1136 ),
1137 ApdbReassignDiaSourceToSSObjectRecord(
1138 update_time_ns=update_time_ns1,
1139 update_order=3,
1140 diaSourceId=2,
1141 ssObjectId=3,
1142 ssObjectReassocTimeMjdTai=60000.0,
1143 ra=45.0,
1144 dec=-45.0,
1145 midpointMjdTai=60000.0,
1146 ),
1147 ApdbWithdrawDiaSourceRecord(
1148 update_time_ns=update_time_ns2,
1149 update_order=0,
1150 diaSourceId=123456,
1151 timeWithdrawnMjdTai=61000.0,
1152 ra=45.0,
1153 dec=-45.0,
1154 midpointMjdTai=60000.0,
1155 ),
1156 ]
1158 update_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
1159 chunk = ReplicaChunk.make_replica_chunk(update_time, 600)
1161 if not self.enable_replica:
1162 with self.assertRaises(TypeError):
1163 self.store_update_records(apdb, records, chunk)
1164 else:
1165 self.store_update_records(apdb, records, chunk)
1167 apdb_replica = ApdbReplica.from_config(config)
1168 records_returned = apdb_replica.getUpdateRecordChunks([chunk.id])
1170 # Input records are ordered, output will be ordered too.
1171 self.assertEqual(records_returned, records)
1173 @abstractmethod
1174 def store_update_records(self, apdb: Apdb, records: list[ApdbUpdateRecord], chunk: ReplicaChunk) -> None:
1175 """Store update records in database, must be overriden in subclass."""
1176 raise NotImplementedError()
1178 def test_midpointMjdTai_src(self) -> None:
1179 """Test for time filtering of DiaSources."""
1180 config = self.make_instance()
1181 apdb = Apdb.from_config(config)
1183 region = self.make_region()
1184 # 2021-01-01 plus 360 days is 2021-12-27
1185 src_time1 = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
1186 src_time2 = astropy.time.Time("2021-01-01T00:00:02", format="isot", scale="tai")
1187 visit_time0 = astropy.time.Time("2021-12-26T23:59:59", format="isot", scale="tai")
1188 visit_time1 = astropy.time.Time("2021-12-27T00:00:01", format="isot", scale="tai")
1189 visit_time2 = astropy.time.Time("2021-12-27T00:00:03", format="isot", scale="tai")
1190 one_sec = astropy.time.TimeDelta(1.0, format="sec")
1192 objects = makeObjectCatalog(region, 100)
1193 oids = list(objects["diaObjectId"])
1194 sources = makeSourceCatalog(objects, src_time1, 0, use_mjd=self.use_mjd)
1195 apdb.store(src_time1, objects, sources)
1197 sources = makeSourceCatalog(objects, src_time2, 100, use_mjd=self.use_mjd)
1198 apdb.store(src_time2, objects, sources)
1200 # reading at time of last save should read all
1201 res = apdb.getDiaSources(region, oids, src_time2)
1202 self.assert_catalog(res, 200, ApdbTables.DiaSource)
1204 # one second before 12 months
1205 res = apdb.getDiaSources(region, oids, visit_time0)
1206 self.assert_catalog(res, 200, ApdbTables.DiaSource)
1208 # reading at later time of last save should only read a subset
1209 res = apdb.getDiaSources(region, oids, visit_time1)
1210 self.assert_catalog(res, 100, ApdbTables.DiaSource)
1212 # reading at later time of last save should only read a subset
1213 res = apdb.getDiaSources(region, oids, visit_time2)
1214 self.assert_catalog(res, 0, ApdbTables.DiaSource)
1216 # Use explicit start time argument instead of 12 month window, visit
1217 # time does not matter in this case, set it to before all data.
1218 res = apdb.getDiaSources(region, oids, src_time1 - one_sec, src_time1 - one_sec)
1219 self.assert_catalog(res, 200, ApdbTables.DiaSource)
1221 res = apdb.getDiaSources(region, oids, src_time1 - one_sec, src_time2 - one_sec)
1222 self.assert_catalog(res, 100, ApdbTables.DiaSource)
1224 res = apdb.getDiaSources(region, oids, src_time1 - one_sec, src_time2 + one_sec)
1225 self.assert_catalog(res, 0, ApdbTables.DiaSource)
1227 def test_midpointMjdTai_fsrc(self) -> None:
1228 """Test for time filtering of DiaForcedSources."""
1229 config = self.make_instance()
1230 apdb = Apdb.from_config(config)
1232 region = self.make_region()
1233 src_time1 = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
1234 src_time2 = astropy.time.Time("2021-01-01T00:00:02", format="isot", scale="tai")
1235 visit_time0 = astropy.time.Time("2021-12-26T23:59:59", format="isot", scale="tai")
1236 visit_time1 = astropy.time.Time("2021-12-27T00:00:01", format="isot", scale="tai")
1237 visit_time2 = astropy.time.Time("2021-12-27T00:00:03", format="isot", scale="tai")
1238 one_sec = astropy.time.TimeDelta(1.0, format="sec")
1240 objects = makeObjectCatalog(region, 100)
1241 oids = list(objects["diaObjectId"])
1242 sources = makeForcedSourceCatalog(objects, src_time1, 1, use_mjd=self.use_mjd)
1243 apdb.store(src_time1, objects, forced_sources=sources)
1245 sources = makeForcedSourceCatalog(objects, src_time2, 2, use_mjd=self.use_mjd)
1246 apdb.store(src_time2, objects, forced_sources=sources)
1248 # reading at time of last save should read all
1249 res = apdb.getDiaForcedSources(region, oids, src_time2)
1250 self.assert_catalog(res, 200, ApdbTables.DiaForcedSource)
1252 # one second before 12 months
1253 res = apdb.getDiaForcedSources(region, oids, visit_time0)
1254 self.assert_catalog(res, 200, ApdbTables.DiaForcedSource)
1256 # reading at later time of last save should only read a subset
1257 res = apdb.getDiaForcedSources(region, oids, visit_time1)
1258 self.assert_catalog(res, 100, ApdbTables.DiaForcedSource)
1260 # reading at later time of last save should only read a subset
1261 res = apdb.getDiaForcedSources(region, oids, visit_time2)
1262 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
1264 # Use explicit start time argument instead of 12 month window, visit
1265 # time does not matter in this case, set it to before all data.
1266 res = apdb.getDiaForcedSources(region, oids, src_time1 - one_sec, src_time1 - one_sec)
1267 self.assert_catalog(res, 200, ApdbTables.DiaForcedSource)
1269 res = apdb.getDiaForcedSources(region, oids, src_time1 - one_sec, src_time2 - one_sec)
1270 self.assert_catalog(res, 100, ApdbTables.DiaForcedSource)
1272 res = apdb.getDiaForcedSources(region, oids, src_time1 - one_sec, src_time2 + one_sec)
1273 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource)
1275 def test_metadata(self) -> None:
1276 """Simple test for writing/reading metadata table"""
1277 config = self.make_instance()
1278 apdb = Apdb.from_config(config)
1279 metadata = apdb.metadata
1281 # APDB should write two or three metadata items with version numbers
1282 # and a frozen JSON config.
1283 self.assertFalse(metadata.empty())
1284 self.assertEqual(len(list(metadata.items())), self.meta_row_count)
1286 metadata.set("meta", "data")
1287 metadata.set("data", "meta")
1289 self.assertFalse(metadata.empty())
1290 self.assertTrue(set(metadata.items()) >= {("meta", "data"), ("data", "meta")})
1292 with self.assertRaisesRegex(KeyError, "Metadata key 'meta' already exists"):
1293 metadata.set("meta", "data1")
1295 metadata.set("meta", "data2", force=True)
1296 self.assertTrue(set(metadata.items()) >= {("meta", "data2"), ("data", "meta")})
1298 self.assertTrue(metadata.delete("meta"))
1299 self.assertIsNone(metadata.get("meta"))
1300 self.assertFalse(metadata.delete("meta"))
1302 self.assertEqual(metadata.get("data"), "meta")
1303 self.assertEqual(metadata.get("meta", "meta"), "meta")
1305 def test_schemaVersionFromYaml(self) -> None:
1306 """Check version number handling for reading schema from YAML."""
1307 config = self.make_instance()
1308 default_schema = config.schema_file
1309 apdb = Apdb.from_config(config)
1310 self.assertEqual(apdb.schema.schemaVersion(), VersionTuple(0, 1, 1))
1312 with update_schema_yaml(default_schema, version="") as schema_file:
1313 config = self.make_instance(schema_file=schema_file)
1314 apdb = Apdb.from_config(config)
1315 self.assertEqual(
1316 apdb.schema.schemaVersion(),
1317 VersionTuple(0, 1, 0),
1318 )
1320 with update_schema_yaml(default_schema, version="99.0.0") as schema_file:
1321 config = self.make_instance(schema_file=schema_file)
1322 apdb = Apdb.from_config(config)
1323 self.assertEqual(
1324 apdb.schema.schemaVersion(),
1325 VersionTuple(99, 0, 0),
1326 )
1328 def test_config_freeze(self) -> None:
1329 """Test that some config fields are correctly frozen in database."""
1330 config = self.make_instance()
1332 # `enable_replica` is the only parameter that is frozen in all
1333 # implementations.
1334 config.enable_replica = not self.enable_replica
1335 apdb = Apdb.from_config(config)
1336 frozen_config = apdb.getConfig()
1337 self.assertEqual(frozen_config.enable_replica, self.enable_replica)
1340class ApdbSchemaUpdateTest(TestCaseMixin, ABC):
1341 """Base class for unit tests that verify how schema changes work."""
1343 visit_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai")
1345 @abstractmethod
1346 def make_instance(self, **kwargs: Any) -> ApdbConfig:
1347 """Make config class instance used in all tests.
1349 This method should return configuration that point to the identical
1350 database instance on each call (i.e. ``db_url`` must be the same,
1351 which also means for sqlite it has to use on-disk storage).
1352 """
1353 raise NotImplementedError()
1355 def make_region(self, xyz: tuple[float, float, float] | LonLat = (1.0, 1.0, -1.0)) -> Region:
1356 """Make a region to use in tests"""
1357 return _make_region(xyz)
1359 def test_schema_add_replica(self) -> None:
1360 """Check that new code can work with old schema without replica
1361 tables.
1362 """
1363 # Make schema without replica tables.
1364 config = self.make_instance(enable_replica=False)
1365 apdb = Apdb.from_config(config)
1366 apdb_replica = ApdbReplica.from_config(config)
1368 # Make APDB instance configured for replication.
1369 config.enable_replica = True
1370 apdb = Apdb.from_config(config)
1372 # Try to insert something, should work OK.
1373 region = self.make_region()
1374 visit_time = self.visit_time
1376 # have to store Objects first
1377 objects = makeObjectCatalog(region, 100)
1378 sources = makeSourceCatalog(objects, visit_time)
1379 fsources = makeForcedSourceCatalog(objects, visit_time)
1380 apdb.store(visit_time, objects, sources, fsources)
1382 # There should be no replica chunks.
1383 replica_chunks = apdb_replica.getReplicaChunks()
1384 self.assertIsNone(replica_chunks)
1386 def test_schemaVersionCheck(self) -> None:
1387 """Check version number compatibility."""
1388 config = self.make_instance()
1389 apdb = Apdb.from_config(config)
1391 self.assertEqual(apdb.schema.schemaVersion(), VersionTuple(0, 1, 1))
1393 # Claim that schema version is now 99.0.0, must raise an exception.
1394 with update_schema_yaml(config.schema_file, version="99.0.0") as schema_file:
1395 config.schema_file = schema_file
1396 with self.assertRaises(IncompatibleVersionError):
1397 apdb = Apdb.from_config(config)
1398 # Version is checked only when we try to do connect.
1399 apdb.metadata.items()