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:24 +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/>. 

21 

22from __future__ import annotations 

23 

24__all__ = ["ApdbSchemaUpdateTest", "ApdbTest", "update_schema_yaml"] 

25 

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 

35 

36import astropy.time 

37import felis.datamodel 

38import numpy 

39import pandas 

40import yaml 

41 

42from lsst.sphgeom import Angle, Circle, LonLat, Region, UnitVector3d 

43 

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 

68 

69if TYPE_CHECKING: 

70 from ..pixelization import Pixelization 

71 

72 

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) 

76 

77 

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 

87 

88 

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. 

96 

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. 

107 

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 

124 

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 

130 

131 

132class ApdbTest(TestCaseMixin, ABC): 

133 """Base class for Apdb tests that can be specialized for concrete 

134 implementation. 

135 

136 This can only be used as a mixin class for a unittest.TestCase and it 

137 calls various assert methods. 

138 """ 

139 

140 visit_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai") 

141 

142 processing_time = astropy.time.Time("2021-01-01T12:00:00", format="isot", scale="tai") 

143 

144 fsrc_requires_id_list = False 

145 """Should be set to True if getDiaForcedSources requires object IDs""" 

146 

147 enable_replica: bool = False 

148 """Set to true when support for replication is configured""" 

149 

150 use_mjd: bool = True 

151 """If True then timestamp columns are MJD TAI.""" 

152 

153 extra_chunk_columns = 1 

154 """Number of additional columns in chunk tables.""" 

155 

156 meta_row_count = 3 

157 """Initial row count in metadata table.""" 

158 

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 } 

167 

168 @abstractmethod 

169 def make_instance(self, **kwargs: Any) -> ApdbConfig: 

170 """Make database instance and return configuration for it.""" 

171 raise NotImplementedError() 

172 

173 @abstractmethod 

174 def getDiaObjects_table(self) -> ApdbTables: 

175 """Return type of table returned from getDiaObjects method.""" 

176 raise NotImplementedError() 

177 

178 @abstractmethod 

179 def pixelization(self, config: ApdbConfig) -> Pixelization: 

180 """Return pixelization used by implementation.""" 

181 raise NotImplementedError() 

182 

183 def assert_catalog(self, catalog: Any, rows: int, table: ApdbTables) -> None: 

184 """Validate catalog type and size 

185 

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

198 

199 def assert_table_data(self, catalog: Any, rows: int, table: ApdbTables) -> None: 

200 """Validate catalog type and size 

201 

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 ) 

220 

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) 

225 

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) 

229 

230 def test_makeSchema(self) -> None: 

231 """Test for making APDB schema.""" 

232 config = self.make_instance() 

233 apdb = Apdb.from_config(config) 

234 

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

243 

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) 

248 

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

257 

258 def test_empty_gets(self) -> None: 

259 """Test for getting data from empty database. 

260 

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) 

267 

268 region = self.make_region() 

269 visit_time = self.visit_time 

270 

271 res: pandas.DataFrame | None 

272 

273 # get objects by region 

274 res = apdb.getDiaObjects(region) 

275 self.assert_catalog(res, 0, self.getDiaObjects_table()) 

276 

277 # get sources by region 

278 res = apdb.getDiaSources(region, None, visit_time) 

279 self.assert_catalog(res, 0, ApdbTables.DiaSource) 

280 

281 res = apdb.getDiaSources(region, [], visit_time) 

282 self.assert_catalog(res, 0, ApdbTables.DiaSource) 

283 

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) 

287 

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) 

291 

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) 

295 

296 # data_factory's ccdVisitId generation corresponds to (1, 1) 

297 res = apdb.containsVisitDetector(visit=1, detector=1) 

298 self.assertFalse(res) 

299 

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) 

307 

308 def test_empty_gets_0months(self) -> None: 

309 """Test for getting data from empty database. 

310 

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) 

316 

317 region = self.make_region() 

318 visit_time = self.visit_time 

