Coverage for python/lsst/dax/apdb/sql/apdbSql.py: 88%

756 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 

22"""Module defining Apdb class and related methods.""" 

23 

24from __future__ import annotations 

25 

26__all__ = ["ApdbSql"] 

27 

28import datetime 

29import json 

30import logging 

31import urllib.parse 

32import uuid 

33import warnings 

34from collections import Counter 

35from collections.abc import Iterable, Mapping, MutableMapping 

36from contextlib import closing 

37from typing import TYPE_CHECKING, Any 

38 

39import astropy.time 

40import numpy as np 

41import pandas 

42import sqlalchemy 

43import sqlalchemy.dialects.postgresql 

44import sqlalchemy.dialects.sqlite 

45from sqlalchemy import func, sql 

46from sqlalchemy.pool import NullPool 

47 

48from lsst.sphgeom import HtmPixelization, LonLat, Region, UnitVector3d 

49from lsst.utils.db_auth import DbAuth, DbAuthNotFoundError 

50from lsst.utils.iteration import chunk_iterable 

51 

52from ..apdb import Apdb 

53from ..apdbConfigFreezer import ApdbConfigFreezer 

54from ..apdbReplica import ReplicaChunk 

55from ..apdbSchema import ApdbSchema, ApdbTables 

56from ..apdbUpdateRecord import ( 

57 ApdbCloseDiaObjectValidityRecord, 

58 ApdbReassignDiaSourceToDiaObjectRecord, 

59 ApdbUpdateNDiaSourcesRecord, 

60 ApdbUpdateRecord, 

61 ApdbWithdrawDiaForcedSourceRecord, 

62 ApdbWithdrawDiaSourceRecord, 

63) 

64from ..config import ApdbConfig 

65from ..monitor import MonAgent 

66from ..recordIds import DiaForcedSourceId, DiaObjectId, DiaSourceId 

67from ..schema_model import Table 

68from ..timer import Timer 

69from ..versionTuple import IncompatibleVersionError, VersionTuple 

70from .apdbMetadataSql import ApdbMetadataSql 

71from .apdbSqlAdmin import ApdbSqlAdmin 

72from .apdbSqlReplica import ApdbSqlReplica, ApdbSqlTableData 

73from .apdbSqlSchema import ApdbSqlSchema, ExtraTables 

74from .config import ApdbSqlConfig 

75 

76if TYPE_CHECKING: 

77 import sqlite3 

78 

79 from ..apdbMetadata import ApdbMetadata 

80 from ..apdbUpdateRecord import ApdbUpdateRecord 

81 

82_LOG = logging.getLogger(__name__) 

83 

84_MON = MonAgent(__name__) 

85 

86VERSION = VersionTuple(1, 2, 1) 

87"""Version for the code controlling non-replication tables. This needs to be 

88updated following compatibility rules when schema produced by this code 

89changes. 

90""" 

91 

92 

93def _coerce_uint64(df: pandas.DataFrame) -> pandas.DataFrame: 

94 """Change the type of uint64 columns to int64, and return copy of data 

95 frame. 

96 """ 

97 names = [c[0] for c in df.dtypes.items() if c[1] == np.uint64] 

98 return df.astype(dict.fromkeys(names, np.int64)) 

99 

100 

101def _make_midpointMjdTai_start(visit_time: astropy.time.Time, months: int) -> float: 

102 """Calculate starting point for time-based source search. 

103 

104 Parameters 

105 ---------- 

106 visit_time : `astropy.time.Time` 

107 Time of current visit. 

108 months : `int` 

109 Number of months in the sources history. 

110 

111 Returns 

112 ------- 

113 time : `float` 

114 A ``midpointMjdTai`` starting point, MJD time. 

115 """ 

116 # TODO: Use of MJD must be consistent with the code in ap_association 

117 # (see DM-31996) 

118 return float(visit_time.tai.mjd) - months * 30 

119 

120 

121def _onSqlite3Connect( 

122 dbapiConnection: sqlite3.Connection, connectionRecord: sqlalchemy.pool._ConnectionRecord 

123) -> None: 

124 # Enable foreign keys 

125 with closing(dbapiConnection.cursor()) as cursor: 

126 cursor.execute("PRAGMA foreign_keys=ON;") 

127 

128 

129class ApdbSql(Apdb): 

130 """Implementation of APDB interface based on SQL database. 

131 

132 The implementation is configured via standard ``pex_config`` mechanism 

133 using `ApdbSqlConfig` configuration class. For an example of different 

134 configurations check ``config/`` folder. 

135 

136 Parameters 

137 ---------- 

138 config : `ApdbSqlConfig` 

139 Configuration object. 

140 """ 

141 

142 metadataSchemaVersionKey = "version:schema" 

143 """Name of the metadata key to store schema version number.""" 

144 

145 metadataCodeVersionKey = "version:ApdbSql" 

146 """Name of the metadata key to store code version number.""" 

147 

148 metadataReplicaVersionKey = "version:ApdbSqlReplica" 

149 """Name of the metadata key to store replica code version number.""" 

150 

151 metadataConfigKey = "config:apdb-sql.json" 

152 """Name of the metadata key to store code version number.""" 

153 

154 metadataDedupKey = "status:deduplication.json" 

155 """Name of the metadata key to store code version number.""" 

156 

157 _frozen_parameters = ( 

158 "enable_replica", 

159 "dia_object_index", 

160 "pixelization.htm_level", 

161 "pixelization.htm_index_column", 

162 "ra_dec_columns", 

163 ) 

164 """Names of the config parameters to be frozen in metadata table.""" 

165 

166 def __init__(self, config: ApdbSqlConfig): 

167 self._engine = self._makeEngine(config, create=False) 

168 

169 sa_metadata = sqlalchemy.MetaData(schema=config.namespace) 

170 meta_table_name = ApdbTables.metadata.table_name(prefix=config.prefix) 

171 meta_table = sqlalchemy.schema.Table(meta_table_name, sa_metadata, autoload_with=self._engine) 

172 self._metadata = ApdbMetadataSql(self._engine, meta_table) 

173 

174 # Get tables schemas. 

175 self._table_schema = ApdbSchema(config.schema_file, config.ss_schema_file) 

176 

177 # Check that versions are compatible, must be the first thing before 

178 # reading frozen config. 

179 self._db_schema_version = self._versionCheck(self._metadata, self._table_schema.schemaVersion()) 

180 

181 # Read frozen config from metadata. 

182 config_json = self._metadata.get(self.metadataConfigKey) 

183 if config_json is not None: 183 ↛ 188line 183 didn't jump to line 188 because the condition on line 183 was always true

184 # Update config from metadata. 

185 freezer = ApdbConfigFreezer[ApdbSqlConfig](self._frozen_parameters) 

186 self.config = freezer.update(config, config_json) 

187 else: 

188 self.config = config 

189 

190 self._schema = ApdbSqlSchema( 

191 table_schema=self._table_schema, 

192 engine=self._engine, 

193 dia_object_index=self.config.dia_object_index, 

194 prefix=self.config.prefix, 

195 namespace=self.config.namespace, 

196 htm_index_column=self.config.pixelization.htm_index_column, 

197 enable_replica=self.config.enable_replica, 

198 ) 

199 

200 self.pixelator = HtmPixelization(self.config.pixelization.htm_level) 

201 

202 if _LOG.isEnabledFor(logging.DEBUG): 202 ↛ exitline 202 didn't return from function '__init__' because the condition on line 202 was always true

203 _LOG.debug("ApdbSql Configuration: %s", self.config.model_dump()) 

204 

205 def _timer(self, name: str, *, tags: Mapping[str, str | int] | None = None) -> Timer: 

206 """Create `Timer` instance given its name.""" 

207 return Timer(name, _MON, tags=tags) 

208 

209 @classmethod 

210 def _makeEngine(cls, config: ApdbSqlConfig, *, create: bool) -> sqlalchemy.engine.Engine: 

211 """Make SQLALchemy engine based on configured parameters. 

212 

213 Parameters 

214 ---------- 

215 config : `ApdbSqlConfig` 

216 Configuration object. 

217 create : `bool` 

218 Whether to try to create new database file, only relevant for 

219 SQLite backend which always creates new files by default. 

220 """ 

221 # engine is reused between multiple processes, make sure that we don't 

222 # share connections by disabling pool (by using NullPool class) 

223 kw: MutableMapping[str, Any] = dict(config.connection_config.extra_parameters) 

224 conn_args: dict[str, Any] = {} 

225 if not config.connection_config.connection_pool: 225 ↛ 226line 225 didn't jump to line 226 because the condition on line 225 was never true

226 kw.update(poolclass=NullPool) 

227 if config.connection_config.isolation_level is not None: 227 ↛ 228line 227 didn't jump to line 228 because the condition on line 227 was never true

228 kw.update(isolation_level=config.connection_config.isolation_level) 

229 elif config.db_url.startswith("sqlite"): 229 ↛ 232line 229 didn't jump to line 232 because the condition on line 229 was always true

230 # Use READ_UNCOMMITTED as default value for sqlite. 

231 kw.update(isolation_level="READ_UNCOMMITTED") 

232 if config.connection_config.connection_timeout is not None: 

233 if config.db_url.startswith("sqlite"): 233 ↛ 235line 233 didn't jump to line 235 because the condition on line 233 was always true

234 conn_args.update(timeout=config.connection_config.connection_timeout) 

235 elif config.db_url.startswith(("postgresql", "mysql")): 

236 conn_args.update(connect_timeout=int(config.connection_config.connection_timeout)) 

237 kw.update(connect_args=conn_args) 

238 engine = sqlalchemy.create_engine(cls._connection_url(config.db_url, create=create), **kw) 

239 

240 if engine.dialect.name == "sqlite": 240 ↛ 244line 240 didn't jump to line 244 because the condition on line 240 was always true

241 # Need to enable foreign keys on every new connection. 

