Coverage for python/lsst/dax/apdb/sql/apdbSql.py: 88%
756 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-29 02:07 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-29 02:07 -0700
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/>.
22"""Module defining Apdb class and related methods."""
24from __future__ import annotations
26__all__ = ["ApdbSql"]
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
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
48from lsst.sphgeom import HtmPixelization, LonLat, Region, UnitVector3d
49from lsst.utils.db_auth import DbAuth, DbAuthNotFoundError
50from lsst.utils.iteration import chunk_iterable
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
76if TYPE_CHECKING:
77 import sqlite3
79 from ..apdbMetadata import ApdbMetadata
80 from ..apdbUpdateRecord import ApdbUpdateRecord
82_LOG = logging.getLogger(__name__)
84_MON = MonAgent(__name__)
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"""
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))
101def _make_midpointMjdTai_start(visit_time: astropy.time.Time, months: int) -> float:
102 """Calculate starting point for time-based source search.
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.
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
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;")
129class ApdbSql(Apdb):
130 """Implementation of APDB interface based on SQL database.
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.
136 Parameters
137 ----------
138 config : `ApdbSqlConfig`
139 Configuration object.
140 """
142 metadataSchemaVersionKey = "version:schema"
143 """Name of the metadata key to store schema version number."""
145 metadataCodeVersionKey = "version:ApdbSql"
146 """Name of the metadata key to store code version number."""
148 metadataReplicaVersionKey = "version:ApdbSqlReplica"
149 """Name of the metadata key to store replica code version number."""
151 metadataConfigKey = "config:apdb-sql.json"
152 """Name of the metadata key to store code version number."""
154 metadataDedupKey = "status:deduplication.json"
155 """Name of the metadata key to store code version number."""
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."""
166 def __init__(self, config: ApdbSqlConfig):
167 self._engine = self._makeEngine(config, create=False)
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)
174 # Get tables schemas.
175 self._table_schema = ApdbSchema(config.schema_file, config.ss_schema_file)
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())
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
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 )
200 self.pixelator = HtmPixelization(self.config.pixelization.htm_level)
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())
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)
209 @classmethod
210 def _makeEngine(cls, config: ApdbSqlConfig, *, create: bool) -> sqlalchemy.engine.Engine:
211 """Make SQLALchemy engine based on configured parameters.
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)
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)
244 return engine
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.
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.
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
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)
281 return config_url
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.
288 Parameters
289 ----------
290 url_string : `str`
291 Connection string.
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
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()
328 return url_string
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 """
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)
344 db_schema_version = _get_version(cls.metadataSchemaVersionKey)
345 db_code_version = _get_version(cls.metadataCodeVersionKey)
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 )
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 )
374 return db_schema_version
376 @classmethod
377 def apdbImplementationVersion(cls) -> VersionTuple:
378 """Return version number for current APDB implementation.
380 Returns
381 -------
382 version : `VersionTuple`
383 Version of the code defined in implementation class.
384 """
385 return VERSION
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.
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.
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
471 cls._makeSchema(config, drop=drop)
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)
477 return config
479 def get_replica(self) -> ApdbSqlReplica:
480 """Return `ApdbReplica` instance for this database."""
481 return ApdbSqlReplica(self._schema, self._engine, self._db_schema_version)
483 def tableRowCount(self) -> dict[str, int]:
484 """Return dictionary with the table names and row counts.
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.
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
505 return res
507 def getConfig(self) -> ApdbSqlConfig:
508 # docstring is inherited from a base class
509 return self.config
511 def tableDef(self, table: ApdbTables) -> Table | None:
512 # docstring is inherited from a base class
513 return self._schema.tableSchemas.get(table)
515 @classmethod
516 def _makeSchema(cls, config: ApdbConfig, drop: bool = False) -> None:
517 # docstring is inherited from a base class
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)}")
522 engine = cls._makeEngine(config, create=True)
524 table_schema = ApdbSchema(config.schema_file, config.ss_schema_file)
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)
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)
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 )
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)
557 def getDiaObjects(self, region: Region) -> pandas.DataFrame:
558 # docstring is inherited from a base class
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)
572 # build selection
573 query = query.where(self._filterRegion(table, region))
575 validity_end_column = self._timestamp_column_name("validityEnd")
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
581 # _LOG.debug("query: %s", query)
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)
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
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)
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)
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
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")
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)
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))
649 _LOG.debug("found %s DiaForcedSources", len(sources))
650 return sources
652 def getDiaObjectsForDedup(self, since: astropy.time.Time | None = None) -> pandas.DataFrame:
653 # docstring is inherited from a base class
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")
663 validity_start_column = self._timestamp_column_name("validityStart")
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)
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]
682 query = sql.select(*columns)
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)
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)
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))
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)
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
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)
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)
751 # fill pixelId column for DiaObjects
752 objects = self._add_spatial_index(objects)
753 self._storeDiaObjects(objects, visit_time, replica_chunk, connection)
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)
760 if forced_sources is not None:
761 self._storeDiaForcedSources(forced_sources, replica_chunk, connection)
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
772 new_object_ids = set(idMap.values())
773 source_ids = {source.diaSourceId for source in idMap}
775 current_time = self._current_time()
776 current_time_ns = int(current_time.unix_tai * 1e9)
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}
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}")
793 found_objects_by_id = {row.diaObjectId: row for row in found_objects}
795 update_records: list[ApdbUpdateRecord] = []
796 update_order = 0
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)
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
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))
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)
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)
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
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)
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
871 with self._engine.begin() as conn:
872 return self._setValidityEnd(conn, objects, validityEnd, raise_on_missing_id=raise_on_missing_id)
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
885 requested_ids = {obj.diaObjectId for obj in objects}
887 validity_end_column = self._timestamp_column_name("validityEnd")
888 validityEnd_value = self._timestamp_column_value(validityEnd)
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 )
899 result = conn.execute(query)
900 found_ids = set(result.scalars())
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}")
907 # Filter existing records.
908 if len(objects) != len(found_ids):
909 objects = [obj for obj in objects if obj.diaObjectId in found_ids]
911 if not objects:
912 return 0
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 )
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 )
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)
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 ]
960 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True, connection=conn)
962 return len(objects)
964 def resetDedup(self, dedup_time: astropy.time.Time | None = None) -> None:
965 # docstring is inherited from a base class
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)
975 def reassignDiaSources(self, idMap: Mapping[int, int]) -> None:
976 # docstring is inherited from a base class
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)
983 table = self._schema.get_table(ApdbTables.DiaSource)
984 query = table.update().where(table.columns["diaSourceId"] == sql.bindparam("srcId"))
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}")
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
1013 diaSourceIds = list(diaSourceIds)
1014 source_ids = {source.diaSourceId for source in diaSourceIds}
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")
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}")
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]
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)
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 )
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)
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
1068 def _fsrc_id(fsource: Any) -> tuple[int, int, int]:
1069 return (fsource.diaObjectId, fsource.visit, fsource.detector)
1071 diaForcedSourceIds = list(diaForcedSourceIds)
1072 source_ids = {_fsrc_id(source) for source in diaForcedSourceIds}
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")
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}")
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]
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)
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 )
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 ]
1120 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True, connection=conn)
1122 def countUnassociatedObjects(self) -> int:
1123 # docstring is inherited from a base class
1125 # Retrieve the DiaObject table.
1126 table: sqlalchemy.schema.Table = self._schema.get_table(ApdbTables.DiaObject)
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 )
1141 # Return the count.
1142 with self._engine.begin() as conn:
1143 count = conn.execute(stmt).scalar_one()
1145 return count
1147 @property
1148 def schema(self) -> ApdbSchema:
1149 # docstring is inherited from a base class
1150 return self._table_schema
1152 @property
1153 def metadata(self) -> ApdbMetadata:
1154 # docstring is inherited from a base class
1155 return self._metadata
1157 @property
1158 def admin(self) -> ApdbSqlAdmin:
1159 # docstring is inherited from a base class
1160 return ApdbSqlAdmin(self.pixelator)
1162 def _getDiaSourcesInRegion(self, region: Region, start_time_mjdTai: float) -> pandas.DataFrame:
1163 """Return catalog of DiaSource instances from given region.
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.
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)
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)
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)
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.
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.
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))
1216 _LOG.debug("found %s DiaSources", len(sources))
1217 return sources
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.
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.
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
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)
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]
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 )
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())
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)
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)
1293 table = self._schema.get_table(ExtraTables.ApdbReplicaChunks)
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.")
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.
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
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])
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)
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)
1345 # DiaObjectLast did not have this column in the past.
1346 use_validity_start = self._schema.check_column(ApdbTables.DiaObjectLast, validity_start_column)
1348 # Drop the previous objects (pandas cannot upsert).
1349 query = table.delete().where(table.columns["diaObjectId"].in_(ids))
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)
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)
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")
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))
1382 # truncate existing validity intervals
1383 table = self._schema.get_table(ApdbTables.DiaObject)
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 )
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)
1401 objs = _coerce_uint64(objs)
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")
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()
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))
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.
1449 Parameters
1450 ----------
1451 sources : `pandas.DataFrame`
1452 Catalog containing DiaSource records
1453 """
1454 table = self._schema.get_table(ApdbTables.DiaSource)
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()
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))
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.
1488 Parameters
1489 ----------
1490 sources : `pandas.DataFrame`
1491 Catalog containing DiaForcedSource records
1492 """
1493 table = self._schema.get_table(ApdbTables.DiaForcedSource)
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()
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))
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.
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.
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.")
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()
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 )
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
1570 # TODO: Need to check that table exists.
1571 table = self._schema.get_table(ExtraTables.ApdbUpdateRecordChunks)
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))
1580 if connection is None:
1581 with self._engine.begin() as connection:
1582 _do_store(connection)
1583 else:
1584 _do_store(connection)
1586 def _htm_indices(self, region: Region) -> list[tuple[int, int]]:
1587 """Generate a set of HTM indices covering specified region.
1589 Parameters
1590 ----------
1591 region: `sphgeom.Region`
1592 Region that needs to be indexed.
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)
1601 return indices.ranges()
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))
1615 return sql.expression.or_(*exprlist)
1617 def _add_spatial_index(self, df: pandas.DataFrame) -> pandas.DataFrame:
1618 """Calculate spatial index for each record and add it to a DataFrame.
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.
1626 Returns
1627 -------
1628 df : `pandas.DataFrame`
1629 DataFrame with ``pixelId`` column which contains pixel index
1630 for ra/dec coordinates.
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
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.
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
1673 def _fix_result_timestamps(self, df: pandas.DataFrame) -> pandas.DataFrame:
1674 """Update timestamp columns to be naive datetime type in returned
1675 DataFrame.
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
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)
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)
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)
1720 return list(result)
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)
1734 return list(result)