319 

320 res: pandas.DataFrame | None 

321 

322 # get objects by region 

323 res = apdb.getDiaObjects(region) 

324 self.assert_catalog(res, 0, self.getDiaObjects_table()) 

325 

326 # get sources by region 

327 res = apdb.getDiaSources(region, None, visit_time) 

328 self.assertIs(res, None) 

329 

330 # get sources by object ID, empty object list 

331 res = apdb.getDiaSources(region, [], visit_time) 

332 self.assertIs(res, None) 

333 

334 # get forced sources by object ID, empty object list 

335 res = apdb.getDiaForcedSources(region, [], visit_time) 

336 self.assertIs(res, None) 

337 

338 # Database is empty, no images exist. 

339 res = apdb.containsVisitDetector(visit=1, detector=1) 

340 self.assertFalse(res) 

341 

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) 

347 

348 region = self.make_region() 

349 visit_time = self.visit_time 

350 

351 # make catalog with Objects 

352 catalog = makeObjectCatalog(region, 100) 

353 

354 # store catalog 

355 apdb.store(visit_time, catalog) 

356 

357 # read it back and check sizes 

358 res = apdb.getDiaObjects(region) 

359 self.assert_catalog(res, len(catalog), self.getDiaObjects_table()) 

360 

361 # TODO: test apdb.contains with generic implementation from DM-41671 

362 

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) 

371 

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

375 

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) 

382 

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) 

388 

389 # Check that they fall into different pixels. 

390 self.assertNotEqual(pixelization.pixel(uv1), pixelization.pixel(uv2)) 

391 

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) 

396 

397 visit_time2 = visit_time1 + astropy.time.TimeDelta(120.0, format="sec") 

398 catalog1 = makeObjectCatalog(lonlat2, 1) 

399 apdb.store(visit_time2, catalog1) 

400 

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

405 

406 # Read it back, must return the latest one. 

407 res = apdb.getDiaObjects(region) 

408 self.assert_catalog(res, 1, self.getDiaObjects_table()) 

409 

410 def test_storeSources(self) -> None: 

411 """Store and retrieve DiaSources.""" 

412 config = self.make_instance() 

413 apdb = Apdb.from_config(config) 

414 

415 region = self.make_region() 

416 visit_time = self.visit_time 

417 

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) 

422 

423 # save the objects and sources 

424 apdb.store(visit_time, objects, sources) 

425 

426 # read it back, no ID filtering 

427 res = apdb.getDiaSources(region, None, visit_time) 

428 self.assert_catalog(res, len(sources), ApdbTables.DiaSource) 

429 

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) 

433 

434 # read it back to get schema 

435 res = apdb.getDiaSources(region, [], visit_time) 

436 self.assert_catalog(res, 0, ApdbTables.DiaSource) 

437 

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) 

445 

446 def test_storeForcedSources(self) -> None: 

447 """Store and retrieve DiaForcedSources.""" 

448 config = self.make_instance() 

449 apdb = Apdb.from_config(config) 

450 

451 region = self.make_region() 

452 visit_time = self.visit_time 

453 

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) 

458 

459 apdb.store(visit_time, objects, forced_sources=catalog) 

460 

461 # read it back and check sizes 

462 res = apdb.getDiaForcedSources(region, oids, visit_time) 

463 self.assert_catalog(res, len(catalog), ApdbTables.DiaForcedSource) 

464 

465 # read it back to get schema 

466 res = apdb.getDiaForcedSources(region, [], visit_time) 

467 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource) 

468 

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) 

475 

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) 

480 

481 region = self.make_region() 

482 visit_time = self.visit_time 

483 

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 

489 

490 # save the objects and sources 

491 apdb.store(visit_time, objects, sources) 

492 

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

498 

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) 

503 

504 region = self.make_region() 

505 visit_time = self.visit_time 

506 

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) 

516 

517 apdb.store(visit_time, objects, forced_sources=catalog) 