242 sqlalchemy.event.listen(engine, "connect", _onSqlite3Connect) 

243 

244 return engine 

245 

246 @classmethod 

247 def _connection_url(cls, config_url: str, *, create: bool) -> sqlalchemy.engine.URL | str: 

248 """Generate a complete URL for database with proper credentials. 

249 

250 Parameters 

251 ---------- 

252 config_url : `str` 

253 Database URL as specified in configuration. 

254 create : `bool` 

255 Whether to try to create new database file, only relevant for 

256 SQLite backend which always creates new files by default. 

257 

258 Returns 

259 ------- 

260 connection_url : `sqlalchemy.engine.URL` or `str` 

261 Connection URL including credentials. 

262 """ 

263 # Allow 3rd party authentication mechanisms by assuming connection 

264 # string is correct when we can not recognize (dialect, host, database) 

265 # matching keys. 

266 components = urllib.parse.urlparse(config_url) 

267 if all((components.scheme is not None, components.hostname is not None, components.path is not None)): 267 ↛ 268line 267 didn't jump to line 268 because the condition on line 267 was never true

268 try: 

269 db_auth = DbAuth() 

270 config_url = db_auth.getUrl(config_url) 

271 except DbAuthNotFoundError: 

272 # Credentials file doesn't exist or no matching credentials, 

273 # use default auth. 

274 pass 

275 

276 # SQLite has a nasty habit creating empty databases when they do not 

277 # exist, tell it not to do that unless we do need to create it. 

278 if not create: 

279 config_url = cls._update_sqlite_url(config_url) 

280 

281 return config_url 

282 

283 @classmethod 

284 def _update_sqlite_url(cls, url_string: str) -> str: 

285 """If URL refers to sqlite dialect, update it so that the backend does 

286 not try to create database file if it does not exist already. 

287 

288 Parameters 

289 ---------- 

290 url_string : `str` 

291 Connection string. 

292 

293 Returns 

294 ------- 

295 url_string : `str` 

296 Possibly updated connection string. 

297 """ 

298 try: 

299 url = sqlalchemy.make_url(url_string) 

300 except sqlalchemy.exc.SQLAlchemyError: 

301 # If parsing fails it means some special format, likely not 

302 # sqlite so we just return it unchanged. 

303 return url_string 

304 

305 if url.get_backend_name() == "sqlite": 305 ↛ 328line 305 didn't jump to line 328 because the condition on line 305 was always true

306 # Massage url so that database name starts with "file:" and 

307 # option string has "mode=rw&uri=true". Database name 

308 # should look like a path (:memory: is not supported by 

309 # Apdb, but someone could still try to use it). 

310 database = url.database 

311 if database and not database.startswith((":", "file:")): 

312 query = dict(url.query, mode="rw", uri="true") 

313 # If ``database`` is an absolute path then original URL should 

314 # include four slashes after "sqlite:". Humans are bad at 

315 # counting things beyond four and sometimes an extra slash gets 

316 # added unintentionally, which causes sqlite to treat initial 

317 # element as "authority" and to complain. Strip extra slashes 

318 # at the start of the path to avoid that (DM-46077). 

319 if database.startswith("//"): 319 ↛ 320line 319 didn't jump to line 320 because the condition on line 319 was never true

320 warnings.warn( 

321 f"Database URL contains extra leading slashes which will be removed: {url}", 

322 stacklevel=3, 

323 ) 

324 database = "/" + database.lstrip("/") 

325 url = url.set(database=f"file:{database}", query=query) 

326 url_string = url.render_as_string() 

327 

328 return url_string 

329 

330 @classmethod 

331 def _versionCheck(cls, metadata: ApdbMetadataSql, schema_version: VersionTuple) -> VersionTuple: 

332 """Check schema version compatibility and return the database schema 

333 version. 

334 """ 

335 

336 def _get_version(key: str) -> VersionTuple: 

337 """Retrieve version number from given metadata key.""" 

338 version_str = metadata.get(key) 

339 if version_str is None: 339 ↛ 341line 339 didn't jump to line 341 because the condition on line 339 was never true

340 # Should not happen with existing metadata table. 

341 raise RuntimeError(f"Version key {key!r} does not exist in metadata table.") 

342 return VersionTuple.fromString(version_str) 

343 

344 db_schema_version = _get_version(cls.metadataSchemaVersionKey) 

345 db_code_version = _get_version(cls.metadataCodeVersionKey) 

346 

347 # For now there is no way to make read-only APDB instances, assume that 

348 # any access can do updates. 

349 if not schema_version.checkCompatibility(db_schema_version): 

350 raise IncompatibleVersionError( 

351 f"Configured schema version {schema_version} " 

352 f"is not compatible with database version {db_schema_version}" 

353 ) 

354 if not cls.apdbImplementationVersion().checkCompatibility(db_code_version): 354 ↛ 355line 354 didn't jump to line 355 because the condition on line 354 was never true

355 raise IncompatibleVersionError( 

356 f"Current code version {cls.apdbImplementationVersion()} " 

357 f"is not compatible with database version {db_code_version}" 

358 ) 

359 

360 # Check replica code version only if replica is enabled. Sort of 

361 # chicken and egg problem - `enable_replica` is a part of frozen 

362 # configuration, but we cannot read frozen configuration until we 

363 # validate versions. Assume that if the replica version is present 

364 # then replication is enabled. 

365 if metadata.get(cls.metadataReplicaVersionKey) is not None: 

366 db_replica_version = _get_version(cls.metadataReplicaVersionKey) 

367 code_replica_version = ApdbSqlReplica.apdbReplicaImplementationVersion() 

368 if not code_replica_version.checkCompatibility(db_replica_version): 368 ↛ 369line 368 didn't jump to line 369 because the condition on line 368 was never true

369 raise IncompatibleVersionError( 

370 f"Current replication code version {code_replica_version} " 

371 f"is not compatible with database version {db_replica_version}" 

372 ) 

373 

374 return db_schema_version 

375 

376 @classmethod 

377 def apdbImplementationVersion(cls) -> VersionTuple: 

378 """Return version number for current APDB implementation. 

379 

380 Returns 

381 ------- 

382 version : `VersionTuple` 

383 Version of the code defined in implementation class. 

384 """ 

385 return VERSION 

386 

387 @classmethod 

388 def init_database( 

389 cls, 

390 db_url: str, 

391 *, 

392 schema_file: str | None = None, 

393 ss_schema_file: str | None = None, 

394 read_sources_months: int | None = None, 

395 read_forced_sources_months: int | None = None, 

396 enable_replica: bool = False, 

397 connection_timeout: int | None = None, 

398 dia_object_index: str | None = None, 

399 htm_level: int | None = None, 

400 htm_index_column: str | None = None, 

401 ra_dec_columns: tuple[str, str] | None = None, 

402 prefix: str | None = None, 

403 namespace: str | None = None, 

404 drop: bool = False, 

405 ) -> ApdbSqlConfig: 

406 """Initialize new APDB instance and make configuration object for it. 

407 

408 Parameters 

409 ---------- 

410 db_url : `str` 

411 SQLAlchemy database URL. 

412 schema_file : `str`, optional 

413 Location of (YAML) configuration file with APDB schema. If not 

414 specified then default location will be used. 

415 ss_schema_file : `str`, optional 

416 Location of (YAML) configuration file with SSO schema. If not 

417 specified then default location will be used. 

418 read_sources_months : `int`, optional 

419 Number of months of history to read from DiaSource. 

420 read_forced_sources_months : `int`, optional 

421 Number of months of history to read from DiaForcedSource. 

422 enable_replica : `bool`, optional 

423 If True, make additional tables used for replication to PPDB. 

424 connection_timeout : `int`, optional 

425 Database connection timeout in seconds. 

426 dia_object_index : `str`, optional 

427 Indexing mode for DiaObject table. 

428 htm_level : `int`, optional 

429 HTM indexing level. 

430 htm_index_column : `str`, optional 

431 Name of a HTM index column for DiaObject and DiaSource tables. 

432 ra_dec_columns : `tuple` [`str`, `str`], optional 

433 Names of ra/dec columns in DiaObject table. 

434 prefix : `str`, optional 

435 Optional prefix for all table names. 

436 namespace : `str`, optional 

437 Name of the database schema for all APDB tables. If not specified 

438 then default schema is used. 

439 drop : `bool`, optional 

440 If `True` then drop existing tables before re-creating the schema. 

441 

442 Returns 

443 ------- 

444 config : `ApdbSqlConfig` 

445 Resulting configuration object for a created APDB instance. 

446 """ 

447 config = ApdbSqlConfig(db_url=db_url, enable_replica=enable_replica) 

448 if schema_file is not None: 448 ↛ 450line 448 didn't jump to line 450 because the condition on line 448 was always true

449 config.schema_file = schema_file 

450 if ss_schema_file is not None: 

451 config.ss_schema_file = ss_schema_file 

452 if read_sources_months is not None: 

453 config.read_sources_months = read_sources_months 

454 if read_forced_sources_months is not None: 

455 config.read_forced_sources_months = read_forced_sources_months 

456 if connection_timeout is not None: 456 ↛ 457line 456 didn't jump to line 457 because the condition on line 456 was never true

457 config.connection_config.connection_timeout = connection_timeout 

458 if dia_object_index is not None: 

459 config.dia_object_index = dia_object_index 

460 if htm_level is not None: 460 ↛ 461line 460 didn't jump to line 461 because the condition on line 460 was never true

461 config.pixelization.htm_level = htm_level 

462 if htm_index_column is not None: 462 ↛ 463line 462 didn't jump to line 463 because the condition on line 462 was never true

463 config.pixelization.htm_index_column = htm_index_column 

464 if ra_dec_columns is not None: 464 ↛ 465line 464 didn't jump to line 465 because the condition on line 464 was never true

465 config.ra_dec_columns = ra_dec_columns 

466 if prefix is not None: 466 ↛ 467line 466 didn't jump to line 467 because the condition on line 466 was never true

467 config.prefix = prefix 

468 if namespace is not None: 468 ↛ 469line 468 didn't jump to line 469 because the condition on line 468 was never true

469 config.namespace = namespace 

470 

471 cls._makeSchema(config, drop=drop) 

472 

473 # SQLite has a nasty habit of creating empty database by default, 

474 # update URL in config file to disable that behavior. 

475 config.db_url = cls._update_sqlite_url(config.db_url) 

476 

477 return config 

478 

479 def get_replica(self) -> ApdbSqlReplica: 

480 """Return `ApdbReplica` instance for this database.""" 

481 return ApdbSqlReplica(self._schema, self._engine, self._db_schema_version) 

482 

483 def tableRowCount(self) -> dict[str, int]: 

484 """Return dictionary with the table names and row counts. 

485 

486 Used by ``ap_proto`` to keep track of the size of the database tables. 

487 Depending on database technology this could be expensive operation. 

488 

489 Returns 

490 ------- 

491 row_counts : `dict` 

492 Dict where key is a table name and value is a row count. 

493 """ 

494 res = {} 

495 tables = [ApdbTables.DiaObject, ApdbTables.DiaSource, ApdbTables.DiaForcedSource] 

496 if self.config.dia_object_index == "last_object_table": 

497 tables.append(ApdbTables.DiaObjectLast) 

498 with self._engine.begin() as conn: 

499 for table in tables: 

500 sa_table = self._schema.get_table(table) 

501 stmt = sql.select(func.count()).select_from(sa_table) 

502 count: int = conn.execute(stmt).scalar_one() 

503 res[table.name] = count 

504 

505 return res 

506 

507 def getConfig(self) -> ApdbSqlConfig: 

508 # docstring is inherited from a base class 

509 return self.config 

510 

511 def tableDef(self, table: ApdbTables) -> Table | None: 

512 # docstring is inherited from a base class 

513 return self._schema.tableSchemas.get(table) 

514 

515 @classmethod 

516 def _makeSchema(cls, config: ApdbConfig, drop: bool = False) -> None: 

517 # docstring is inherited from a base class 

518 

519 if not isinstance(config, ApdbSqlConfig): 519 ↛ 520line 519 didn't jump to line 520 because the condition on line 519 was never true

520 raise TypeError(f"Unexpected type of configuration object: {type(config)}") 

521 

522 engine = cls._makeEngine(config, create=True) 

523 

524 table_schema = ApdbSchema(config.schema_file, config.ss_schema_file) 

525 

526 # Ask schema class to create all tables. 

527 schema = ApdbSqlSchema( 

528 table_schema=table_schema, 

529 engine=engine, 

530 dia_object_index=config.dia_object_index, 

531 prefix=config.prefix, 

532 namespace=config.namespace, 

533 htm_index_column=config.pixelization.htm_index_column, 

534 enable_replica=config.enable_replica, 

535 ) 

536 schema.makeSchema(drop=drop) 

537 

538 # Need metadata table to store few items in it. 

539 meta_table = schema.get_table(ApdbTables.metadata) 

540 apdb_meta = ApdbMetadataSql(engine, meta_table) 

541 

542 # Fill version numbers, overwrite if they are already there. 

543 apdb_meta.set(cls.metadataSchemaVersionKey, str(table_schema.schemaVersion()), force=True) 

544 apdb_meta.set(cls.metadataCodeVersionKey, str(cls.apdbImplementationVersion()), force=True) 

545 if config.enable_replica: 

546 # Only store replica code version if replica is enabled. 

547 apdb_meta.set( 

548 cls.metadataReplicaVersionKey, 

549 str(ApdbSqlReplica.apdbReplicaImplementationVersion()), 

550 force=True, 

551 ) 

552 

553 # Store frozen part of a configuration in metadata. 

554 freezer = ApdbConfigFreezer[ApdbSqlConfig](cls._frozen_parameters) 

555 apdb_meta.set(cls.metadataConfigKey, freezer.to_json(config), force=True) 

556 

557 def getDiaObjects(self, region: Region) -> pandas.DataFrame: 

558 # docstring is inherited from a base class 

559 

560 # decide what columns we need 

561 if self.config.dia_object_index == "last_object_table": 

562 table_enum = ApdbTables.DiaObjectLast 

563 else: 

564 table_enum = ApdbTables.DiaObject 

565 table = self._schema.get_table(table_enum) 

566 if not self.config.dia_object_columns: 566 ↛ 569line 566 didn't jump to line 569 because the condition on line 566 was always true

567 columns = self._schema.get_apdb_columns(table_enum) 

568 else: 

569 columns = [table.c[col] for col in self.config.dia_object_columns] 

570 query = sql.select(*columns) 

571 

572 # build selection 

573 query = query.where(self._filterRegion(table, region)) 

574 

575 validity_end_column = self._timestamp_column_name("validityEnd") 

576 

577 # select latest version of objects 

578 if self.config.dia_object_index != "last_object_table": 

579 query = query.where(table.columns[validity_end_column] == None) # noqa: E711 

580 

581 # _LOG.debug("query: %s", query) 

582 

583 # execute select 

584 with self._timer("select_time", tags={"table": "DiaObject"}) as timer: 

585 with self._engine.begin() as conn: 

586 result = conn.execute(query) 

587 column_defs = self._table_schema.tableSchemas[table_enum].columns 

588 table_data = ApdbSqlTableData(result, column_defs) 

589 objects = table_data.to_pandas() 

590 timer.add_values(row_count=len(objects)) 

591 _LOG.debug("found %s DiaObjects", len(objects)) 

592 return self._fix_result_timestamps(objects) 

593 

594 def getDiaSources( 

595 self, 

596 region: Region, 

597 object_ids: Iterable[int] | None, 

598 visit_time: astropy.time.Time, 

599 start_time: astropy.time.Time | None = None, 

600 ) -> pandas.DataFrame | None: 

601 # docstring is inherited from a base class 

602 if start_time is None and self.config.read_sources_months == 0: 

603 _LOG.debug("Skip DiaSources fetching") 

604 return None 

605 

606 if start_time is None: 

607 start_time_mjdTai = _make_midpointMjdTai_start(visit_time, self.config.read_sources_months) 

608 else: 

609 start_time_mjdTai = float(start_time.tai.mjd) 

610 _LOG.debug("start_time_mjdTai = %.6f", start_time_mjdTai) 

611 

612 if object_ids is None: 

613 # region-based select 

614 return self._getDiaSourcesInRegion(region, start_time_mjdTai) 

615 else: 

616 return self._getDiaSourcesByIDs(list(object_ids), start_time_mjdTai) 

617 

618 def getDiaForcedSources( 

619 self, 

620 region: Region, 

621 object_ids: Iterable[int] | None, 

622 visit_time: astropy.time.Time, 

623 start_time: astropy.time.Time | None = None, 

624 ) -> pandas.DataFrame | None: 

625 # docstring is inherited from a base class 

626 if start_time is None and self.config.read_forced_sources_months == 0: 

627 _LOG.debug("Skip DiaForceSources fetching") 

628 return None 

629 

630 if object_ids is None: 

631 # This implementation does not support region-based selection. In 

632 # the past DiaForcedSource schema did not have ra/dec columns (it 

633 # had x/y columns). ra/dec were added at some point, so we could 

634 # add pixelOd column to this table if/when needed. 

635 raise NotImplementedError("Region-based selection is not supported") 

636 

637 # TODO: DateTime.MJD must be consistent with code in ap_association, 

638 # alternatively we can fill midpointMjdTai ourselves in store() 

639 if start_time is None: 

640 start_time_mjdTai = _make_midpointMjdTai_start(visit_time, self.config.read_forced_sources_months) 

641 else: 

642 start_time_mjdTai = float(start_time.tai.mjd) 

643 _LOG.debug("start_time_mjdTai = %.6f", start_time_mjdTai) 

644 

645 with self._timer("select_time", tags={"table": "DiaForcedSource"}) as timer: 

646 sources = self._getSourcesByIDs(ApdbTables.DiaForcedSource, list(object_ids), start_time_mjdTai) 

647 timer.add_values(row_count=len(sources)) 

648 

649 _LOG.debug("found %s DiaForcedSources", len(sources)) 

650 return sources 

651 

652 def getDiaObjectsForDedup(self, since: astropy.time.Time | None = None) -> pandas.DataFrame: 

653 # docstring is inherited from a base class 

654 

655 if since is None: 

656 # Read last deduplication time from metadata. 

657 dedup_str = self._metadata.get(self.metadataDedupKey) 

658 if dedup_str is not None: 

659 dedup_state = json.loads(dedup_str) 

660 dedup_time_str = dedup_state["dedup_time_iso_tai"] 

661 since = astropy.time.Time(dedup_time_str, format="iso", scale="tai") 

662 

663 validity_start_column = self._timestamp_column_name("validityStart") 

664 

665 # decide what columns we need 

666 if self.config.dia_object_index == "last_object_table": 

667 table_enum = ApdbTables.DiaObjectLast 

668 else: 

669 table_enum = ApdbTables.DiaObject 

670 table = self._schema.get_table(table_enum) 

671 

672 if not self.config.dia_object_columns_for_dedup: 672 ↛ 673line 672 didn't jump to line 673 because the condition on line 672 was never true

673 columns = self._schema.get_apdb_columns(table_enum) 

674 else: 

675 column_names = list(self.config.dia_object_columns_for_dedup) 

676 if validity_start_column not in column_names: 676 ↛ 678line 676 didn't jump to line 678 because the condition on line 676 was always true

677 column_names.insert(0, validity_start_column) 

678 if "diaObjectId" not in column_names: 678 ↛ 680line 678 didn't jump to line 680 because the condition on line 678 was always true