518 

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) 

523 

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

533 

534 def test_getDiaObjectsForDedup(self) -> None: 

535 """Test getDiaObjectsForDedup() method.""" 

536 config = self.make_instance() 

537 apdb = Apdb.from_config(config) 

538 

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) 

546 

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 ] 

552 

553 for visit_time, objects in visits: 

554 apdb.store(visit_time, objects) 

555 

556 catalog = apdb.getDiaObjectsForDedup() 

557 self.assertEqual(len(catalog), 300) 

558 

559 catalog = apdb.getDiaObjectsForDedup(visits[0][0]) 

560 self.assertEqual(len(catalog), 300) 

561 

562 catalog = apdb.getDiaObjectsForDedup(visits[1][0]) 

563 self.assertEqual(len(catalog), 200) 

564 

565 catalog = apdb.getDiaObjectsForDedup(visits[2][0]) 

566 self.assertEqual(len(catalog), 100) 

567 

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) 

571 

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] 

578 

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) 

586 

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 ] 

592 

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 

598 

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 ] 

605 

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

610 

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

615 

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) 

622 

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) 

629 

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) 

634 

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) 

638 

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

647 

648 dia_source_ids = [DiaSourceId.from_named_tuple(row) for row in sources.itertuples()] 

649 

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) 

657 

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

664 

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

669 

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 ) 

679 

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

686 

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 

692 

693 # There could be one or two chunks. 

694 self.assertTrue(1 <= len(replica_chunks) <= 2) 

695 

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) 

701 

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) 

709 

710 region = self.make_region() 

711 visit_time = self.visit_time 

712 

713 # make catalog with Objects 

714 catalog = makeObjectCatalog(region, 100) 

715 

716 # store catalog 

717 apdb.store(visit_time, catalog) 

718 

719 # read it back and check sizes 

720 res = apdb.getDiaObjects(region) 

721 self.assert_catalog(res, 100, self.getDiaObjects_table()) 

722 

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) 

727 

728 res = apdb.getDiaObjects(region) 

729 self.assert_catalog(res, 90, self.getDiaObjects_table()) 

730 

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 

737 

738 # There could be one or two chunks. 

739 self.assertTrue(1 <= len(replica_chunks) <= 2) 

740 

741 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks]) 

742 self.assertEqual(len(update_records), 10) 

743 

744 # Check that empty list works. 

745 count = apdb.setValidityEnd(object_ids, self.processing_time) 

746 self.assertEqual(count, 0) 

747 

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) 

753 

754 count = apdb.setValidityEnd(object_ids, self.processing_time) 

755 self.assertEqual(count, 2) 

756 

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) 

762 

763 region = self.make_region() 

764 

765 # make catalog with Objects 

766 objects = makeObjectCatalog(region, 100) 

767 

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

772 

773 # store catalog 

774 apdb.store(visit_time1, objects) 

775 

776 catalog = apdb.getDiaObjectsForDedup() 

777 self.assertEqual(len(catalog), 100) 

778 

779 catalog = apdb.getDiaObjectsForDedup(visit_time1) 

780 self.assertEqual(len(catalog), 100) 

781 

782 apdb.resetDedup(dedup_time1) 

783 

784 catalog = apdb.getDiaObjectsForDedup(visit_time1) 

785 self.assertEqual(len(catalog), self._count_after_reset_dedup(100)) 

786 

787 apdb.store(visit_time2, objects) 

788 

789 catalog = apdb.getDiaObjectsForDedup() 

790 self.assertEqual(len(catalog), 100) 

791 

792 catalog = apdb.getDiaObjectsForDedup(dedup_time1) 

793 self.assertEqual(len(catalog), 100) 

794 

795 apdb.resetDedup(dedup_time2) 

796 

797 catalog = apdb.getDiaObjectsForDedup(dedup_time1) 

798 self.assertEqual(len(catalog), self._count_after_reset_dedup(100)) 

799 