679 column_names.insert(0, "diaObjectId") 

680 columns = [table.columns[col] for col in column_names] 

681 

682 query = sql.select(*columns) 

683 

684 # build selection 

685 if since is not None: 

686 timestamp = self._timestamp_column_value(since) 

687 query = query.where(table.columns[validity_start_column] >= timestamp) 

688 

689 # execute select 

690 with self._timer("select_time", tags={"table": "DiaObject"}) as timer: 

691 with self._engine.begin() as conn: 

692 result = conn.execute(query) 

693 column_defs = self._table_schema.tableSchemas[table_enum].columns 

694 table_data = ApdbSqlTableData(result, column_defs) 

695 objects = table_data.to_pandas() 

696 timer.add_values(row_count=len(objects)) 

697 _LOG.debug("found %s DiaObjects", len(objects)) 

698 return self._fix_result_timestamps(objects) 

699 

700 def getDiaSourcesForDiaObjects( 

701 self, objects: list[DiaObjectId], start_time: astropy.time.Time, max_dist_arcsec: float = 1.0 

702 ) -> pandas.DataFrame: 

703 # docstring is inherited from a base class 

704 object_ids = {object_id.diaObjectId for object_id in objects} 

705 return self._getDiaSourcesByIDs(list(object_ids), float(start_time.tai.mjd)) 

706 

707 def containsVisitDetector( 

708 self, 

709 visit: int, 

710 detector: int, 

711 region: Region | None = None, 

712 visit_time: astropy.time.Time | None = None, 

713 ) -> bool: 

714 # docstring is inherited from a base class 

715 src_table: sqlalchemy.schema.Table = self._schema.get_table(ApdbTables.DiaSource) 

716 frcsrc_table: sqlalchemy.schema.Table = self._schema.get_table(ApdbTables.DiaForcedSource) 

717 # Query should load only one leaf page of the index 

718 query1 = sql.select(src_table.c.visit).filter_by(visit=visit, detector=detector).limit(1) 

719 

720 with self._engine.begin() as conn: 

721 result = conn.execute(query1).scalar_one_or_none() 

722 if result is not None: 

723 return True 

724 else: 

725 # Backup query if an image was processed but had no diaSources 

726 query2 = sql.select(frcsrc_table.c.visit).filter_by(visit=visit, detector=detector).limit(1) 

727 result = conn.execute(query2).scalar_one_or_none() 

728 return result is not None 

729 

730 def store( 

731 self, 

732 visit_time: astropy.time.Time, 

733 objects: pandas.DataFrame, 

734 sources: pandas.DataFrame | None = None, 

735 forced_sources: pandas.DataFrame | None = None, 

736 ) -> None: 

737 # docstring is inherited from a base class 

738 objects = self._fix_input_timestamps(objects) 

739 if sources is not None: 

740 sources = self._fix_input_timestamps(sources) 

741 if forced_sources is not None: 

742 forced_sources = self._fix_input_timestamps(forced_sources) 

743 

744 # We want to run all inserts in one transaction. 

745 with self._engine.begin() as connection: 

746 replica_chunk: ReplicaChunk | None = None 

747 if self._schema.replication_enabled: 

748 replica_chunk = ReplicaChunk.make_replica_chunk(visit_time, self.config.replica_chunk_seconds) 

749 self._storeReplicaChunk(replica_chunk, connection) 

750 

751 # fill pixelId column for DiaObjects 

752 objects = self._add_spatial_index(objects) 

753 self._storeDiaObjects(objects, visit_time, replica_chunk, connection) 

754 

755 if sources is not None: 

756 # fill pixelId column for DiaSources 

757 sources = self._add_spatial_index(sources) 

758 self._storeDiaSources(sources, replica_chunk, connection) 

759 

760 if forced_sources is not None: 

761 self._storeDiaForcedSources(forced_sources, replica_chunk, connection) 

762 

763 def reassignDiaSourcesToDiaObjects( 

764 self, 

765 idMap: Mapping[DiaSourceId, int], 

766 *, 

767 increment_nDiaSources: bool = True, 

768 decrement_nDiaSources: bool = True, 

769 ) -> None: 

770 # docstring is inherited from a base class 

771 

772 new_object_ids = set(idMap.values()) 

773 source_ids = {source.diaSourceId for source in idMap} 

774 

775 current_time = self._current_time() 

776 current_time_ns = int(current_time.unix_tai * 1e9) 

777 

778 with self._engine.begin() as conn: 

779 # Make sure that all DiaSources exist. 

780 found_sources = self._get_diasource_data(conn, source_ids, "diaObjectId") 

781 if missing_ids := (source_ids - {row.diaSourceId for row in found_sources}): 781 ↛ 782line 781 didn't jump to line 782 because the condition on line 781 was never true

782 raise LookupError(f"Some source IDs are missing from DiaSource table: {missing_ids}") 

783 original_object_ids = {row.diaSourceId: row.diaObjectId for row in found_sources} 

784 

785 # Make sure that all DiaObjects exist, we also want to know 

786 # nDiaSources count for current and new records because we want to 

787 # send updated values to replica. 

788 all_object_ids = new_object_ids | set(original_object_ids.values()) 

789 found_objects = self._get_diaobject_data(conn, all_object_ids, "ra", "dec", "nDiaSources") 

790 if missing_ids := (new_object_ids - {row.diaObjectId for row in found_objects}): 790 ↛ 791line 790 didn't jump to line 791 because the condition on line 790 was never true

791 raise LookupError(f"Some object IDs are missing from DiaObject table: {missing_ids}") 

792 

793 found_objects_by_id = {row.diaObjectId: row for row in found_objects} 

794 

795 update_records: list[ApdbUpdateRecord] = [] 

796 update_order = 0 

797 

798 # Update DiaSources. 

799 source_table = self._schema.get_table(ApdbTables.DiaSource) 

800 for source, diaObjectId in idMap.items(): 

801 update = ( 

802 source_table.update() 

803 .where(source_table.columns["diaSourceId"] == source.diaSourceId) 

804 .values(diaObjectId=diaObjectId) 

805 ) 

806 conn.execute(update) 

807 

808 if self._schema.replication_enabled: 

809 update_records.append( 

810 ApdbReassignDiaSourceToDiaObjectRecord( 

811 diaSourceId=source.diaSourceId, 

812 ra=source.ra, 

813 dec=source.dec, 

814 midpointMjdTai=source.midpointMjdTai, 

815 diaObjectId=diaObjectId, 

816 update_time_ns=current_time_ns, 

817 update_order=update_order, 

818 ) 

819 ) 

820 update_order += 1 

821 

822 # DiaObject tables to update. 

823 object_tables = [self._schema.get_table(ApdbTables.DiaObject)] 

824 if self.config.dia_object_index == "last_object_table": 

825 object_tables.append(self._schema.get_table(ApdbTables.DiaObjectLast)) 

826 

827 # Things to increment/decrement. 

828 increments: Counter = Counter() 

829 if increment_nDiaSources: 

830 increments.update(idMap.values()) 

831 if decrement_nDiaSources: 

832 increments.subtract(original_object_ids[source_id.diaSourceId] for source_id in idMap) 

833 

834 if increments: 

835 for table in object_tables: 

836 for diaObjectId, increment in increments.items(): 

837 update = ( 

838 table.update() 

839 .where(table.columns["diaObjectId"] == diaObjectId) 

840 .values(nDiaSources=table.columns["nDiaSources"] + increment) 

841 ) 

842 conn.execute(update) 

843 

844 # Also send updated values to replica. 

845 if self._schema.replication_enabled: 

846 for diaObjectId, increment in increments.items(): 

847 dia_object = found_objects_by_id[diaObjectId] 

848 update_records.append( 

849 ApdbUpdateNDiaSourcesRecord( 

850 diaObjectId=diaObjectId, 

851 ra=dia_object.ra, 

852 dec=dia_object.dec, 

853 nDiaSources=dia_object.nDiaSources + increment, 

854 update_time_ns=current_time_ns, 

855 update_order=update_order, 

856 ) 

857 ) 

858 update_order += 1 

859 

860 if update_records: 

861 replica_chunk = ReplicaChunk.make_replica_chunk( 

862 current_time, self.config.replica_chunk_seconds 

863 ) 

864 self._storeUpdateRecords(update_records, replica_chunk, connection=conn, store_chunk=True) 

865 

866 def setValidityEnd( 

867 self, objects: list[DiaObjectId], validityEnd: astropy.time.Time, raise_on_missing_id: bool = False 

868 ) -> int: 

869 # docstring is inherited from a base class 

870 

871 with self._engine.begin() as conn: 

872 return self._setValidityEnd(conn, objects, validityEnd, raise_on_missing_id=raise_on_missing_id) 

873 

874 def _setValidityEnd( 

875 self, 

876 conn: sqlalchemy.Connection, 

877 objects: list[DiaObjectId], 

878 validityEnd: astropy.time.Time, 

879 raise_on_missing_id: bool = False, 

880 ) -> int: 

881 # docstring is inherited from a base class 

882 if not objects: 882 ↛ 883line 882 didn't jump to line 883 because the condition on line 882 was never true

883 return 0 

884 

885 requested_ids = {obj.diaObjectId for obj in objects} 

886 

887 validity_end_column = self._timestamp_column_name("validityEnd") 

888 validityEnd_value = self._timestamp_column_value(validityEnd) 

889 

890 # Find all matching DiaObjects with validityEnd = NULL. 

891 table = self._schema.get_table(ApdbTables.DiaObject) 

892 query = sql.select(table.columns["diaObjectId"]).where( 

893 sqlalchemy.and_( 

894 table.columns["diaObjectId"].in_(sorted(requested_ids)), 

895 table.columns[validity_end_column].is_(None), 

896 ) 

897 ) 

898 

899 result = conn.execute(query) 

900 found_ids = set(result.scalars()) 

901 

902 # Check that we found all that is requested. 

903 if raise_on_missing_id: 

904 if missing_ids := (requested_ids - found_ids): 904 ↛ 908line 904 didn't jump to line 908 because the condition on line 904 was always true

905 raise LookupError(f"Some object IDs are missing from DiaObject table: {missing_ids}") 

906 

907 # Filter existing records. 

908 if len(objects) != len(found_ids): 

909 objects = [obj for obj in objects if obj.diaObjectId in found_ids] 

910 

911 if not objects: 

912 return 0 

913 

914 values = {validity_end_column: validityEnd_value} 

915 update = ( 

916 table.update() 

917 .where( 

918 sqlalchemy.and_( 

919 table.columns["diaObjectId"].in_(sorted(found_ids)), 

920 table.columns[validity_end_column].is_(None), 

921 ) 

922 ) 

923 .values(**values) 

924 ) 

925 result = conn.execute(update) 

926 if result.rowcount != len(found_ids): 926 ↛ 927line 926 didn't jump to line 927 because the condition on line 926 was never true

927 raise RuntimeError( 

928 f"Unexpected mismatch in the number of records updated. Object IDs = {found_ids}" 

929 ) 

930 

931 # Also drop them from DiaObjectLast. 

932 if self.config.dia_object_index == "last_object_table": 

933 last_table = self._schema.get_table(ApdbTables.DiaObjectLast) 

934 delete = last_table.delete().where(last_table.columns["diaObjectId"].in_(sorted(found_ids))) 

935 result = conn.execute(delete) 

936 if result.rowcount != len(found_ids): 936 ↛ 937line 936 didn't jump to line 937 because the condition on line 936 was never true

937 raise RuntimeError( 

938 f"Unexpected mismatch in the number of records deleted. Object IDs = {found_ids}" 

939 ) 

940 

941 # If replication is enabled then send all updates. 

942 if self._schema.replication_enabled: 

943 current_time = self._current_time() 

944 current_time_ns = int(current_time.unix_tai * 1e9) 

945 replica_chunk = ReplicaChunk.make_replica_chunk(current_time, self.config.replica_chunk_seconds) 

946 

947 update_records = [ 

948 ApdbCloseDiaObjectValidityRecord( 

949 diaObjectId=obj.diaObjectId, 

950 ra=obj.ra, 

951 dec=obj.dec, 

952 update_time_ns=current_time_ns, 

953 update_order=index, 

954 validityEndMjdTai=float(validityEnd.tai.mjd), 

955 nDiaSources=None, 

956 ) 

957 for index, obj in enumerate(objects) 

958 ] 

959 

960 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True, connection=conn) 

961 

962 return len(objects) 

963 

964 def resetDedup(self, dedup_time: astropy.time.Time | None = None) -> None: 

965 # docstring is inherited from a base class 

966 

967 # SQL backend does not have separate dedup tables, nothing to delete, 

968 # only save last dedup time in metadata. 

969 if dedup_time is None: 969 ↛ 970line 969 didn't jump to line 970 because the condition on line 969 was never true

970 dedup_time = self._current_time() 

971 data = {"dedup_time_iso_tai": dedup_time.tai.to_value("iso")} 

972 data_json = json.dumps(data) 

973 self._metadata.set(self.metadataDedupKey, data_json, force=True) 

974 

975 def reassignDiaSources(self, idMap: Mapping[int, int]) -> None: 

976 # docstring is inherited from a base class 

977 

978 timestamp: float | datetime.datetime 

979 now = self._current_time() 

980 timestamp_column = self._timestamp_column_name("ssObjectReassocTime") 

981 timestamp = self._timestamp_column_value(now) 

982 

983 table = self._schema.get_table(ApdbTables.DiaSource) 

984 query = table.update().where(table.columns["diaSourceId"] == sql.bindparam("srcId")) 

985 

986 with self._engine.begin() as conn: 

987 # Need to make sure that every ID exists in the database, but 

988 # executemany may not support rowcount, so iterate and check what 

989 # is missing. 

990 missing_ids: list[int] = [] 

991 for key, value in idMap.items(): 

992 params = { 

993 "srcId": key, 

994 "diaObjectId": 0, 

995 "ssObjectId": value, 

996 timestamp_column: timestamp, 

997 } 

998 result = conn.execute(query, params) 

999 if result.rowcount == 0: 

1000 missing_ids.append(key) 

1001 if missing_ids: 

1002 missing = ",".join(str(item) for item in missing_ids) 

1003 raise ValueError(f"Following DiaSource IDs do not exist in the database: {missing}") 

1004 

1005 def withdrawDiaSources( 

1006 self, 

1007 diaSourceIds: Iterable[DiaSourceId], 

1008 *, 

1009 timeWithdrawn: astropy.time.Time | None = None, 

1010 ) -> None: 

1011 # docstring is inherited from a base class 

1012 

1013 diaSourceIds = list(diaSourceIds) 

1014 source_ids = {source.diaSourceId for source in diaSourceIds} 

1015 

1016 if timeWithdrawn is None: 1016 ↛ 1017line 1016 didn't jump to line 1017 because the condition on line 1016 was never true

1017 timeWithdrawn = self._current_time() 

1018 time_value = self._timestamp_column_value(timeWithdrawn) 

1019 column_name = self._timestamp_column_name("time_withdrawn") 

1020 

1021 with self._engine.begin() as conn: 

1022 # Make sure that all DiaSources exist. 

1023 found_sources = self._get_diasource_data(conn, source_ids, "diaObjectId", column_name) 

1024 if missing_ids := (source_ids - {row.diaSourceId for row in found_sources}): 1024 ↛ 1025line 1024 didn't jump to line 1025 because the condition on line 1024 was never true

1025 raise LookupError(f"Some source IDs are missing from DiaSource table: {missing_ids}") 

1026 

1027 # Ignore sources already withdrawn. 

1028 source_ids = {row.diaSourceId for row in found_sources if getattr(row, column_name) is None} 

1029 diaSourceIds = [source for source in diaSourceIds if source.diaSourceId in source_ids] 

1030 

1031 # Set time_withdrawn for sources. 

1032 table = self._schema.get_table(ApdbTables.DiaSource) 

1033 where = table.columns["diaSourceId"].in_(sorted(source_ids)) 

1034 update = table.update().where(where).values({column_name: time_value}) 

1035 conn.execute(update) 

1036 

1037 # If replication is enabled then send all updates. 

1038 if self._schema.replication_enabled: 

1039 current_time = self._current_time() 

1040 current_time_ns = int(current_time.unix_tai * 1e9) 

1041 replica_chunk = ReplicaChunk.make_replica_chunk( 

1042 current_time, self.config.replica_chunk_seconds 

1043 ) 

1044 

1045 update_records = [ 

1046 ApdbWithdrawDiaSourceRecord( 

1047 diaSourceId=source.diaSourceId, 

1048 ra=source.ra, 

1049 dec=source.dec, 

1050 midpointMjdTai=source.midpointMjdTai, 

1051 update_time_ns=current_time_ns, 

1052 update_order=update_order, 

1053 timeWithdrawnMjdTai=float(timeWithdrawn.tai.mjd), 

1054 ) 

1055 for update_order, source in enumerate(diaSourceIds) 

1056 ] 

1057 if update_records: 1057 ↛ exitline 1057 didn't jump to the function exit

1058 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True, connection=conn) 

1059 

1060 def withdrawDiaForcedSources( 

1061 self, 

1062 diaForcedSourceIds: Iterable[DiaForcedSourceId], 

1063 *, 

1064 timeWithdrawn: astropy.time.Time | None = None, 

1065 ) -> None: 

1066 # docstring is inherited from a base class 

1067 

1068 def _fsrc_id(fsource: Any) -> tuple[int, int, int]: 

1069 return (fsource.diaObjectId, fsource.visit, fsource.detector) 

1070 

1071 diaForcedSourceIds = list(diaForcedSourceIds) 

1072 source_ids = {_fsrc_id(source) for source in diaForcedSourceIds} 

1073 

1074 if timeWithdrawn is None: 1074 ↛ 1075line 1074 didn't jump to line 1075 because the condition on line 1074 was never true

1075 timeWithdrawn = self._current_time() 

1076 time_value = self._timestamp_column_value(timeWithdrawn) 

1077 column_name = self._timestamp_column_name("time_withdrawn") 

1078 

1079 with self._engine.begin() as conn: 

1080 # Make sure that all DiaForcedSources exist. 

1081 table = self._schema.get_table(ApdbTables.DiaForcedSource) 

1082 id_columns = [table.columns["diaObjectId"], table.columns["visit"], table.columns["detector"]] 

1083 where = sqlalchemy.tuple_(*id_columns).in_(sorted(source_ids)) 

1084 query = sql.select(table.columns[column_name], *id_columns).where(where) 

1085 result = list(conn.execute(query)) 

1086 if missing_ids := (source_ids - {_fsrc_id(row) for row in result}): 1086 ↛ 1087line 1086 didn't jump to line 1087 because the condition on line 1086 was never true

1087 raise LookupError(f"Some source IDs are missing from DiaForcedSource table: {missing_ids}") 

1088 

1089 # Ignore sources already withdrawn. 

1090 source_ids = {_fsrc_id(row) for row in result if getattr(row, column_name) is None} 

1091 diaForcedSourceIds = [source for source in diaForcedSourceIds if _fsrc_id(source) in source_ids] 

1092 

1093 where = sqlalchemy.tuple_(*id_columns).in_(sorted(source_ids)) 

1094 update = table.update().where(where).values({column_name: time_value}) 

1095 conn.execute(update) 

1096 

1097 # If replication is enabled then send all updates. 

1098 if self._schema.replication_enabled: 

1099 current_time = self._current_time() 

1100 current_time_ns = int(current_time.unix_tai * 1e9) 