800 catalog = apdb.getDiaObjectsForDedup() 

801 self.assertEqual(len(catalog), 0) 

802 

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

810 

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) 

817 

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) 

822 

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) 

830 

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) 

835 

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

841 

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} 

850 

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

857 

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) 

864 

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 ) 

886 

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 

894 

895 # There could be one or two chunks. 

896 self.assertTrue(1 <= len(replica_chunks) <= 2) 

897 

898 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks]) 

899 self.assertEqual(len(update_records), 4) 

900 

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) 

907 

908 lonlat = LonLat.fromDegrees(0.0, 0.0) 

909 region = self.make_region(lonlat) 

910 

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) 

915 

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} 

920 

921 if self.use_mjd: 

922 self.assertTrue(all(pandas.isna(fsources["timeWithdrawnMjdTai"]))) 

923 else: 

924 self.assertTrue(all(pandas.isnull(fsources["time_withdrawn"]))) 

925 

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) 

932 

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

944 

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 

952 

953 # There could be one or two chunks. 

954 self.assertTrue(1 <= len(replica_chunks) <= 2) 

955 

956 update_records = apdb_replica.getUpdateRecordChunks([chunk.id for chunk in replica_chunks]) 

957 self.assertEqual(len(update_records), 3) 

958 

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 

966 

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) 

972 

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 ] 

985 

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 

992 

993 replica_chunks = apdb_replica.getReplicaChunks() 

994 if not self.enable_replica: 

995 self.assertIsNone(replica_chunks) 

996 

997 with self.assertRaisesRegex(ValueError, "APDB is not configured for replication"): 

998 apdb_replica.getTableDataChunks(ApdbTables.DiaObject, []) 

999 

1000 else: 

1001 assert replica_chunks is not None 

1002 self.assertEqual(len(replica_chunks), 4) 

1003 

1004 with self.assertRaisesRegex(ValueError, "does not support replica chunks"): 

1005 apdb_replica.getTableDataChunks(ApdbTables.SSObject, []) 

1006 

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 ) 

1030 

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 ) 

1044 

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 ) 

1058 

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

1065 

1066 # try to remove some of those 

1067 deleted_chunks = replica_chunks[:1] 

1068 apdb_replica.deleteReplicaChunks(chunk.id for chunk in deleted_chunks) 

1069 

1070 # All queries on deleted ids should return empty set. 

1071 _check_chunks(deleted_chunks, 0) 

1072 

1073 replica_chunks = apdb_replica.getReplicaChunks() 

1074 assert replica_chunks is not None 

1075 self.assertEqual(len(replica_chunks), 3) 

1076 

1077 _check_chunks(replica_chunks, 600) 

1078 

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) 

1084 

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) 

1091 

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) 

1095 

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) 

1099 

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) 

1108 

1109 def test_storeUpdateRecord(self) -> None: 

1110 """Test _storeUpdateRecord() method.""" 

1111 config = self.make_instance() 

1112 apdb = Apdb.from_config(config) 

1113 

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 ] 

1157 

1158 update_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai") 

1159 chunk = ReplicaChunk.make_replica_chunk(update_time, 600) 

1160 

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) 

1166 

1167 apdb_replica = ApdbReplica.from_config(config) 

1168 records_returned = apdb_replica.getUpdateRecordChunks([chunk.id]) 

1169 

1170 # Input records are ordered, output will be ordered too. 

1171 self.assertEqual(records_returned, records) 

1172 

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

1177 

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) 

1182 

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

1191 

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) 

1196 

1197 sources = makeSourceCatalog(objects, src_time2, 100, use_mjd=self.use_mjd) 

1198 apdb.store(src_time2, objects, sources) 

1199 

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) 

1203 

1204 # one second before 12 months 

1205 res = apdb.getDiaSources(region, oids, visit_time0) 

1206 self.assert_catalog(res, 200, ApdbTables.DiaSource) 

1207 

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) 

1211 

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) 

1215 

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) 