1101 replica_chunk = ReplicaChunk.make_replica_chunk( 

1102 current_time, self.config.replica_chunk_seconds 

1103 ) 

1104 

1105 update_records = [ 

1106 ApdbWithdrawDiaForcedSourceRecord( 

1107 diaObjectId=source.diaObjectId, 

1108 visit=source.visit, 

1109 detector=source.detector, 

1110 ra=source.ra, 

1111 dec=source.dec, 

1112 midpointMjdTai=source.midpointMjdTai, 

1113 update_time_ns=current_time_ns, 

1114 update_order=index, 

1115 timeWithdrawnMjdTai=float(timeWithdrawn.tai.mjd), 

1116 ) 

1117 for index, source in enumerate(diaForcedSourceIds) 

1118 ] 

1119 

1120 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True, connection=conn) 

1121 

1122 def countUnassociatedObjects(self) -> int: 

1123 # docstring is inherited from a base class 

1124 

1125 # Retrieve the DiaObject table. 

1126 table: sqlalchemy.schema.Table = self._schema.get_table(ApdbTables.DiaObject) 

1127 

1128 # Construct the sql statement. 

1129 validity_end_column = self._timestamp_column_name("validityEnd") 

1130 stmt = ( 

1131 sql.select(func.count()) 

1132 .select_from(table) 

1133 .where( 

1134 sqlalchemy.and_( 

1135 table.columns["nDiaSources"] == 1, 

1136 table.columns[validity_end_column].is_(None), 

1137 ) 

1138 ) 

1139 ) 

1140 

1141 # Return the count. 

1142 with self._engine.begin() as conn: 

1143 count = conn.execute(stmt).scalar_one() 

1144 

1145 return count 

1146 

1147 @property 

1148 def schema(self) -> ApdbSchema: 

1149 # docstring is inherited from a base class 

1150 return self._table_schema 

1151 

1152 @property 

1153 def metadata(self) -> ApdbMetadata: 

1154 # docstring is inherited from a base class 

1155 return self._metadata 

1156 

1157 @property 

1158 def admin(self) -> ApdbSqlAdmin: 

1159 # docstring is inherited from a base class 

1160 return ApdbSqlAdmin(self.pixelator) 

1161 

1162 def _getDiaSourcesInRegion(self, region: Region, start_time_mjdTai: float) -> pandas.DataFrame: 

1163 """Return catalog of DiaSource instances from given region. 

1164 

1165 Parameters 

1166 ---------- 

1167 region : `lsst.sphgeom.Region` 

1168 Region to search for DIASources. 

1169 start_time_mjdTai : `float` 

1170 Lower bound of time window for the query. 

1171 

1172 Returns 

1173 ------- 

1174 catalog : `pandas.DataFrame` 

1175 Catalog containing DiaSource records. 

1176 """ 

1177 table = self._schema.get_table(ApdbTables.DiaSource) 

1178 columns = self._schema.get_apdb_columns(ApdbTables.DiaSource) 

1179 query = sql.select(*columns) 

1180 

1181 # build selection 

1182 time_filter = table.columns["midpointMjdTai"] > start_time_mjdTai 

1183 where = sql.expression.and_(self._filterRegion(table, region), time_filter) 

1184 query = query.where(where) 

1185 

1186 # execute select 

1187 with self._timer("DiaSource_select_time", tags={"table": "DiaSource"}) as timer: 

1188 with self._engine.begin() as conn: 

1189 result = conn.execute(query) 

1190 column_defs = self._table_schema.tableSchemas[ApdbTables.DiaSource].columns 

1191 table_data = ApdbSqlTableData(result, column_defs) 

1192 sources = table_data.to_pandas() 

1193 timer.add_values(row_counts=len(sources)) 

1194 _LOG.debug("found %s DiaSources", len(sources)) 

1195 return self._fix_result_timestamps(sources) 

1196 

1197 def _getDiaSourcesByIDs(self, object_ids: list[int], start_time_mjdTai: float) -> pandas.DataFrame: 

1198 """Return catalog of DiaSource instances given set of DiaObject IDs. 

1199 

1200 Parameters 

1201 ---------- 

1202 object_ids : 

1203 Collection of DiaObject IDs 

1204 start_time_mjdTai : `float` 

1205 Lower bound of time window for the query. 

1206 

1207 Returns 

1208 ------- 

1209 catalog : `pandas.DataFrame` 

1210 Catalog containing DiaSource records. 

1211 """ 

1212 with self._timer("select_time", tags={"table": "DiaSource"}) as timer: 

1213 sources = self._getSourcesByIDs(ApdbTables.DiaSource, object_ids, start_time_mjdTai) 

1214 timer.add_values(row_count=len(sources)) 

1215 

1216 _LOG.debug("found %s DiaSources", len(sources)) 

1217 return sources 

1218 

1219 def _getSourcesByIDs( 

1220 self, table_enum: ApdbTables, object_ids: list[int], midpointMjdTai_start: float 

1221 ) -> pandas.DataFrame: 

1222 """Return catalog of DiaSource or DiaForcedSource instances given set 

1223 of DiaObject IDs. 

1224 

1225 Parameters 

1226 ---------- 

1227 table : `sqlalchemy.schema.Table` 

1228 Database table. 

1229 object_ids : 

1230 Collection of DiaObject IDs 

1231 midpointMjdTai_start : `float` 

1232 Earliest midpointMjdTai to retrieve. 

1233 

1234 Returns 

1235 ------- 

1236 catalog : `pandas.DataFrame` 

1237 Catalog contaning DiaSource records. `None` is returned if 

1238 ``read_sources_months`` configuration parameter is set to 0 or 

1239 when ``object_ids`` is empty. 

1240 """ 

1241 table = self._schema.get_table(table_enum) 

1242 columns = self._schema.get_apdb_columns(table_enum) 

1243 column_defs = self._table_schema.tableSchemas[table_enum].columns 

1244 

1245 sources: pandas.DataFrame | None = None 

1246 if len(object_ids) <= 0: 

1247 _LOG.debug("ID list is empty, just fetch empty result") 

1248 query = sql.select(*columns).where(sql.literal(False)) 

1249 with self._engine.begin() as conn: 

1250 result = conn.execute(query) 

1251 table_data = ApdbSqlTableData(result, column_defs) 

1252 sources = table_data.to_pandas() 

1253 else: 

1254 data_frames: list[pandas.DataFrame] = [] 

1255 for ids in chunk_iterable(sorted(object_ids), 1000): 

1256 query = sql.select(*columns) 

1257 

1258 # Some types like np.int64 can cause issues with 

1259 # sqlalchemy, convert them to int. 

1260 int_ids = [int(oid) for oid in ids] 

1261 

1262 # select by object id 

1263 query = query.where( 

1264 sql.expression.and_( 

1265 table.columns["diaObjectId"].in_(int_ids), 

1266 table.columns["midpointMjdTai"] >= midpointMjdTai_start, 

1267 ) 

1268 ) 

1269 

1270 # execute select 

1271 with self._engine.begin() as conn: 

1272 result = conn.execute(query) 

1273 table_data = ApdbSqlTableData(result, column_defs) 

1274 data_frames.append(table_data.to_pandas()) 

1275 

1276 if len(data_frames) == 1: 1276 ↛ 1279line 1276 didn't jump to line 1279 because the condition on line 1276 was always true

1277 sources = data_frames[0] 

1278 else: 

1279 sources = pandas.concat(data_frames) 

1280 assert sources is not None, "Catalog cannot be None" 

1281 return self._fix_result_timestamps(sources) 

1282 

1283 def _storeReplicaChunk( 

1284 self, 

1285 replica_chunk: ReplicaChunk, 

1286 connection: sqlalchemy.engine.Connection, 

1287 ) -> None: 

1288 # `visit_time.datetime` returns naive datetime, even though all astropy 

1289 # times are in UTC. Add UTC timezone to timestamp so that database 

1290 # can store a correct value. 

1291 dt = datetime.datetime.fromtimestamp(replica_chunk.last_update_time.unix_tai, tz=datetime.UTC) 

1292 

1293 table = self._schema.get_table(ExtraTables.ApdbReplicaChunks) 

1294 

1295 # We need UPSERT which is dialect-specific construct 

1296 values = {"last_update_time": dt, "unique_id": replica_chunk.unique_id} 

1297 row = {"apdb_replica_chunk": replica_chunk.id} | values 

1298 if connection.dialect.name == "sqlite": 1298 ↛ 1302line 1298 didn't jump to line 1302 because the condition on line 1298 was always true

1299 insert_sqlite = sqlalchemy.dialects.sqlite.insert(table) 

1300 insert_sqlite = insert_sqlite.on_conflict_do_update(index_elements=table.primary_key, set_=values) 

1301 connection.execute(insert_sqlite, row) 

1302 elif connection.dialect.name == "postgresql": 

1303 insert_pg = sqlalchemy.dialects.postgresql.dml.insert(table) 

1304 insert_pg = insert_pg.on_conflict_do_update(constraint=table.primary_key, set_=values) 

1305 connection.execute(insert_pg, row) 

1306 else: 

1307 raise TypeError(f"Unsupported dialect {connection.dialect.name} for upsert.") 

1308 

1309 def _storeDiaObjects( 

1310 self, 

1311 objs: pandas.DataFrame, 

1312 visit_time: astropy.time.Time, 

1313 replica_chunk: ReplicaChunk | None, 

1314 connection: sqlalchemy.engine.Connection, 

1315 ) -> None: 

1316 """Store catalog of DiaObjects from current visit. 

1317 

1318 Parameters 

1319 ---------- 

1320 objs : `pandas.DataFrame` 

1321 Catalog with DiaObject records. 

1322 visit_time : `astropy.time.Time` 

1323 Time of the visit. 

1324 replica_chunk : `ReplicaChunk` 

1325 Insert identifier. 

1326 """ 

1327 if len(objs) == 0: 

1328 _LOG.debug("No objects to write to database.") 

1329 return 

1330 

1331 # Some types like np.int64 can cause issues with sqlalchemy, convert 

1332 # them to int. 

1333 ids = sorted(int(oid) for oid in objs["diaObjectId"]) 

1334 _LOG.debug("first object ID: %d", ids[0]) 

1335 

1336 validity_start_column = self._timestamp_column_name("validityStart") 

1337 validity_end_column = self._timestamp_column_name("validityEnd") 

1338 timestamp = self._timestamp_column_value(visit_time) 

1339 

1340 # everything to be done in single transaction 

1341 if self.config.dia_object_index == "last_object_table": 

1342 # Insert and replace all records in LAST table. 

1343 table = self._schema.get_table(ApdbTables.DiaObjectLast) 

1344 

1345 # DiaObjectLast did not have this column in the past. 

1346 use_validity_start = self._schema.check_column(ApdbTables.DiaObjectLast, validity_start_column) 

1347 

1348 # Drop the previous objects (pandas cannot upsert). 

1349 query = table.delete().where(table.columns["diaObjectId"].in_(ids)) 

1350 

1351 with self._timer("delete_time", tags={"table": table.name}) as timer: 

1352 res = connection.execute(query) 

1353 timer.add_values(row_count=res.rowcount) 

1354 _LOG.debug("deleted %s objects", res.rowcount) 

1355 

1356 # DiaObjectLast is a subset of DiaObject, strip missing columns 

1357 last_column_names = [column.name for column in table.columns] 

1358 if validity_start_column in last_column_names and validity_start_column not in objs.columns: 1358 ↛ 1360line 1358 didn't jump to line 1360 because the condition on line 1358 was always true

1359 last_column_names.remove(validity_start_column) 

1360 last_objs = objs[last_column_names] 

1361 last_objs = _coerce_uint64(last_objs) 

1362 

1363 # Fill validityStart, only when it is in the schema. 

1364 if use_validity_start: 1364 ↛ 1372line 1364 didn't jump to line 1372 because the condition on line 1364 was always true

1365 if validity_start_column in last_objs: 1365 ↛ 1366line 1365 didn't jump to line 1366 because the condition on line 1365 was never true

1366 last_objs[validity_start_column] = timestamp 

1367 else: 

1368 extra_column = pandas.Series([timestamp] * len(last_objs), name=validity_start_column) 

1369 last_objs.set_index(extra_column.index, inplace=True) 

1370 last_objs = pandas.concat([last_objs, extra_column], axis="columns") 

1371 

1372 with self._timer("insert_time", tags={"table": "DiaObjectLast"}) as timer: 

1373 last_objs.to_sql( 

1374 table.name, 

1375 connection, 

1376 if_exists="append", 

1377 index=False, 

1378 schema=table.schema, 

1379 ) 

1380 timer.add_values(row_count=len(last_objs)) 

1381 

1382 # truncate existing validity intervals 

1383 table = self._schema.get_table(ApdbTables.DiaObject) 

1384 

1385 update = ( 

1386 table.update() 

1387 .values(**{validity_end_column: timestamp}) 

1388 .where( 

1389 sql.expression.and_( 

1390 table.columns["diaObjectId"].in_(ids), 

1391 table.columns[validity_end_column].is_(None), 

1392 ) 

1393 ) 

1394 ) 

1395 

1396 with self._timer("truncate_time", tags={"table": table.name}) as timer: 

1397 res = connection.execute(update) 

1398 timer.add_values(row_count=res.rowcount) 

1399 _LOG.debug("truncated %s intervals", res.rowcount) 

1400 

1401 objs = _coerce_uint64(objs) 

1402 

1403 # Fill additional columns 

1404 extra_columns: list[pandas.Series] = [] 

1405 if validity_start_column in objs.columns: 1405 ↛ 1406line 1405 didn't jump to line 1406 because the condition on line 1405 was never true

1406 objs[validity_start_column] = timestamp 

1407 else: 

1408 extra_columns.append(pandas.Series([timestamp] * len(objs), name=validity_start_column)) 

1409 if validity_end_column in objs.columns: 1409 ↛ 1410line 1409 didn't jump to line 1410 because the condition on line 1409 was never true

1410 objs[validity_end_column] = None 

1411 else: 

1412 extra_columns.append(pandas.Series([None] * len(objs), name=validity_end_column)) 

1413 if extra_columns: 1413 ↛ 1418line 1413 didn't jump to line 1418 because the condition on line 1413 was always true

1414 objs.set_index(extra_columns[0].index, inplace=True) 

1415 objs = pandas.concat([objs] + extra_columns, axis="columns") 

1416 

1417 # Insert replica data 

1418 table = self._schema.get_table(ApdbTables.DiaObject) 

1419 replica_data: list[dict] = [] 

1420 replica_stmt: Any = None 

1421 replica_table_name = "" 

1422 if replica_chunk is not None: 

1423 pk_names = [column.name for column in table.primary_key] 

1424 replica_data = objs[pk_names].to_dict("records") 

1425 if replica_data: 1425 ↛ 1433line 1425 didn't jump to line 1433 because the condition on line 1425 was always true

1426 for row in replica_data: 

1427 row["apdb_replica_chunk"] = replica_chunk.id 

1428 replica_table = self._schema.get_table(ExtraTables.DiaObjectChunks) 

1429 replica_table_name = replica_table.name 

1430 replica_stmt = replica_table.insert() 

1431 

1432 # insert new versions 

1433 with self._timer("insert_time", tags={"table": table.name}) as timer: 

1434 objs.to_sql(table.name, connection, if_exists="append", index=False, schema=table.schema) 

1435 timer.add_values(row_count=len(objs)) 

1436 if replica_stmt is not None: 

1437 with self._timer("insert_time", tags={"table": replica_table_name}) as timer: 

1438 connection.execute(replica_stmt, replica_data) 

1439 timer.add_values(row_count=len(replica_data)) 

1440 

1441 def _storeDiaSources( 

1442 self, 

1443 sources: pandas.DataFrame, 

1444 replica_chunk: ReplicaChunk | None, 

1445 connection: sqlalchemy.engine.Connection, 

1446 ) -> None: 

1447 """Store catalog of DiaSources from current visit. 

1448 

1449 Parameters 

1450 ---------- 

1451 sources : `pandas.DataFrame` 

1452 Catalog containing DiaSource records 

1453 """ 

1454 table = self._schema.get_table(ApdbTables.DiaSource) 

1455 

1456 # Insert replica data 

1457 replica_data: list[dict] = [] 

1458 replica_stmt: Any = None 

1459 replica_table_name = "" 

1460 if replica_chunk is not None: 

1461 pk_names = [column.name for column in table.primary_key] 

1462 replica_data = sources[pk_names].to_dict("records") 

1463 if replica_data: 1463 ↛ 1471line 1463 didn't jump to line 1471 because the condition on line 1463 was always true

1464 for row in replica_data: 

1465 row["apdb_replica_chunk"] = replica_chunk.id 

1466 replica_table = self._schema.get_table(ExtraTables.DiaSourceChunks) 

1467 replica_table_name = replica_table.name 

1468 replica_stmt = replica_table.insert() 

1469 

1470 # everything to be done in single transaction 

1471 with self._timer("insert_time", tags={"table": table.name}) as timer: 

1472 sources = _coerce_uint64(sources) 

1473 sources.to_sql(table.name, connection, if_exists="append", index=False, schema=table.schema) 

1474 timer.add_values(row_count=len(sources)) 

1475 if replica_stmt is not None: 

1476 with self._timer("replica_insert_time", tags={"table": replica_table_name}) as timer: 

1477 connection.execute(replica_stmt, replica_data) 

1478 timer.add_values(row_count=len(replica_data)) 

1479 

1480 def _storeDiaForcedSources( 

1481 self, 

1482 sources: pandas.DataFrame, 

1483 replica_chunk: ReplicaChunk | None, 

1484 connection: sqlalchemy.engine.Connection, 

1485 ) -> None: 

1486 """Store a set of DiaForcedSources from current visit. 

1487 

1488 Parameters 

1489 ---------- 

1490 sources : `pandas.DataFrame` 

1491 Catalog containing DiaForcedSource records 

1492 """ 

1493 table = self._schema.get_table(ApdbTables.DiaForcedSource) 

1494 

1495 # Insert replica data 

1496 replica_data: list[dict] = [] 

1497 replica_stmt: Any = None 

1498 replica_table_name = "" 

1499 if replica_chunk is not None: 

1500 pk_names = [column.name for column in table.primary_key] 

1501 replica_data = sources[pk_names].to_dict("records") 

1502 if replica_data: 1502 ↛ 1510line 1502 didn't jump to line 1510 because the condition on line 1502 was always true

1503 for row in replica_data: 

1504 row["apdb_replica_chunk"] = replica_chunk.id 

1505 replica_table = self._schema.get_table(ExtraTables.DiaForcedSourceChunks) 

1506 replica_table_name = replica_table.name 

1507 replica_stmt = replica_table.insert() 

1508 

1509 # everything to be done in single transaction 

1510 with self._timer("insert_time", tags={"table": table.name}) as timer: 

1511 sources = _coerce_uint64(sources) 

1512 sources.to_sql(table.name, connection, if_exists="append", index=False, schema=table.schema) 

1513 timer.add_values(row_count=len(sources)) 

1514 if replica_stmt is not None: 

1515 with self._timer("insert_time", tags={"table": replica_table_name}) as timer: 

1516 connection.execute(replica_stmt, replica_data) 

1517 timer.add_values(row_count=len(replica_data)) 

1518 

1519 def _storeUpdateRecords( 

1520 self, 

1521 records: Iterable[ApdbUpdateRecord], 

1522 chunk: ReplicaChunk, 

1523 *, 

1524 store_chunk: bool = False, 

1525 connection: sqlalchemy.engine.Connection | None = None, 

1526 ) -> None: 