1220 

1221 res = apdb.getDiaSources(region, oids, src_time1 - one_sec, src_time2 - one_sec) 

1222 self.assert_catalog(res, 100, ApdbTables.DiaSource) 

1223 

1224 res = apdb.getDiaSources(region, oids, src_time1 - one_sec, src_time2 + one_sec) 

1225 self.assert_catalog(res, 0, ApdbTables.DiaSource) 

1226 

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) 

1231 

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

1239 

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) 

1244 

1245 sources = makeForcedSourceCatalog(objects, src_time2, 2, use_mjd=self.use_mjd) 

1246 apdb.store(src_time2, objects, forced_sources=sources) 

1247 

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) 

1251 

1252 # one second before 12 months 

1253 res = apdb.getDiaForcedSources(region, oids, visit_time0) 

1254 self.assert_catalog(res, 200, ApdbTables.DiaForcedSource) 

1255 

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) 

1259 

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) 

1263 

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) 

1268 

1269 res = apdb.getDiaForcedSources(region, oids, src_time1 - one_sec, src_time2 - one_sec) 

1270 self.assert_catalog(res, 100, ApdbTables.DiaForcedSource) 

1271 

1272 res = apdb.getDiaForcedSources(region, oids, src_time1 - one_sec, src_time2 + one_sec) 

1273 self.assert_catalog(res, 0, ApdbTables.DiaForcedSource) 

1274 

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 

1280 

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) 

1285 

1286 metadata.set("meta", "data") 

1287 metadata.set("data", "meta") 

1288 

1289 self.assertFalse(metadata.empty()) 

1290 self.assertTrue(set(metadata.items()) >= {("meta", "data"), ("data", "meta")}) 

1291 

1292 with self.assertRaisesRegex(KeyError, "Metadata key 'meta' already exists"): 

1293 metadata.set("meta", "data1") 

1294 

1295 metadata.set("meta", "data2", force=True) 

1296 self.assertTrue(set(metadata.items()) >= {("meta", "data2"), ("data", "meta")}) 

1297 

1298 self.assertTrue(metadata.delete("meta")) 

1299 self.assertIsNone(metadata.get("meta")) 

1300 self.assertFalse(metadata.delete("meta")) 

1301 

1302 self.assertEqual(metadata.get("data"), "meta") 

1303 self.assertEqual(metadata.get("meta", "meta"), "meta") 

1304 

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

1311 

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 ) 

1319 

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 ) 

1327 

1328 def test_config_freeze(self) -> None: 

1329 """Test that some config fields are correctly frozen in database.""" 

1330 config = self.make_instance() 

1331 

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) 

1338 

1339 

1340class ApdbSchemaUpdateTest(TestCaseMixin, ABC): 

1341 """Base class for unit tests that verify how schema changes work.""" 

1342 

1343 visit_time = astropy.time.Time("2021-01-01T00:00:00", format="isot", scale="tai") 

1344 

1345 @abstractmethod 

1346 def make_instance(self, **kwargs: Any) -> ApdbConfig: 

1347 """Make config class instance used in all tests. 

1348 

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

1354 

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) 

1358 

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) 

1367 

1368 # Make APDB instance configured for replication. 

1369 config.enable_replica = True 

1370 apdb = Apdb.from_config(config) 

1371 

1372 # Try to insert something, should work OK. 

1373 region = self.make_region() 

1374 visit_time = self.visit_time 

1375 

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) 

1381 

1382 # There should be no replica chunks. 

1383 replica_chunks = apdb_replica.getReplicaChunks() 

1384 self.assertIsNone(replica_chunks) 

1385 

1386 def test_schemaVersionCheck(self) -> None: 

1387 """Check version number compatibility.""" 

1388 config = self.make_instance() 

1389 apdb = Apdb.from_config(config) 

1390 

1391 self.assertEqual(apdb.schema.schemaVersion(), VersionTuple(0, 1, 1)) 

1392 

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