1527 """Store ApdbUpdateRecords in the replica table for those records. 

1528 

1529 Parameters 

1530 ---------- 

1531 records : `list` [`ApdbUpdateRecord`] 

1532 Records to store. 

1533 chunk : `ReplicaChunk` 

1534 Replica chunk for these records. 

1535 store_chunk : `bool` 

1536 If True then also store replica chunk. 

1537 connection : `sqlalchemy.engine.Connection` 

1538 SQLALchemy connection to use, if `None` the new connection will be 

1539 made. `None` is useful for tests only, regular use will call this 

1540 method in the same transaction that saves other types of records. 

1541 

1542 Raises 

1543 ------ 

1544 TypeError 

1545 Raised if replication is not enabled for this instance. 

1546 """ 

1547 if not self._schema.replication_enabled: 

1548 raise TypeError("Replication is not enabled for this APDB instance.") 

1549 

1550 apdb_replica_chunk = chunk.id 

1551 # Do not use unique_if from ReplicaChunk as it could be reused in 

1552 # multiple calls to this method. 

1553 update_unique_id = uuid.uuid4() 

1554 

1555 record_dicts = [] 

1556 for record in records: 

1557 record_dicts.append( 

1558 { 

1559 "apdb_replica_chunk": apdb_replica_chunk, 

1560 "update_time_ns": record.update_time_ns, 

1561 "update_order": record.update_order, 

1562 "update_unique_id": update_unique_id, 

1563 "update_payload": record.to_json(), 

1564 } 

1565 ) 

1566 

1567 if not record_dicts: 1567 ↛ 1568line 1567 didn't jump to line 1568 because the condition on line 1567 was never true

1568 return 

1569 

1570 # TODO: Need to check that table exists. 

1571 table = self._schema.get_table(ExtraTables.ApdbUpdateRecordChunks) 

1572 

1573 def _do_store(connection: sqlalchemy.engine.Connection) -> None: 

1574 if store_chunk: 1574 ↛ 1576line 1574 didn't jump to line 1576 because the condition on line 1574 was always true

1575 self._storeReplicaChunk(chunk, connection) 

1576 with self._timer("insert_time", tags={"table": table.name}) as timer: 

1577 connection.execute(table.insert(), record_dicts) 

1578 timer.add_values(row_count=len(record_dicts)) 

1579 

1580 if connection is None: 

1581 with self._engine.begin() as connection: 

1582 _do_store(connection) 

1583 else: 

1584 _do_store(connection) 

1585 

1586 def _htm_indices(self, region: Region) -> list[tuple[int, int]]: 

1587 """Generate a set of HTM indices covering specified region. 

1588 

1589 Parameters 

1590 ---------- 

1591 region: `sphgeom.Region` 

1592 Region that needs to be indexed. 

1593 

1594 Returns 

1595 ------- 

1596 Sequence of ranges, range is a tuple (minHtmID, maxHtmID). 

1597 """ 

1598 _LOG.debug("region: %s", region) 

1599 indices = self.pixelator.envelope(region, self.config.pixelization.htm_max_ranges) 

1600 

1601 return indices.ranges() 

1602 

1603 def _filterRegion(self, table: sqlalchemy.schema.Table, region: Region) -> sql.ColumnElement: 

1604 """Make SQLAlchemy expression for selecting records in a region.""" 

1605 htm_index_column = table.columns[self.config.pixelization.htm_index_column] 

1606 exprlist = [] 

1607 pixel_ranges = self._htm_indices(region) 

1608 for low, upper in pixel_ranges: 

1609 upper -= 1 

1610 if low == upper: 1610 ↛ 1611line 1610 didn't jump to line 1611 because the condition on line 1610 was never true

1611 exprlist.append(htm_index_column == low) 

1612 else: 

1613 exprlist.append(sql.expression.between(htm_index_column, low, upper)) 

1614 

1615 return sql.expression.or_(*exprlist) 

1616 

1617 def _add_spatial_index(self, df: pandas.DataFrame) -> pandas.DataFrame: 

1618 """Calculate spatial index for each record and add it to a DataFrame. 

1619 

1620 Parameters 

1621 ---------- 

1622 df : `pandas.DataFrame` 

1623 DataFrame which has to contain ra/dec columns, names of these 

1624 columns are defined by configuration ``ra_dec_columns`` field. 

1625 

1626 Returns 

1627 ------- 

1628 df : `pandas.DataFrame` 

1629 DataFrame with ``pixelId`` column which contains pixel index 

1630 for ra/dec coordinates. 

1631 

1632 Notes 

1633 ----- 

1634 This overrides any existing column in a DataFrame with the same name 

1635 (pixelId). Original DataFrame is not changed, copy of a DataFrame is 

1636 returned. 

1637 """ 

1638 # calculate HTM index for every DiaObject 

1639 htm_index = np.zeros(df.shape[0], dtype=np.int64) 

1640 ra_col, dec_col = self.config.ra_dec_columns 

1641 for i, (ra, dec) in enumerate(zip(df[ra_col], df[dec_col])): 

1642 uv3d = UnitVector3d(LonLat.fromDegrees(ra, dec)) 

1643 idx = self.pixelator.index(uv3d) 

1644 htm_index[i] = idx 

1645 df = df.copy() 

1646 df[self.config.pixelization.htm_index_column] = htm_index 

1647 return df 

1648 

1649 def _fix_input_timestamps(self, df: pandas.DataFrame) -> pandas.DataFrame: 

1650 """Update timestamp columns in input DataFrame to be aware datetime 

1651 type in in UTC. 

1652 

1653 AP pipeline generates naive datetime instances, we want them to be 

1654 aware before they go to database. All naive timestamps are assumed to 

1655 be in UTC timezone (they should be TAI). 

1656 """ 

1657 # Find all columns with aware non-UTC timestamps and convert to UTC. 

1658 columns = [ 

1659 column 

1660 for column, dtype in df.dtypes.items() 

1661 if isinstance(dtype, pandas.DatetimeTZDtype) and dtype.tz is not datetime.UTC 

1662 ] 

1663 for column in columns: 1663 ↛ 1664line 1663 didn't jump to line 1664 because the loop on line 1663 never started

1664 df[column] = df[column].dt.tz_convert(datetime.UTC) 

1665 # Find all columns with naive timestamps and add UTC timezone. 

1666 columns = [ 

1667 column for column, dtype in df.dtypes.items() if pandas.api.types.is_datetime64_dtype(dtype) 

1668 ] 

1669 for column in columns: 

1670 df[column] = df[column].dt.tz_localize(datetime.UTC) 

1671 return df 

1672 

1673 def _fix_result_timestamps(self, df: pandas.DataFrame) -> pandas.DataFrame: 

1674 """Update timestamp columns to be naive datetime type in returned 

1675 DataFrame. 

1676 

1677 AP pipeline code expects DataFrames to contain naive datetime columns, 

1678 while Postgres queries return timezone-aware type. This method converts 

1679 those columns to naive datetime in UTC timezone. 

1680 """ 

1681 # Find all columns with aware timestamps. 

1682 columns = [column for column, dtype in df.dtypes.items() if isinstance(dtype, pandas.DatetimeTZDtype)] 

1683 for column in columns: 1683 ↛ 1685line 1683 didn't jump to line 1685 because the loop on line 1683 never started

1684 # tz_convert(None) will convert to UTC and drop timezone. 

1685 df[column] = df[column].dt.tz_convert(None) 

1686 return df 

1687 

1688 def _timestamp_column_name(self, column: str) -> str: 

1689 """Return column name before/after schema migration to MJD TAI.""" 

1690 return self.schema.timestamp_column_name(column) 

1691 

1692 def _timestamp_column_value(self, time: astropy.time.Time) -> float | datetime.datetime: 

1693 """Return column value before/after schema migration to MJD TAI.""" 

1694 if self.schema.has_mjd_timestamps: 

1695 return float(time.tai.mjd) 

1696 else: 

1697 return time.datetime.astimezone(tz=datetime.UTC) 

1698 

1699 def _get_diaobject_data( 

1700 self, conn: sqlalchemy.engine.Connection, object_ids: Iterable[int], *columns: str 

1701 ) -> list: 

1702 """Select records from either DiaObject or DiaObjectLast and return 

1703 selected rows as names tuples. 

1704 """ 

1705 where: sqlalchemy.ColumnElement[bool] 

1706 if self.config.dia_object_index == "last_object_table": 

1707 table = self._schema.get_table(ApdbTables.DiaObjectLast) 

1708 where = table.columns["diaObjectId"].in_(sorted(object_ids)) 

1709 else: 

1710 table = self._schema.get_table(ApdbTables.DiaObject) 

1711 validity_end_column = self._timestamp_column_name("validityEnd") 

1712 where = sqlalchemy.and_( 

1713 table.columns["diaObjectId"].in_(sorted(object_ids)), 

1714 table.columns[validity_end_column].is_(None), 

1715 ) 

1716 column_list = [table.columns["diaObjectId"]] + [table.columns[column] for column in columns] 

1717 query = sql.select(*column_list).where(where) 

1718 result = conn.execute(query) 

1719 

1720 return list(result) 

1721 

1722 def _get_diasource_data( 

1723 self, conn: sqlalchemy.engine.Connection, source_ids: Iterable[int], *columns: str 

1724 ) -> list: 

1725 """Select records from DiaSource table by diaSourceId and return 

1726 selected rows as named tuples. 

1727 """ 

1728 table = self._schema.get_table(ApdbTables.DiaSource) 

1729 where = table.columns["diaSourceId"].in_(sorted(source_ids)) 

1730 column_list = [table.columns["diaSourceId"]] + [table.columns[column] for column in columns] 

1731 query = sql.select(*column_list).where(where) 

1732 result = conn.execute(query) 

1733 

1734 return list(result)