Coverage for python/lsst/dax/apdb/cassandra/apdbCassandra.py: 8%
912 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-29 09:48 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-29 09:48 +0000
1# This file is part of dax_apdb.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (http://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <http://www.gnu.org/licenses/>.
22from __future__ import annotations
24__all__ = ["ApdbCassandra"]
26import datetime
27import json
28import logging
29import random
30import uuid
31import warnings
32from collections import Counter, defaultdict
33from collections.abc import Iterable, Mapping, Set
34from typing import TYPE_CHECKING, Any, cast
36import numpy as np
37import pandas
39# If cassandra-driver is not there the module can still be imported
40# but ApdbCassandra cannot be instantiated.
41try:
42 import cassandra
43 import cassandra.query
44 from cassandra.query import UNSET_VALUE
46 CASSANDRA_IMPORTED = True
47except ImportError:
48 CASSANDRA_IMPORTED = False
50import astropy.time
51import felis.datamodel
53from lsst import sphgeom
54from lsst.utils.iteration import chunk_iterable
56from ..apdb import Apdb, ApdbConfig
57from ..apdbConfigFreezer import ApdbConfigFreezer
58from ..apdbReplica import ApdbTableData, ReplicaChunk
59from ..apdbSchema import ApdbSchema, ApdbTables
60from ..apdbUpdateRecord import (
61 ApdbCloseDiaObjectValidityRecord,
62 ApdbReassignDiaSourceToDiaObjectRecord,
63 ApdbUpdateNDiaSourcesRecord,
64 ApdbWithdrawDiaForcedSourceRecord,
65 ApdbWithdrawDiaSourceRecord,
66)
67from ..monitor import MonAgent
68from ..recordIds import DiaForcedSourceId, DiaObjectId, DiaSourceId
69from ..schema_model import Table
70from ..timer import Timer
71from ..versionTuple import VersionTuple
72from .apdbCassandraAdmin import ApdbCassandraAdmin
73from .apdbCassandraReplica import ApdbCassandraReplica
74from .apdbCassandraSchema import ApdbCassandraSchema, CreateTableOptions, ExtraTables
75from .apdbMetadataCassandra import ApdbMetadataCassandra
76from .cassandra_utils import (
77 ApdbCassandraTableData,
78 execute_concurrent,
79 literal,
80 select_concurrent,
81)
82from .config import ApdbCassandraConfig, ApdbCassandraConnectionConfig, ApdbCassandraTimePartitionRange
83from .connectionContext import ConnectionContext, DbVersions
84from .exceptions import CassandraMissingError
85from .partitioner import Partitioner
86from .queries import Column as C # noqa: N817
87from .queries import ColumnExpr, Delete, Insert, QExpr, Select, Update
88from .sessionFactory import SessionContext, SessionFactory
90if TYPE_CHECKING:
91 from ..apdbMetadata import ApdbMetadata
92 from ..apdbUpdateRecord import ApdbUpdateRecord
94_LOG = logging.getLogger(__name__)
96_MON = MonAgent(__name__)
98VERSION = VersionTuple(1, 3, 0)
99"""Version for the code controlling non-replication tables. This needs to be
100updated following compatibility rules when schema produced by this code
101changes.
102"""
105class ApdbCassandra(Apdb):
106 """Implementation of APDB database with Apache Cassandra backend.
108 Parameters
109 ----------
110 config : `ApdbCassandraConfig`
111 Configuration object.
112 """
114 def __init__(self, config: ApdbCassandraConfig):
115 if not CASSANDRA_IMPORTED:
116 raise CassandraMissingError()
118 self._config = config
119 self._keyspace = config.keyspace
120 self._schema = ApdbSchema(config.schema_file, config.ss_schema_file)
122 self._session_factory = SessionFactory(config)
123 self._connection_context: ConnectionContext | None = None
125 @property
126 def _context(self) -> ConnectionContext:
127 """Establish connection if not established and return context."""
128 if self._connection_context is None:
129 current_versions = DbVersions(
130 schema_version=self.schema.schemaVersion(),
131 code_version=self.apdbImplementationVersion(),
132 replica_version=ApdbCassandraReplica.apdbReplicaImplementationVersion(),
133 )
134 _LOG.debug("Current versions: %s", current_versions)
136 session = self._session_factory.session()
137 self._connection_context = ConnectionContext(
138 session, self._config, self.schema.tableSchemas, current_versions
139 )
141 if _LOG.isEnabledFor(logging.DEBUG):
142 _LOG.debug("ApdbCassandra Configuration: %s", self._connection_context.config.model_dump())
144 return self._connection_context
146 def _timer(self, name: str, *, tags: Mapping[str, str | int] | None = None) -> Timer:
147 """Create `Timer` instance given its name."""
148 return Timer(name, _MON, tags=tags)
150 @classmethod
151 def apdbImplementationVersion(cls) -> VersionTuple:
152 """Return version number for current APDB implementation.
154 Returns
155 -------
156 version : `VersionTuple`
157 Version of the code defined in implementation class.
158 """
159 return VERSION
161 def getConfig(self) -> ApdbCassandraConfig:
162 # docstring is inherited from a base class
163 return self._context.config
165 def tableDef(self, table: ApdbTables) -> Table | None:
166 # docstring is inherited from a base class
167 return self.schema.tableSchemas.get(table)
169 @classmethod
170 def init_database(
171 cls,
172 hosts: tuple[str, ...],
173 keyspace: str,
174 *,
175 schema_file: str | None = None,
176 ss_schema_file: str | None = None,
177 read_sources_months: int | None = None,
178 read_forced_sources_months: int | None = None,
179 enable_replica: bool = False,
180 replica_skips_diaobjects: bool = False,
181 port: int | None = None,
182 username: str | None = None,
183 dbauth_alias: str | None = None,
184 prefix: str | None = None,
185 part_pixelization: str | None = None,
186 part_pix_level: int | None = None,
187 time_partition_tables: bool = True,
188 time_partition_start: str | None = None,
189 time_partition_end: str | None = None,
190 read_consistency: str | None = None,
191 write_consistency: str | None = None,
192 read_timeout: int | None = None,
193 write_timeout: int | None = None,
194 ra_dec_columns: tuple[str, str] | None = None,
195 replication_factor: int | None = None,
196 drop: bool = False,
197 table_options: CreateTableOptions | None = None,
198 ) -> ApdbCassandraConfig:
199 """Initialize new APDB instance and make configuration object for it.
201 Parameters
202 ----------
203 hosts : `tuple` [`str`, ...]
204 List of host names or IP addresses for Cassandra cluster.
205 keyspace : `str`
206 Name of the keyspace for APDB tables.
207 schema_file : `str`, optional
208 Location of (YAML) configuration file with APDB schema. If not
209 specified then default location will be used.
210 ss_schema_file : `str`, optional
211 Location of (YAML) configuration file with SSO schema. If not
212 specified then default location will be used.
213 read_sources_months : `int`, optional
214 Number of months of history to read from DiaSource.
215 read_forced_sources_months : `int`, optional
216 Number of months of history to read from DiaForcedSource.
217 enable_replica : `bool`, optional
218 If True, make additional tables used for replication to PPDB.
219 replica_skips_diaobjects : `bool`, optional
220 If `True` then do not fill regular ``DiaObject`` table when
221 ``enable_replica`` is `True`.
222 port : `int`, optional
223 Port number to use for Cassandra connections.
224 username : `str`, optional
225 User name for Cassandra connections.
226 dbauth_alias : `str`, optional
227 If specified then this string will be used to as a host name when
228 checking credentials in db-auth.yaml in addition to regular host
229 names in contact_points. For example if
230 dbauth_alias='pp_apdb_prod_cluster' then the entry
231 'cassandra://pp_apdb_prod_cluster/' will match. Port number should
232 not be used in that entry. Alias has higher priority than host
233 names.
234 prefix : `str`, optional
235 Optional prefix for all table names.
236 part_pixelization : `str`, optional
237 Name of the MOC pixelization used for partitioning.
238 part_pix_level : `int`, optional
239 Pixelization level.
240 time_partition_tables : `bool`, optional
241 Create per-partition tables.
242 time_partition_start : `str`, optional
243 Starting time for per-partition tables, in yyyy-mm-ddThh:mm:ss
244 format, in TAI.
245 time_partition_end : `str`, optional
246 Ending time for per-partition tables, in yyyy-mm-ddThh:mm:ss
247 format, in TAI.
248 read_consistency : `str`, optional
249 Name of the consistency level for read operations.
250 write_consistency : `str`, optional
251 Name of the consistency level for write operations.
252 read_timeout : `int`, optional
253 Read timeout in seconds.
254 write_timeout : `int`, optional
255 Write timeout in seconds.
256 ra_dec_columns : `tuple` [`str`, `str`], optional
257 Names of ra/dec columns in DiaObject table.
258 replication_factor : `int`, optional
259 Replication factor used when creating new keyspace, if keyspace
260 already exists its replication factor is not changed.
261 drop : `bool`, optional
262 If `True` then drop existing tables before re-creating the schema.
263 table_options : `CreateTableOptions`, optional
264 Options used when creating Cassandra tables.
266 Returns
267 -------
268 config : `ApdbCassandraConfig`
269 Resulting configuration object for a created APDB instance.
270 """
271 # Some non-standard defaults for connection parameters, these can be
272 # changed later in generated config. Check Cassandra driver
273 # documentation for what these parameters do. These parameters are not
274 # used during database initialization, but they will be saved with
275 # generated config.
276 connection_config = ApdbCassandraConnectionConfig(
277 extra_parameters={
278 "idle_heartbeat_interval": 0,
279 "idle_heartbeat_timeout": 30,
280 "control_connection_timeout": 100,
281 },
282 )
283 config = ApdbCassandraConfig(
284 contact_points=hosts,
285 keyspace=keyspace,
286 enable_replica=enable_replica,
287 replica_skips_diaobjects=replica_skips_diaobjects,
288 connection_config=connection_config,
289 )
290 config.partitioning.time_partition_tables = time_partition_tables
291 if schema_file is not None:
292 config.schema_file = schema_file
293 if ss_schema_file is not None:
294 config.ss_schema_file = ss_schema_file
295 if read_sources_months is not None:
296 config.read_sources_months = read_sources_months
297 if read_forced_sources_months is not None:
298 config.read_forced_sources_months = read_forced_sources_months
299 if port is not None:
300 config.connection_config.port = port
301 if username is not None:
302 config.connection_config.username = username
303 if dbauth_alias is not None:
304 config.connection_config.dbauth_alias = dbauth_alias
305 if prefix is not None:
306 config.prefix = prefix
307 if part_pixelization is not None:
308 config.partitioning.part_pixelization = part_pixelization
309 if part_pix_level is not None:
310 config.partitioning.part_pix_level = part_pix_level
311 if time_partition_start is not None:
312 config.partitioning.time_partition_start = time_partition_start
313 if time_partition_end is not None:
314 config.partitioning.time_partition_end = time_partition_end
315 if read_consistency is not None:
316 config.connection_config.read_consistency = read_consistency
317 if write_consistency is not None:
318 config.connection_config.write_consistency = write_consistency
319 if read_timeout is not None:
320 config.connection_config.read_timeout = read_timeout
321 if write_timeout is not None:
322 config.connection_config.write_timeout = write_timeout
323 if ra_dec_columns is not None:
324 config.ra_dec_columns = ra_dec_columns
326 cls._makeSchema(config, drop=drop, replication_factor=replication_factor, table_options=table_options)
328 return config
330 def get_replica(self) -> ApdbCassandraReplica:
331 """Return `ApdbReplica` instance for this database."""
332 # Note that this instance has to stay alive while replica exists, so
333 # we pass reference to self.
334 return ApdbCassandraReplica(self)
336 @classmethod
337 def _makeSchema(
338 cls,
339 config: ApdbConfig,
340 *,
341 drop: bool = False,
342 replication_factor: int | None = None,
343 table_options: CreateTableOptions | None = None,
344 ) -> None:
345 # docstring is inherited from a base class
347 if not isinstance(config, ApdbCassandraConfig):
348 raise TypeError(f"Unexpected type of configuration object: {type(config)}")
350 simple_schema = ApdbSchema(config.schema_file, config.ss_schema_file)
352 with SessionContext(config) as session:
353 schema = ApdbCassandraSchema(
354 session=session,
355 keyspace=config.keyspace,
356 table_schemas=simple_schema.tableSchemas,
357 prefix=config.prefix,
358 time_partition_tables=config.partitioning.time_partition_tables,
359 enable_replica=config.enable_replica,
360 replica_skips_diaobjects=config.replica_skips_diaobjects,
361 )
363 # Ask schema to create all tables.
364 part_range_config: ApdbCassandraTimePartitionRange | None = None
365 if config.partitioning.time_partition_tables:
366 partitioner = Partitioner(config)
367 time_partition_start = astropy.time.Time(
368 config.partitioning.time_partition_start, format="isot", scale="tai"
369 )
370 time_partition_end = astropy.time.Time(
371 config.partitioning.time_partition_end, format="isot", scale="tai"
372 )
373 part_range_config = ApdbCassandraTimePartitionRange(
374 start=partitioner.time_partition(time_partition_start),
375 end=partitioner.time_partition(time_partition_end),
376 )
377 schema.makeSchema(
378 drop=drop,
379 part_range=part_range_config,
380 replication_factor=replication_factor,
381 table_options=table_options,
382 )
383 else:
384 schema.makeSchema(
385 drop=drop, replication_factor=replication_factor, table_options=table_options
386 )
388 meta_table_name = ApdbTables.metadata.table_name(config.prefix)
389 metadata = ApdbMetadataCassandra(
390 session, meta_table_name, config.keyspace, "read_tuples", "write"
391 )
393 # Fill version numbers, overrides if they existed before.
394 metadata.set(
395 ConnectionContext.metadataSchemaVersionKey, str(simple_schema.schemaVersion()), force=True
396 )
397 metadata.set(
398 ConnectionContext.metadataCodeVersionKey, str(cls.apdbImplementationVersion()), force=True
399 )
401 if config.enable_replica:
402 # Only store replica code version if replica is enabled.
403 metadata.set(
404 ConnectionContext.metadataReplicaVersionKey,
405 str(ApdbCassandraReplica.apdbReplicaImplementationVersion()),
406 force=True,
407 )
409 # Store frozen part of a configuration in metadata.
410 freezer = ApdbConfigFreezer[ApdbCassandraConfig](ConnectionContext.frozen_parameters)
411 metadata.set(ConnectionContext.metadataConfigKey, freezer.to_json(config), force=True)
413 # Store time partition range.
414 if part_range_config:
415 part_range_config.save_to_meta(metadata)
417 def getDiaObjects(self, region: sphgeom.Region) -> pandas.DataFrame:
418 # docstring is inherited from a base class
419 context = self._context
420 config = context.config
422 sp_where, num_sp_part = context.partitioner.spatial_where(region)
423 _LOG.debug("getDiaObjects: #partitions: %s", len(sp_where))
425 # We need to exclude extra partitioning columns from result.
426 column_names = context.schema.apdbColumnNames(ApdbTables.DiaObjectLast)
427 table_name = context.schema.tableName(ApdbTables.DiaObjectLast)
428 query = Select(self._keyspace, table_name, column_names)
429 statements: list[tuple] = []
430 for where_clause in sp_where:
431 full_query = query.where(where_clause)
432 statements.append(context.stmt_factory.with_params(full_query, prepare=True))
433 _LOG.debug("getDiaObjects: #queries: %s", len(statements))
435 with self._timer("select_time", tags={"table": "DiaObject", "method": "getDiaObjects"}) as timer:
436 raw_objects = cast(
437 ApdbCassandraTableData,
438 select_concurrent(
439 context.session,
440 statements,
441 "read_raw_multi",
442 config.connection_config.read_concurrency,
443 ),
444 )
445 objects = raw_objects.to_pandas(context.schema._table_schema(ApdbTables.DiaObjectLast))
446 timer.add_values(row_count=len(objects), num_sp_part=num_sp_part, num_queries=len(statements))
448 _LOG.debug("found %s DiaObjects", objects.shape[0])
449 return objects
451 def getDiaSources(
452 self,
453 region: sphgeom.Region,
454 object_ids: Iterable[int] | None,
455 visit_time: astropy.time.Time,
456 start_time: astropy.time.Time | None = None,
457 ) -> pandas.DataFrame | None:
458 # docstring is inherited from a base class
459 context = self._context
460 config = context.config
462 months = config.read_sources_months
463 if start_time is None and months == 0:
464 return None
466 mjd_end = float(visit_time.tai.mjd)
467 if start_time is None:
468 mjd_start = mjd_end - months * 30
469 else:
470 mjd_start = float(start_time.tai.mjd)
472 return self._getSources(region, object_ids, mjd_start, mjd_end, ApdbTables.DiaSource)
474 def getDiaForcedSources(
475 self,
476 region: sphgeom.Region,
477 object_ids: Iterable[int] | None,
478 visit_time: astropy.time.Time,
479 start_time: astropy.time.Time | None = None,
480 ) -> pandas.DataFrame | None:
481 # docstring is inherited from a base class
482 context = self._context
483 config = context.config
485 months = config.read_forced_sources_months
486 if start_time is None and months == 0:
487 return None
489 mjd_end = float(visit_time.tai.mjd)
490 if start_time is None:
491 mjd_start = mjd_end - months * 30
492 else:
493 mjd_start = float(start_time.tai.mjd)
495 return self._getSources(region, object_ids, mjd_start, mjd_end, ApdbTables.DiaForcedSource)
497 def getDiaObjectsForDedup(self, since: astropy.time.Time | None = None) -> pandas.DataFrame:
498 # docstring is inherited from a base class
499 context = self._context
500 config = context.config
502 if not context.has_dedup_table:
503 raise TypeError("DiaObjectDedup table does not exist in this APDB instance.")
505 if since is None:
506 # Read last deduplication time from metadata.
507 dedup_str = context.metadata.get(context.metadataDedupKey)
508 if dedup_str is not None:
509 dedup_state = json.loads(dedup_str)
510 dedup_time_str = dedup_state["dedup_time_iso_tai"]
511 since = astropy.time.Time(dedup_time_str, format="iso", scale="tai")
513 column_names = context.schema.apdbColumnNames(ExtraTables.DiaObjectDedup)
515 validity_start_column = self._timestamp_column_name("validityStart")
516 timestamp = None if since is None else self._timestamp_column_value(since)
518 table_name = context.schema.tableName(ExtraTables.DiaObjectDedup)
519 query = Select(self._keyspace, table_name, column_names, extra_clause="ALLOW FILTERING")
520 query = query.where(C("dedup_part") == 0)
521 if since is not None:
522 query = query.where(C(validity_start_column) >= 0)
524 statement = context.stmt_factory(query, prepare=False)
526 num_part = config.partitioning.num_part_dedup
527 statements = []
528 for dedup_part in range(num_part):
529 params = (dedup_part,) if timestamp is None else (dedup_part, timestamp)
530 statements.append((statement, params))
532 with self._timer(
533 "select_time", tags={"table": "DiaObjectDedup", "method": "getDiaObjectsForDedup"}
534 ) as timer:
535 objects_raw = cast(
536 ApdbCassandraTableData,
537 select_concurrent(
538 context.session,
539 statements,
540 "read_raw_multi_dedup",
541 config.connection_config.read_concurrency,
542 ),
543 )
544 objects = objects_raw.to_pandas(context.schema._table_schema(ExtraTables.DiaObjectDedup))
545 timer.add_values(row_count=len(objects), num_queries=num_part)
547 _LOG.debug("found %s DiaObjectDedup records", objects.shape[0])
548 return objects
550 def getDiaSourcesForDiaObjects(
551 self, objects: list[DiaObjectId], start_time: astropy.time.Time, max_dist_arcsec: float = 1.0
552 ) -> pandas.DataFrame:
553 # docstring is inherited from a base class
554 context = self._context
555 config = context.config
557 # Which tables to query and temporal constraints.
558 end_time = self._current_time()
559 tables, temporal_where = context.partitioner.temporal_where(
560 ApdbTables.DiaSource,
561 start_time,
562 end_time,
563 partitons_range=context.time_partitions_range,
564 query_per_time_part=False,
565 )
566 if not tables:
567 warnings.warn(
568 f"Query time range ({start_time.isot} - {end_time.isot}) does not overlap database "
569 "time partitions."
570 )
572 # Group DiaObjects by partition.
573 partitioned_object_ids = self._group_dia_objects_by_partition(
574 context.partitioner, objects, max_dist_arcsec
575 )
577 # Columns to return.
578 column_names = context.schema.apdbColumnNames(ApdbTables.DiaSource)
580 # Make a bunch of queries.
581 statements = []
582 for apdb_part, diaObjectIds in partitioned_object_ids.items():
583 spatial_where = [C("apdb_part") == apdb_part]
584 for table in tables:
585 query = Select(self._keyspace, table, column_names, extra_clause="ALLOW FILTERING")
586 for id_chunk in chunk_iterable(diaObjectIds, 10_000):
587 id_where = C("diaObjectId").in_(id_chunk)
588 for clause in QExpr.combine(spatial_where, temporal_where, extra=id_where):
589 statements.append(
590 context.stmt_factory.with_params(query.where(clause), prepare=False)
591 )
593 _LOG.debug("getDiaSourcesForDiaObjects #queries: %s", len(statements))
595 with self._timer(
596 "select_time", tags={"table": "DiaSource", "method": "getDiaSourcesForDiaObjects"}
597 ) as timer:
598 table_data_raw = cast(
599 ApdbCassandraTableData,
600 select_concurrent(
601 context.session,
602 statements,
603 "read_raw_multi",
604 config.connection_config.read_concurrency,
605 ),
606 )
607 catalog = table_data_raw.to_pandas(context.schema._table_schema(ApdbTables.DiaSource))
608 timer.add_values(row_count_from_db=len(catalog), num_queries=len(statements))
610 # precise filtering on midpointMjdTai
611 catalog = cast(pandas.DataFrame, catalog[catalog["midpointMjdTai"] >= start_time.tai.mjd])
613 timer.add_values(row_count=len(catalog))
615 _LOG.debug("found %d DiaSources", len(catalog))
616 return catalog
618 def containsVisitDetector(
619 self,
620 visit: int,
621 detector: int,
622 region: sphgeom.Region | None = None,
623 visit_time: astropy.time.Time | None = None,
624 ) -> bool:
625 # docstring is inherited from a base class
626 context = self._context
628 table_name = context.schema.tableName(ExtraTables.ApdbVisitDetector)
629 query = Select(self._keyspace, table_name, [ColumnExpr("count(*)")])
630 query = query.where((C("visit") == visit) & (C("detector") == detector))
631 stmt, params = context.stmt_factory.with_params(query, prepare=False)
633 with self._timer("contains_visit_detector_time", tags={"table": table_name}):
634 result = context.session.execute(stmt, params)
635 return bool(result.one()[0])
637 def store(
638 self,
639 visit_time: astropy.time.Time,
640 objects: pandas.DataFrame,
641 sources: pandas.DataFrame | None = None,
642 forced_sources: pandas.DataFrame | None = None,
643 ) -> None:
644 # docstring is inherited from a base class
645 context = self._context
646 config = context.config
648 # Store visit/detector in a special table, this has to be done
649 # before all other writes so if there is a failure at any point
650 # later we still have a record for attempted write.
651 visit_detector: set[tuple[int, int]] = set()
652 for df in sources, forced_sources:
653 if df is not None and not df.empty:
654 df = df[["visit", "detector"]]
655 for visit, detector in df.itertuples(index=False):
656 visit_detector.add((visit, detector))
658 if visit_detector:
659 # Typically there is only one entry, do not bother with
660 # concurrency.
661 table_name = context.schema.tableName(ExtraTables.ApdbVisitDetector)
662 query = Insert(self._keyspace, table_name, ("visit", "detector"))
663 stmt = context.stmt_factory(query)
664 for item in visit_detector:
665 context.session.execute(stmt, item, execution_profile="write")
667 objects = self._fix_input_timestamps(objects)
668 if sources is not None:
669 sources = self._fix_input_timestamps(sources)
670 if forced_sources is not None:
671 forced_sources = self._fix_input_timestamps(forced_sources)
673 replica_chunk: ReplicaChunk | None = None
674 if context.schema.replication_enabled:
675 replica_chunk = ReplicaChunk.make_replica_chunk(visit_time, config.replica_chunk_seconds)
676 self._storeReplicaChunk(replica_chunk)
678 # fill region partition column for DiaObjects
679 objects = self._add_apdb_part(objects)
680 self._storeDiaObjects(objects, visit_time, replica_chunk)
682 if sources is not None and len(sources) > 0:
683 # copy apdb_part column from DiaObjects to DiaSources
684 sources = self._add_apdb_part(sources)
685 subchunk = self._storeDiaSources(ApdbTables.DiaSource, sources, replica_chunk)
686 self._storeDiaSourcesPartitions(sources, visit_time, replica_chunk, subchunk)
688 if forced_sources is not None and len(forced_sources) > 0:
689 forced_sources = self._add_apdb_part(forced_sources)
690 self._storeDiaSources(ApdbTables.DiaForcedSource, forced_sources, replica_chunk)
692 def reassignDiaSourcesToDiaObjects(
693 self,
694 idMap: Mapping[DiaSourceId, int],
695 *,
696 increment_nDiaSources: bool = True,
697 decrement_nDiaSources: bool = True,
698 ) -> None:
699 # docstring is inherited from a base class
700 context = self._context
701 config = context.config
703 source_ids = {source_id.diaSourceId for source_id in idMap}
705 # Find all DiaSources.
706 found_sources = self._get_diasource_data(
707 idMap, "apdb_part", "diaObjectId", "ra", "dec", "midpointMjdTai"
708 )
710 if missing_ids := (source_ids - {row.diaSourceId for row in found_sources}):
711 raise LookupError(f"Some source IDs were not found in DiaSource table: {missing_ids}")
713 found_sources_by_id = {row.diaSourceId: row for row in found_sources}
715 # Make sure that all DiaObjects exist, we also want to know
716 # nDiaSources count for current and new records because we want to
717 # send updated values to replica.
718 current_object_ids = {
719 DiaObjectId(diaObjectId=row.diaObjectId, ra=row.ra, dec=row.dec) for row in found_sources
720 }
721 # Assume that DiaSource ra/dec are very close to re-assigned objects.
722 new_object_ids = {
723 DiaObjectId(diaObjectId=diaObjectId, ra=source_id.ra, dec=source_id.dec)
724 for source_id, diaObjectId in idMap.items()
725 }
726 all_object_ids = new_object_ids | current_object_ids
727 found_objects = self._get_diaobject_data(all_object_ids, "apdb_part", "ra", "dec", "nDiaSources")
729 if missing_ids := (
730 {row.diaObjectId for row in all_object_ids} - {row.diaObjectId for row in found_objects}
731 ):
732 raise LookupError(f"Some object IDs were not found in DiaObjectLast table: {missing_ids}")
734 update_records: list[ApdbUpdateRecord] = []
735 update_order = 0
736 current_time = self._current_time()
737 current_time_ns = int(current_time.unix_tai * 1e9)
739 # Update DiaSources.
740 statements: list[tuple] = []
741 for source_id, diaObjectId in idMap.items():
742 source_row = found_sources_by_id[source_id.diaSourceId]
743 apdb_part = source_row.apdb_part
744 time_part = context.partitioner.time_partition(source_row.midpointMjdTai)
746 if config.partitioning.time_partition_tables:
747 table_name = context.schema.tableName(ApdbTables.DiaSource, time_part)
748 update = (
749 Update(self._keyspace, table_name)
750 .values(C("diaObjectId").update(diaObjectId))
751 .where(C("apdb_part") == apdb_part)
752 .where(C("diaSourceId") == source_id.diaSourceId)
753 )
754 else:
755 table_name = context.schema.tableName(ApdbTables.DiaSource)
756 update = (
757 Update(self._keyspace, table_name)
758 .values(C("diaObjectId").update(diaObjectId))
759 .where(C("apdb_part") == apdb_part)
760 .where(C("apdb_time_part") == time_part)
761 .where(C("diaSourceId") == source_id.diaSourceId)
762 )
763 statements.append(context.stmt_factory.with_params(update, prepare=True))
765 if context.schema.replication_enabled:
766 update_records.append(
767 ApdbReassignDiaSourceToDiaObjectRecord(
768 diaSourceId=source_id.diaSourceId,
769 ra=source_id.ra,
770 dec=source_id.dec,
771 midpointMjdTai=source_id.midpointMjdTai,
772 diaObjectId=diaObjectId,
773 update_time_ns=current_time_ns,
774 update_order=update_order,
775 )
776 )
777 update_order += 1
779 with self._timer(
780 "update_time", tags={"table": "DiaSource", "method": "reassignDiaSourcesToDiaObjects"}
781 ) as timer:
782 execute_concurrent(context.session, statements, execution_profile="write")
783 timer.add_values(num_queries=len(statements))
785 # Update nDiaSources in DiaObjectLast. We do not update DiaObject table
786 # here because it may not even exist. PPDB updates DiaObject from
787 # update records.
788 if increment_nDiaSources or decrement_nDiaSources:
789 table_name = context.schema.tableName(ApdbTables.DiaObjectLast)
790 update = (
791 Update(self._keyspace, table_name)
792 .values(C("nDiaSources").update(-1))
793 .where(C("apdb_part") == -1)
794 .where(C("diaObjectId") == -1)
795 )
796 statement = context.stmt_factory(update, prepare=True)
797 statements = []
799 # Calculate increments/decrements for all affected DiaObjects.
800 increments: Counter = Counter()
801 if increment_nDiaSources:
802 increments.update(idMap.values())
803 if decrement_nDiaSources:
804 increments.subtract(row.diaObjectId for row in found_sources)
806 for row in found_objects:
807 if increments.get(row.diaObjectId):
808 nDiaSources = row.nDiaSources + increments[row.diaObjectId]
809 statements.append((statement, (nDiaSources, row.apdb_part, row.diaObjectId)))
811 # Also send updated values to replica.
812 if context.schema.replication_enabled:
813 update_records.append(
814 ApdbUpdateNDiaSourcesRecord(
815 diaObjectId=row.diaObjectId,
816 ra=row.ra,
817 dec=row.dec,
818 nDiaSources=nDiaSources,
819 update_time_ns=current_time_ns,
820 update_order=update_order,
821 )
822 )
823 update_order += 1
825 if statements:
826 with self._timer(
827 "update_time", tags={"table": table_name, "method": "reassignDiaSourcesToDiaObjects"}
828 ) as timer:
829 execute_concurrent(context.session, statements, execution_profile="write")
830 timer.add_values(num_queries=len(statements))
832 if update_records:
833 replica_chunk = ReplicaChunk.make_replica_chunk(current_time, config.replica_chunk_seconds)
834 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True)
836 def setValidityEnd(
837 self, objects: list[DiaObjectId], validityEnd: astropy.time.Time, raise_on_missing_id: bool = False
838 ) -> int:
839 # docstring is inherited from a base class
840 if not objects:
841 return 0
843 context = self._context
844 config = context.config
846 pad_arcsec = 1.0
847 partitioned_object_ids = self._group_dia_objects_by_partition(
848 context.partitioner, objects, pad_arcsec
849 )
851 # Check that all objects exist.
852 table_name = context.schema.tableName(ApdbTables.DiaObjectLast)
853 statements: list[tuple] = []
854 for apdb_part, diaObjectIds in partitioned_object_ids.items():
855 query = Select(self._keyspace, table_name, ["apdb_part", "diaObjectId"])
856 query = query.where(C("apdb_part") == apdb_part)
857 query = query.where(C("diaObjectId").in_(diaObjectIds))
858 statements.append(context.stmt_factory.with_params(query, prepare=False))
860 with self._timer("select_time", tags={"table": table_name, "method": "setValidityEnd"}) as timer:
861 records = cast(
862 list[tuple[int, int]],
863 select_concurrent(
864 context.session,
865 statements,
866 "read_tuples",
867 config.connection_config.read_concurrency,
868 ),
869 )
870 timer.add_values(row_count=len(objects))
872 requested_ids = {obj.diaObjectId for obj in objects}
873 found_ids = {rec[1] for rec in records}
874 if extra_ids := (found_ids - requested_ids):
875 raise RuntimeError(f"Consistency error - found duplicate records for object IDs: {extra_ids}")
876 if raise_on_missing_id:
877 if missing_ids := (requested_ids - found_ids):
878 raise LookupError(f"Some object IDs are missing from DiaObjectLast table: {missing_ids}")
880 # Filter existing records.
881 if len(objects) != len(found_ids):
882 objects = [obj for obj in objects if obj.diaObjectId in found_ids]
884 if not objects:
885 return 0
887 # Group by partitions again.
888 grouped_object_ids: dict[int, list[int]] = defaultdict(list)
889 for apdb_part, diaObjectId in records:
890 grouped_object_ids[apdb_part].append(diaObjectId)
892 # Remove all matching rows from DiaObjectLast.
893 statements = []
894 for apdb_part, diaObjectIds in grouped_object_ids.items():
895 delete = (
896 Delete(self._keyspace, table_name)
897 .where(C("apdb_part") == apdb_part)
898 .where(C("diaObjectId").in_(diaObjectIds))
899 )
900 statements.append(context.stmt_factory.with_params(delete))
902 # Also remove from DiaObjectLastToPartition.
903 reverse_table_name = context.schema.tableName(ExtraTables.DiaObjectLastToPartition)
904 delete = Delete(self._keyspace, reverse_table_name).where(
905 C("diaObjectId").in_([rec[1] for rec in records])
906 )
907 statements.append(context.stmt_factory.with_params(delete))
909 with self._timer("delete_time", tags={"table": table_name, "method": "setValidityEnd"}) as timer:
910 execute_concurrent(context.session, statements, execution_profile="write")
911 timer.add_values(row_count=len(records))
913 # If repication is enabled then send all updates.
914 if context.schema.replication_enabled:
915 current_time = self._current_time()
916 current_time_ns = int(current_time.unix_tai * 1e9)
917 replica_chunk = ReplicaChunk.make_replica_chunk(current_time, config.replica_chunk_seconds)
919 update_records = [
920 ApdbCloseDiaObjectValidityRecord(
921 diaObjectId=obj.diaObjectId,
922 ra=obj.ra,
923 dec=obj.dec,
924 update_time_ns=current_time_ns,
925 update_order=index,
926 validityEndMjdTai=float(validityEnd.tai.mjd),
927 nDiaSources=None,
928 )
929 for index, obj in enumerate(objects)
930 ]
932 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True)
934 return len(objects)
936 def resetDedup(self, dedup_time: astropy.time.Time | None = None) -> None:
937 # docstring is inherited from a base class
938 context = self._context
940 if not context.has_dedup_table:
941 raise TypeError("DiaObjectDedup table does not exist in this APDB instance.")
943 if dedup_time is None:
944 dedup_time = self._current_time()
946 validity_start_column = self._timestamp_column_name("validityStart")
948 # Find latest timestamp in deduplication table.
949 table_name = context.schema.tableName(ExtraTables.DiaObjectDedup)
950 query = Select(self._keyspace, table_name, [ColumnExpr(f'MAX("{validity_start_column}")')])
951 stmt = context.stmt_factory(query, prepare=False)
953 result = context.session.execute(stmt, execution_profile="read_tuples")
954 max_value = result.one()[0]
955 if self._schema.has_mjd_timestamps:
956 max_validity_start = astropy.time.Time(max_value, format="mjd", scale="tai")
957 else:
958 max_validity_start = astropy.time.Time(max_value, format="datetime", scale="tai")
960 # If max time is lower than dedup time we can do TRUNCATE.
961 if dedup_time >= max_validity_start:
962 query_str = f'TRUNCATE TABLE "{self._keyspace}"."{table_name}"'
963 context.session.execute(query_str, execution_profile="write")
964 else:
965 dedup_time_value = self._timestamp_column_value(dedup_time)
966 delete = Delete(self._keyspace, table_name).where(C(validity_start_column) < dedup_time_value)
967 stmt, params = context.stmt_factory.with_params(delete)
968 context.session.execute(stmt, params, execution_profile="write")
970 # Store dedup time.
971 data = {"dedup_time_iso_tai": dedup_time.tai.to_value("iso")}
972 data_json = json.dumps(data)
973 context.metadata.set(context.metadataDedupKey, data_json, force=True)
975 def reassignDiaSources(self, idMap: Mapping[int, int]) -> None:
976 # docstring is inherited from a base class
977 context = self._context
978 config = context.config
980 now = self._current_time()
981 reassign_time_column = self._timestamp_column_name("ssObjectReassocTime")
982 reassignTime = self._timestamp_column_value(now)
984 # To update a record we need to know its exact primary key (including
985 # partition key) so we start by querying for diaSourceId to find the
986 # primary keys.
988 table_name = context.schema.tableName(ExtraTables.DiaSourceToPartition)
989 # split it into 1k IDs per query
990 selects: list[tuple] = []
991 columns = ["diaSourceId", "apdb_part", "apdb_time_part", "apdb_replica_chunk"]
992 query = Select(self._keyspace, table_name, columns)
993 for ids in chunk_iterable(idMap.keys(), 1_000):
994 full_query = query.where(C("diaSourceId").in_(ids))
995 selects.append(context.stmt_factory.with_params(full_query, prepare=False))
997 # No need for DataFrame here, read data as tuples.
998 result = cast(
999 list[tuple[int, int, int, int | None]],
1000 select_concurrent(
1001 context.session, selects, "read_tuples", config.connection_config.read_concurrency
1002 ),
1003 )
1005 # Make mapping from source ID to its partition.
1006 id2partitions: dict[int, tuple[int, int]] = {}
1007 id2chunk_id: dict[int, int] = {}
1008 for row in result:
1009 id2partitions[row[0]] = row[1:3]
1010 if row[3] is not None:
1011 id2chunk_id[row[0]] = row[3]
1013 # make sure we know partitions for each ID
1014 if set(id2partitions) != set(idMap):
1015 missing = ",".join(str(item) for item in set(idMap) - set(id2partitions))
1016 raise ValueError(f"Following DiaSource IDs do not exist in the database: {missing}")
1018 # Reassign in standard tables
1019 queries: list[tuple[cassandra.query.PreparedStatement, tuple]] = []
1020 for diaSourceId, ssObjectId in idMap.items():
1021 apdb_part, apdb_time_part = id2partitions[diaSourceId]
1022 if config.partitioning.time_partition_tables:
1023 table_name = context.schema.tableName(ApdbTables.DiaSource, apdb_time_part)
1024 update = (
1025 Update(self._keyspace, table_name)
1026 .values(
1027 C("ssObjectId").update(ssObjectId),
1028 C("diaObjectId").update(None),
1029 C(reassign_time_column).update(reassignTime),
1030 )
1031 .where(C("apdb_part") == apdb_part)
1032 .where(C("diaSourceId") == diaSourceId)
1033 )
1034 else:
1035 table_name = context.schema.tableName(ApdbTables.DiaSource)
1036 update = (
1037 Update(self._keyspace, table_name)
1038 .values(
1039 C("ssObjectId").update(ssObjectId),
1040 C("diaObjectId").update(None),
1041 C(reassign_time_column).update(reassignTime),
1042 )
1043 .where(C("apdb_part") == apdb_part)
1044 .where(C("apdb_time_part") == apdb_time_part)
1045 .where(C("diaSourceId") == diaSourceId)
1046 )
1047 queries.append(context.stmt_factory.with_params(update, prepare=True))
1049 # TODO: (DM-50190) Replication for updated records is not implemented.
1050 if id2chunk_id:
1051 warnings.warn("Replication of reassigned DiaSource records is not implemented.", stacklevel=2)
1053 _LOG.debug("%s: will update %d records", table_name, len(idMap))
1054 with self._timer("source_reassign_time") as timer:
1055 execute_concurrent(context.session, queries, execution_profile="write")
1056 timer.add_values(source_count=len(idMap))
1058 def withdrawDiaSources(
1059 self,
1060 diaSourceIds: Iterable[DiaSourceId],
1061 *,
1062 timeWithdrawn: astropy.time.Time | None = None,
1063 ) -> None:
1064 # docstring is inherited from a base class
1065 context = self._context
1066 config = context.config
1068 if timeWithdrawn is None:
1069 timeWithdrawn = self._current_time()
1070 time_value = self._timestamp_column_value(timeWithdrawn)
1071 column_name = self._timestamp_column_name("time_withdrawn")
1073 diaSourceIds = list(diaSourceIds)
1074 source_ids = {source_id.diaSourceId for source_id in diaSourceIds}
1076 # Find all DiaSources.
1077 found_sources = self._get_diasource_data(
1078 diaSourceIds, "apdb_part", "diaObjectId", "ra", "dec", "midpointMjdTai", column_name
1079 )
1081 if missing_ids := (source_ids - {row.diaSourceId for row in found_sources}):
1082 raise LookupError(f"Some source IDs were not found in DiaSource table: {missing_ids}")
1084 found_sources_by_id = {row.diaSourceId: row for row in found_sources}
1086 update_records: list[ApdbUpdateRecord] = []
1087 update_order = 0
1088 current_time = self._current_time()
1089 current_time_ns = int(current_time.unix_tai * 1e9)
1091 # Update DiaSources.
1092 statements: list[tuple] = []
1093 for source_id in diaSourceIds:
1094 source_row = found_sources_by_id[source_id.diaSourceId]
1095 # Ignore sources already withdrawn.
1096 if getattr(source_row, column_name) is not None:
1097 continue
1099 apdb_part = source_row.apdb_part
1100 time_part = context.partitioner.time_partition(source_row.midpointMjdTai)
1102 if config.partitioning.time_partition_tables:
1103 table_name = context.schema.tableName(ApdbTables.DiaSource, time_part)
1104 update = (
1105 Update(self._keyspace, table_name)
1106 .values(C(column_name).update(time_value))
1107 .where(C("apdb_part") == apdb_part)
1108 .where(C("diaSourceId") == source_id.diaSourceId)
1109 )
1110 else:
1111 table_name = context.schema.tableName(ApdbTables.DiaSource)
1112 update = (
1113 Update(self._keyspace, table_name)
1114 .values(C(column_name).update(time_value))
1115 .where(C("apdb_part") == apdb_part)
1116 .where(C("apdb_time_part") == time_part)
1117 .where(C("diaSourceId") == source_id.diaSourceId)
1118 )
1119 statements.append(context.stmt_factory.with_params(update, prepare=True))
1121 if context.schema.replication_enabled:
1122 update_records.append(
1123 ApdbWithdrawDiaSourceRecord(
1124 diaSourceId=source_id.diaSourceId,
1125 ra=source_id.ra,
1126 dec=source_id.dec,
1127 midpointMjdTai=source_id.midpointMjdTai,
1128 update_time_ns=current_time_ns,
1129 update_order=update_order,
1130 timeWithdrawnMjdTai=float(timeWithdrawn.tai.mjd),
1131 )
1132 )
1133 update_order += 1
1135 with self._timer("update_time", tags={"table": "DiaSource", "method": "withdrawDiaSources"}) as timer:
1136 execute_concurrent(context.session, statements, execution_profile="write")
1137 timer.add_values(num_queries=len(statements))
1139 if update_records:
1140 replica_chunk = ReplicaChunk.make_replica_chunk(current_time, config.replica_chunk_seconds)
1141 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True)
1143 def withdrawDiaForcedSources(
1144 self,
1145 diaForcedSourceIds: Iterable[DiaForcedSourceId],
1146 *,
1147 timeWithdrawn: astropy.time.Time | None = None,
1148 ) -> None:
1149 # docstring is inherited from a base class
1150 context = self._context
1151 config = context.config
1153 def _fsrc_id(fsource: Any) -> tuple[int, int, int]:
1154 return (fsource.diaObjectId, fsource.visit, fsource.detector)
1156 if timeWithdrawn is None:
1157 timeWithdrawn = self._current_time()
1158 time_value = self._timestamp_column_value(timeWithdrawn)
1159 column_name = self._timestamp_column_name("time_withdrawn")
1161 diaForcedSourceIds = list(diaForcedSourceIds)
1162 fsource_keys = {_fsrc_id(source) for source in diaForcedSourceIds}
1164 found_fsources = self._get_diaforcedsource_data(
1165 diaForcedSourceIds, "apdb_part", "ra", "dec", "midpointMjdTai", column_name
1166 )
1168 found_keys = {_fsrc_id(row) for row in found_fsources}
1169 if missing_ids := (fsource_keys - found_keys):
1170 raise LookupError(f"Some source IDs were not found in DiaForcedSource table: {missing_ids}")
1172 statements: list[tuple] = []
1173 update_records = []
1174 update_order = 0
1175 current_time = self._current_time()
1176 current_time_ns = int(current_time.unix_tai * 1e9)
1178 for source_row in found_fsources:
1179 # Ignore sources already withdrawn.
1180 if getattr(source_row, column_name) is not None:
1181 continue
1183 apdb_part = source_row.apdb_part
1184 time_part = context.partitioner.time_partition(source_row.midpointMjdTai)
1186 if config.partitioning.time_partition_tables:
1187 table_name = context.schema.tableName(ApdbTables.DiaForcedSource, time_part)
1188 update = (
1189 Update(self._keyspace, table_name)
1190 .values(C(column_name).update(time_value))
1191 .where(C("apdb_part") == apdb_part)
1192 .where(C("diaObjectId") == source_row.diaObjectId)
1193 .where(C("visit") == source_row.visit)
1194 .where(C("detector") == source_row.detector)
1195 )
1196 else:
1197 table_name = context.schema.tableName(ApdbTables.DiaForcedSource)
1198 update = (
1199 Update(self._keyspace, table_name)
1200 .values(C(column_name).update(time_value))
1201 .where(C("apdb_part") == apdb_part)
1202 .where(C("apdb_time_part") == time_part)
1203 .where(C("diaObjectId") == source_row.diaObjectId)
1204 .where(C("visit") == source_row.visit)
1205 .where(C("detector") == source_row.detector)
1206 )
1207 statements.append(context.stmt_factory.with_params(update, prepare=True))
1209 if context.schema.replication_enabled:
1210 update_records.append(
1211 ApdbWithdrawDiaForcedSourceRecord(
1212 diaObjectId=source_row.diaObjectId,
1213 visit=source_row.visit,
1214 detector=source_row.detector,
1215 ra=source_row.ra,
1216 dec=source_row.dec,
1217 midpointMjdTai=source_row.midpointMjdTai,
1218 update_time_ns=current_time_ns,
1219 update_order=update_order,
1220 timeWithdrawnMjdTai=float(timeWithdrawn.tai.mjd),
1221 )
1222 )
1223 update_order += 1
1225 if statements:
1226 with self._timer(
1227 "update_time", tags={"table": "DiaForcedSource", "method": "withdrawDiaForcedSources"}
1228 ) as timer:
1229 execute_concurrent(context.session, statements, execution_profile="write")
1230 timer.add_values(num_queries=len(statements))
1232 if update_records:
1233 replica_chunk = ReplicaChunk.make_replica_chunk(current_time, config.replica_chunk_seconds)
1234 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True)
1236 def countUnassociatedObjects(self) -> int:
1237 # docstring is inherited from a base class
1239 # It's too inefficient to implement it for Cassandra in current schema.
1240 raise NotImplementedError()
1242 @property
1243 def schema(self) -> ApdbSchema:
1244 # docstring is inherited from a base class
1245 return self._schema
1247 @property
1248 def metadata(self) -> ApdbMetadata:
1249 # docstring is inherited from a base class
1250 context = self._context
1251 return context.metadata
1253 @property
1254 def admin(self) -> ApdbCassandraAdmin:
1255 # docstring is inherited from a base class
1256 return ApdbCassandraAdmin(self)
1258 def _getSources(
1259 self,
1260 region: sphgeom.Region,
1261 object_ids: Iterable[int] | None,
1262 mjd_start: float,
1263 mjd_end: float,
1264 table_name: ApdbTables,
1265 ) -> pandas.DataFrame:
1266 """Return catalog of DiaSource instances given set of DiaObject IDs.
1268 Parameters
1269 ----------
1270 region : `lsst.sphgeom.Region`
1271 Spherical region.
1272 object_ids :
1273 Collection of DiaObject IDs
1274 mjd_start : `float`
1275 Lower bound of time interval.
1276 mjd_end : `float`
1277 Upper bound of time interval.
1278 table_name : `ApdbTables`
1279 Name of the table.
1281 Returns
1282 -------
1283 catalog : `pandas.DataFrame`, or `None`
1284 Catalog containing DiaSource records. Empty catalog is returned if
1285 ``object_ids`` is empty.
1286 """
1287 context = self._context
1288 config = context.config
1290 object_id_set: Set[int] = set()
1291 if object_ids is not None:
1292 object_id_set = set(object_ids)
1293 if len(object_id_set) == 0:
1294 return self._make_empty_catalog(table_name)
1296 sp_where, num_sp_part = context.partitioner.spatial_where(region)
1297 tables, temporal_where = context.partitioner.temporal_where(
1298 table_name, mjd_start, mjd_end, partitons_range=context.time_partitions_range
1299 )
1300 if not tables:
1301 start = astropy.time.Time(mjd_start, format="mjd", scale="tai")
1302 end = astropy.time.Time(mjd_end, format="mjd", scale="tai")
1303 warnings.warn(
1304 f"Query time range ({start.isot} - {end.isot}) does not overlap database time partitions."
1305 )
1307 # We need to exclude extra partitioning columns from result.
1308 column_names = context.schema.apdbColumnNames(table_name)
1310 # Build all queries
1311 statements: list[tuple] = []
1312 for table in tables:
1313 query = Select(self._keyspace, table, column_names)
1314 for clause in QExpr.combine(sp_where, temporal_where):
1315 statements.append(context.stmt_factory.with_params(query.where(clause), prepare=True))
1316 _LOG.debug("_getSources %s: #queries: %s", table_name, len(statements))
1318 with self._timer("select_time", tags={"table": table_name.name, "method": "_getSources"}) as timer:
1319 table_data_raw = cast(
1320 ApdbCassandraTableData,
1321 select_concurrent(
1322 context.session,
1323 statements,
1324 "read_raw_multi",
1325 config.connection_config.read_concurrency,
1326 ),
1327 )
1328 catalog = table_data_raw.to_pandas(context.schema._table_schema(table_name))
1329 timer.add_values(
1330 row_count_from_db=len(catalog), num_sp_part=num_sp_part, num_queries=len(statements)
1331 )
1333 # filter by given object IDs
1334 if len(object_id_set) > 0:
1335 catalog = cast(pandas.DataFrame, catalog[catalog["diaObjectId"].isin(object_id_set)])
1337 # precise filtering on midpointMjdTai
1338 catalog = cast(pandas.DataFrame, catalog[catalog["midpointMjdTai"] > mjd_start])
1340 timer.add_values(row_count=len(catalog))
1342 _LOG.debug("found %d %ss", catalog.shape[0], table_name.name)
1343 return catalog
1345 def _storeReplicaChunk(self, replica_chunk: ReplicaChunk) -> None:
1346 context = self._context
1347 config = context.config
1349 # Cassandra timestamp uses milliseconds since epoch
1350 timestamp = int(replica_chunk.last_update_time.unix_tai * 1000)
1352 # everything goes into a single partition
1353 partition = 0
1355 table_name = context.schema.tableName(ExtraTables.ApdbReplicaChunks)
1357 columns = ["partition", "apdb_replica_chunk", "last_update_time", "unique_id"]
1358 values = [partition, replica_chunk.id, timestamp, replica_chunk.unique_id]
1359 if context.has_chunk_sub_partitions:
1360 columns.append("has_subchunks")
1361 values.append(True)
1363 query = Insert(self._keyspace, table_name, columns)
1364 stmt = context.stmt_factory(query)
1366 context.session.execute(
1367 stmt,
1368 values,
1369 timeout=config.connection_config.write_timeout,
1370 execution_profile="write",
1371 )
1373 def _queryDiaObjectLastPartitions(self, ids: Iterable[int]) -> Mapping[int, int]:
1374 """Return existing mapping of diaObjectId to its last partition."""
1375 context = self._context
1376 config = context.config
1378 table_name = context.schema.tableName(ExtraTables.DiaObjectLastToPartition)
1379 queries = []
1380 object_count = 0
1381 for id_chunk in chunk_iterable(ids, 10_000):
1382 id_chunk_list = tuple(id_chunk)
1383 query = Select(self._keyspace, table_name, ("diaObjectId", "apdb_part"))
1384 query = query.where(C("diaObjectId").in_(id_chunk_list))
1385 queries.append(context.stmt_factory.with_params(query, prepare=False))
1386 object_count += len(id_chunk_list)
1388 with self._timer("query_object_last_partitions", tags={"table": table_name}) as timer:
1389 data = cast(
1390 ApdbTableData,
1391 select_concurrent(
1392 context.session,
1393 queries,
1394 "read_raw_multi",
1395 config.connection_config.read_concurrency,
1396 ),
1397 )
1398 timer.add_values(object_count=object_count, row_count=len(data.rows()))
1400 if data.column_names() != ["diaObjectId", "apdb_part"]:
1401 raise RuntimeError(f"Unexpected column names in query result: {data.column_names()}")
1403 return {row[0]: row[1] for row in data.rows()}
1405 def _deleteMovingObjects(self, objs: pandas.DataFrame) -> None:
1406 """Objects in DiaObjectsLast can move from one spatial partition to
1407 another. For those objects inserting new version does not replace old
1408 one, so we need to explicitly remove old versions before inserting new
1409 ones.
1410 """
1411 context = self._context
1413 # Extract all object IDs.
1414 new_partitions = dict(zip(objs["diaObjectId"], objs["apdb_part"]))
1415 old_partitions = self._queryDiaObjectLastPartitions(objs["diaObjectId"])
1417 moved_oids: dict[int, tuple[int, int]] = {}
1418 for oid, old_part in old_partitions.items():
1419 new_part = new_partitions.get(oid, old_part)
1420 if new_part != old_part:
1421 moved_oids[oid] = (old_part, new_part)
1422 _LOG.debug("DiaObject IDs that moved to new partition: %s", moved_oids)
1424 if moved_oids:
1425 # Delete old records from DiaObjectLast.
1426 table_name = context.schema.tableName(ApdbTables.DiaObjectLast)
1427 query = Delete(self._keyspace, table_name)
1428 query = query.where('apdb_part = {} AND "diaObjectId" = {}', (-1, -1))
1429 statement = context.stmt_factory(query, prepare=True)
1430 queries = []
1431 for oid, (old_part, _) in moved_oids.items():
1432 queries.append((statement, (old_part, oid)))
1433 with self._timer("delete_object_last", tags={"table": table_name}) as timer:
1434 execute_concurrent(context.session, queries, execution_profile="write")
1435 timer.add_values(row_count=len(moved_oids))
1437 # Add all new records to the map.
1438 table_name = context.schema.tableName(ExtraTables.DiaObjectLastToPartition)
1439 insert = Insert(self._keyspace, table_name, ("diaObjectId", "apdb_part"))
1440 statement = context.stmt_factory(insert, prepare=True)
1442 queries = []
1443 for oid, new_part in new_partitions.items():
1444 queries.append((statement, (oid, new_part)))
1446 with self._timer("update_object_last_partition", tags={"table": table_name}) as timer:
1447 execute_concurrent(context.session, queries, execution_profile="write")
1448 timer.add_values(row_count=len(queries))
1450 def _storeDiaObjects(
1451 self, objs: pandas.DataFrame, visit_time: astropy.time.Time, replica_chunk: ReplicaChunk | None
1452 ) -> None:
1453 """Store catalog of DiaObjects from current visit.
1455 Parameters
1456 ----------
1457 objs : `pandas.DataFrame`
1458 Catalog with DiaObject records
1459 visit_time : `astropy.time.Time`
1460 Time of the current visit.
1461 replica_chunk : `ReplicaChunk` or `None`
1462 Replica chunk identifier if replication is configured.
1463 """
1464 if len(objs) == 0:
1465 _LOG.debug("No objects to write to database.")
1466 return
1468 context = self._context
1469 config = context.config
1471 self._deleteMovingObjects(objs)
1473 validity_start_column = self._timestamp_column_name("validityStart")
1474 timestamp = self._timestamp_column_value(visit_time)
1476 # DiaObjectLast did not have this column in the past.
1477 extra_columns: dict[str, Any] = {}
1478 if context.schema.check_column(ApdbTables.DiaObjectLast, validity_start_column):
1479 extra_columns[validity_start_column] = timestamp
1481 self._storeObjectsPandas(objs, ApdbTables.DiaObjectLast, extra_columns=extra_columns)
1483 extra_columns[validity_start_column] = timestamp
1484 visit_time_part = context.partitioner.time_partition(visit_time)
1485 time_part: int | None = visit_time_part
1486 if (time_partitions_range := context.time_partitions_range) is not None:
1487 self._check_time_partitions([visit_time_part], time_partitions_range)
1488 if not config.partitioning.time_partition_tables:
1489 extra_columns["apdb_time_part"] = time_part
1490 time_part = None
1492 # Only store DiaObects if not doing replication or explicitly
1493 # configured to always store them.
1494 if replica_chunk is None or not config.replica_skips_diaobjects:
1495 self._storeObjectsPandas(
1496 objs, ApdbTables.DiaObject, extra_columns=extra_columns, time_part=time_part
1497 )
1499 if replica_chunk is not None:
1500 extra_columns = {"apdb_replica_chunk": replica_chunk.id, validity_start_column: timestamp}
1501 table = ExtraTables.DiaObjectChunks
1502 if context.has_chunk_sub_partitions:
1503 table = ExtraTables.DiaObjectChunks2
1504 # Use a random number for a second part of partitioning key so
1505 # that different clients could wrtite to different partitions.
1506 # This makes it not exactly reproducible.
1507 extra_columns["apdb_replica_subchunk"] = random.randrange(config.replica_sub_chunk_count)
1508 self._storeObjectsPandas(objs, table, extra_columns=extra_columns)
1510 # Store copy of the records in dedup table.
1511 if context.has_dedup_table:
1512 table = ExtraTables.DiaObjectDedup
1513 extra_columns = {
1514 "dedup_part": random.randrange(config.partitioning.num_part_dedup),
1515 validity_start_column: timestamp,
1516 }
1517 self._storeObjectsPandas(objs, table, extra_columns=extra_columns)
1519 def _storeDiaSources(
1520 self,
1521 table_name: ApdbTables,
1522 sources: pandas.DataFrame,
1523 replica_chunk: ReplicaChunk | None,
1524 ) -> int | None:
1525 """Store catalog of DIASources or DIAForcedSources from current visit.
1527 Parameters
1528 ----------
1529 table_name : `ApdbTables`
1530 Table where to store the data.
1531 sources : `pandas.DataFrame`
1532 Catalog containing DiaSource records
1533 visit_time : `astropy.time.Time`
1534 Time of the current visit.
1535 replica_chunk : `ReplicaChunk` or `None`
1536 Replica chunk identifier if replication is configured.
1538 Returns
1539 -------
1540 subchunk : `int` or `None`
1541 Subchunk number for resulting replica data, `None` if relication is
1542 not enabled ot subchunking is not enabled.
1543 """
1544 context = self._context
1545 config = context.config
1547 # Time partitioning has to be based on midpointMjdTai, not visit_time
1548 # as visit_time is not really a visit time.
1549 tp_sources = sources.copy(deep=False)
1550 tp_sources["apdb_time_part"] = tp_sources["midpointMjdTai"].apply(context.partitioner.time_partition)
1551 if (time_partitions_range := context.time_partitions_range) is not None:
1552 self._check_time_partitions(tp_sources["apdb_time_part"], time_partitions_range)
1553 extra_columns: dict[str, Any] = {}
1554 if not config.partitioning.time_partition_tables:
1555 self._storeObjectsPandas(tp_sources, table_name)
1556 else:
1557 # Group by time partition
1558 partitions = set(tp_sources["apdb_time_part"])
1559 if len(partitions) == 1:
1560 # Single partition - just save the whole thing.
1561 time_part = partitions.pop()
1562 self._storeObjectsPandas(sources, table_name, time_part=time_part)
1563 else:
1564 # group by time partition.
1565 for time_part, sub_frame in tp_sources.groupby(by="apdb_time_part"):
1566 sub_frame.drop(columns="apdb_time_part", inplace=True)
1567 self._storeObjectsPandas(sub_frame, table_name, time_part=time_part)
1569 subchunk: int | None = None
1570 if replica_chunk is not None:
1571 extra_columns = {"apdb_replica_chunk": replica_chunk.id}
1572 if context.has_chunk_sub_partitions:
1573 subchunk = random.randrange(config.replica_sub_chunk_count)
1574 extra_columns["apdb_replica_subchunk"] = subchunk
1575 if table_name is ApdbTables.DiaSource:
1576 extra_table = ExtraTables.DiaSourceChunks2
1577 else:
1578 extra_table = ExtraTables.DiaForcedSourceChunks2
1579 else:
1580 if table_name is ApdbTables.DiaSource:
1581 extra_table = ExtraTables.DiaSourceChunks
1582 else:
1583 extra_table = ExtraTables.DiaForcedSourceChunks
1584 self._storeObjectsPandas(sources, extra_table, extra_columns=extra_columns)
1586 return subchunk
1588 def _check_time_partitions(
1589 self, partitions: Iterable[int], time_partitions_range: ApdbCassandraTimePartitionRange
1590 ) -> None:
1591 """Check that time partitons for new data actually exist.
1593 Parameters
1594 ----------
1595 partitions : `~collections.abc.Iterable` [`int`]
1596 Time partitions for new data.
1597 time_partitions_range : `ApdbCassandraTimePartitionRange`
1598 Currrent time partition range.
1599 """
1600 partitions = set(partitions)
1601 min_part = min(partitions)
1602 max_part = max(partitions)
1603 if min_part < time_partitions_range.start or max_part > time_partitions_range.end:
1604 raise ValueError(
1605 "Attempt to store data for time partitions that do not yet exist. "
1606 f"Partitons for new records: {min_part}-{max_part}. "
1607 f"Database partitons: {time_partitions_range.start}-{time_partitions_range.end}."
1608 )
1609 # Make a noise when writing to the last partition.
1610 if max_part == time_partitions_range.end:
1611 warnings.warn(
1612 "Writing into the last temporal partition. Partition range needs to be extended soon.",
1613 stacklevel=3,
1614 )
1616 def _storeDiaSourcesPartitions(
1617 self,
1618 sources: pandas.DataFrame,
1619 visit_time: astropy.time.Time,
1620 replica_chunk: ReplicaChunk | None,
1621 subchunk: int | None,
1622 ) -> None:
1623 """Store mapping of diaSourceId to its partitioning values.
1625 Parameters
1626 ----------
1627 sources : `pandas.DataFrame`
1628 Catalog containing DiaSource records
1629 visit_time : `astropy.time.Time`
1630 Time of the current visit.
1631 replica_chunk : `ReplicaChunk` or `None`
1632 Replication chunk, or `None` when replication is disabled.
1633 subchunk : `int` or `None`
1634 Replication sub-chunk, or `None` when replication is disabled or
1635 sub-chunking is not used.
1636 """
1637 context = self._context
1639 id_map = cast(pandas.DataFrame, sources[["diaSourceId", "apdb_part"]])
1640 extra_columns = {
1641 "apdb_time_part": context.partitioner.time_partition(visit_time),
1642 "apdb_replica_chunk": replica_chunk.id if replica_chunk is not None else None,
1643 }
1644 if context.has_chunk_sub_partitions:
1645 extra_columns["apdb_replica_subchunk"] = subchunk
1647 self._storeObjectsPandas(
1648 id_map, ExtraTables.DiaSourceToPartition, extra_columns=extra_columns, time_part=None
1649 )
1651 def _storeObjectsPandas(
1652 self,
1653 records: pandas.DataFrame,
1654 table_name: ApdbTables | ExtraTables,
1655 extra_columns: Mapping | None = None,
1656 time_part: int | None = None,
1657 ) -> None:
1658 """Store generic objects.
1660 Takes Pandas catalog and stores a bunch of records in a table.
1662 Parameters
1663 ----------
1664 records : `pandas.DataFrame`
1665 Catalog containing object records
1666 table_name : `ApdbTables`
1667 Name of the table as defined in APDB schema.
1668 extra_columns : `dict`, optional
1669 Mapping (column_name, column_value) which gives fixed values for
1670 columns in each row, overrides values in ``records`` if matching
1671 columns exist there.
1672 time_part : `int`, optional
1673 If not `None` then insert into a per-partition table.
1675 Notes
1676 -----
1677 If Pandas catalog contains additional columns not defined in table
1678 schema they are ignored. Catalog does not have to contain all columns
1679 defined in a table, but partition and clustering keys must be present
1680 in a catalog or ``extra_columns``.
1681 """
1682 context = self._context
1684 # use extra columns if specified
1685 if extra_columns is None:
1686 extra_columns = {}
1687 extra_fields = list(extra_columns.keys())
1689 # Fields that will come from dataframe.
1690 df_fields = [column for column in records.columns if column not in extra_fields]
1692 column_map = context.schema.getColumnMap(table_name)
1693 # list of columns (as in felis schema)
1694 fields = [column_map[field].name for field in df_fields if field in column_map]
1695 fields += extra_fields
1697 # check that all partitioning and clustering columns are defined
1698 partition_columns = context.schema.partitionColumns(table_name)
1699 required_columns = partition_columns + context.schema.clusteringColumns(table_name)
1700 missing_columns = [column for column in required_columns if column not in fields]
1701 if missing_columns:
1702 raise ValueError(f"Primary key columns are missing from catalog: {missing_columns}")
1704 batch_size = self._batch_size(table_name)
1706 with self._timer("insert_build_time", tags={"table": table_name.name}):
1707 # Multi-partition batches are problematic in general, so we want to
1708 # group records in a batch by their partition key.
1709 values_by_key: dict[tuple, list[list]] = defaultdict(list)
1710 for rec in records.itertuples(index=False):
1711 values = []
1712 partitioning_values: dict[str, Any] = {}
1713 for field in df_fields:
1714 if field not in column_map:
1715 continue
1716 value = getattr(rec, field)
1717 if column_map[field].datatype is felis.datamodel.DataType.timestamp:
1718 if isinstance(value, pandas.Timestamp):
1719 value = value.to_pydatetime()
1720 elif value is pandas.NaT:
1721 value = None
1722 else:
1723 # Assume it's seconds since epoch, Cassandra
1724 # datetime is in milliseconds
1725 value = int(value * 1000)
1726 value = literal(value)
1727 values.append(UNSET_VALUE if value is None else value)
1728 if field in partition_columns:
1729 partitioning_values[field] = value
1730 for field in extra_fields:
1731 value = literal(extra_columns[field])
1732 values.append(UNSET_VALUE if value is None else value)
1733 if field in partition_columns:
1734 partitioning_values[field] = value
1736 key = tuple(partitioning_values[field] for field in partition_columns)
1737 values_by_key[key].append(values)
1739 table = context.schema.tableName(table_name, time_part)
1741 query = Insert(self._keyspace, table, fields)
1742 statement = context.stmt_factory(query, prepare=True)
1743 # Cassandra has 64k limit on batch size, normally that should be
1744 # enough but some tests generate too many forced sources.
1745 queries = []
1746 for key_values in values_by_key.values():
1747 for values_chunk in chunk_iterable(key_values, batch_size):
1748 batch = cassandra.query.BatchStatement()
1749 for row_values in values_chunk:
1750 batch.add(statement, row_values)
1751 queries.append((batch, None))
1752 assert batch.routing_key is not None and batch.keyspace is not None
1754 _LOG.debug("%s: will store %d records", context.schema.tableName(table_name), records.shape[0])
1755 with self._timer(
1756 "insert_time", tags={"table": table_name.name, "method": "_storeObjectsPandas"}
1757 ) as timer:
1758 execute_concurrent(context.session, queries, execution_profile="write")
1759 timer.add_values(row_count=len(records), num_batches=len(queries))
1761 def _storeUpdateRecords(
1762 self, records: Iterable[ApdbUpdateRecord], chunk: ReplicaChunk, *, store_chunk: bool = False
1763 ) -> None:
1764 """Store ApdbUpdateRecords in the replica table for those records.
1766 Parameters
1767 ----------
1768 records : `list` [`ApdbUpdateRecord`]
1769 Records to store.
1770 chunk : `ReplicaChunk`
1771 Replica chunk for these records.
1772 store_chunk : `bool`
1773 If True then also store replica chunk.
1775 Raises
1776 ------
1777 TypeError
1778 Raised if replication is not enabled for this instance.
1779 """
1780 context = self._context
1781 config = context.config
1783 if not context.schema.replication_enabled:
1784 raise TypeError("Replication is not enabled for this APDB instance.")
1786 if store_chunk:
1787 self._storeReplicaChunk(chunk)
1789 apdb_replica_chunk = chunk.id
1790 # Do not use unique_if from ReplicaChunk as it could be reused in
1791 # multiple calls to this method.
1792 update_unique_id = uuid.uuid4()
1794 rows = []
1795 for record in records:
1796 rows.append(
1797 [
1798 apdb_replica_chunk,
1799 record.update_time_ns,
1800 record.update_order,
1801 update_unique_id,
1802 record.to_json(),
1803 ]
1804 )
1805 columns = [
1806 "apdb_replica_chunk",
1807 "update_time_ns",
1808 "update_order",
1809 "update_unique_id",
1810 "update_payload",
1811 ]
1812 if context.has_chunk_sub_partitions:
1813 subchunk = random.randrange(config.replica_sub_chunk_count)
1814 for row in rows:
1815 row.append(subchunk)
1816 columns.append("apdb_replica_subchunk")
1818 table_name = context.schema.tableName(ExtraTables.ApdbUpdateRecordChunks)
1819 query = Insert(self._keyspace, table_name, columns)
1820 stmt = context.stmt_factory(query)
1821 queries = [(stmt, row) for row in rows]
1823 with self._timer("store_update_record", tags={"table": table_name}) as timer:
1824 execute_concurrent(context.session, queries, execution_profile="write")
1825 timer.add_values(row_count=len(queries))
1827 def _add_apdb_part(self, df: pandas.DataFrame) -> pandas.DataFrame:
1828 """Calculate spatial partition for each record and add it to a
1829 DataFrame.
1831 Parameters
1832 ----------
1833 df : `pandas.DataFrame`
1834 DataFrame which has to contain ra/dec columns, names of these
1835 columns are defined by configuration ``ra_dec_columns`` field.
1837 Returns
1838 -------
1839 df : `pandas.DataFrame`
1840 DataFrame with ``apdb_part`` column which contains pixel index
1841 for ra/dec coordinates.
1843 Notes
1844 -----
1845 This overrides any existing column in a DataFrame with the same name
1846 (``apdb_part``). Original DataFrame is not changed, copy of a DataFrame
1847 is returned.
1848 """
1849 context = self._context
1850 config = context.config
1852 # Calculate pixelization index for every record.
1853 apdb_part = np.zeros(df.shape[0], dtype=np.int64)
1854 ra_col, dec_col = config.ra_dec_columns
1855 for i, (ra, dec) in enumerate(zip(df[ra_col], df[dec_col])):
1856 idx = context.partitioner.pixel(ra, dec)
1857 apdb_part[i] = idx
1858 df = df.copy()
1859 df["apdb_part"] = apdb_part
1860 return df
1862 def _make_empty_catalog(self, table_name: ApdbTables) -> pandas.DataFrame:
1863 """Make an empty catalog for a table with a given name.
1865 Parameters
1866 ----------
1867 table_name : `ApdbTables`
1868 Name of the table.
1870 Returns
1871 -------
1872 catalog : `pandas.DataFrame`
1873 An empty catalog.
1874 """
1875 table = self.schema.tableSchemas[table_name]
1877 data = {columnDef.name: pandas.Series(dtype=columnDef.pandas_type) for columnDef in table.columns}
1878 return pandas.DataFrame(data)
1880 def _fix_input_timestamps(self, df: pandas.DataFrame) -> pandas.DataFrame:
1881 """Update timestamp columns in input DataFrame to be naive datetime
1882 type.
1884 Clients may or may not generate aware timestamps, code in this class
1885 assumes that timestamps are naive, so we convert them to UTC and
1886 drop timezone.
1887 """
1888 # Find all columns with aware timestamps.
1889 columns = [column for column, dtype in df.dtypes.items() if isinstance(dtype, pandas.DatetimeTZDtype)]
1890 for column in columns:
1891 # tz_convert(None) will convert to UTC and drop timezone.
1892 df[column] = df[column].dt.tz_convert(None)
1893 return df
1895 def _batch_size(self, table: ApdbTables | ExtraTables) -> int:
1896 """Calculate batch size based on config parameters."""
1897 context = self._context
1898 config = context.config
1900 # Cassandra limit on number of statements in a batch is 64k.
1901 batch_size = 65_535
1902 if 0 < config.batch_statement_limit < batch_size:
1903 batch_size = config.batch_statement_limit
1904 if config.batch_size_limit > 0:
1905 # The purpose of this limit is to try not to exceed batch size
1906 # threshold which is set on server side. Cassandra wire protocol
1907 # for prepared queries (and batches) only sends column values with
1908 # with an additional 4 bytes per value specifying size. Value is
1909 # not included for NULL or NOT_SET values, but the size is always
1910 # there. There is additional small per-query overhead, which we
1911 # ignore.
1912 row_size = context.schema.table_row_size(table)
1913 row_size += 4 * len(context.schema.getColumnMap(table))
1914 batch_size = min(batch_size, (config.batch_size_limit // row_size) + 1)
1915 return batch_size
1917 def _group_dia_objects_by_partition(
1918 self, partitioner: Partitioner, objects: list[DiaObjectId], pad_arcsec: float
1919 ) -> Mapping[int, list[int]]:
1920 """Group DiaObjects by partition.
1922 Parameters
1923 ----------
1924 partitioner : `Partitioner`
1925 Objects which knows how to partition things.
1926 objects : `list` [`DiaObjectId`]
1927 Collection of objects to partition.
1928 pad_arcsec : `float`
1929 Additional padding around object position.
1931 Returns
1932 -------
1933 grouped_objects
1934 Mapping of spatial patition ID to list ob object IDs that it
1935 contains. Some objects may belong to more than one partition.
1936 """
1937 partitioned_object_ids: dict[int, list[int]] = defaultdict(list)
1938 for obj_id in objects:
1939 partitions = partitioner.pixelization.circle_pixels(obj_id.ra, obj_id.dec, pad_arcsec)
1940 for pixel in partitions:
1941 partitioned_object_ids[pixel].append(obj_id.diaObjectId)
1942 return partitioned_object_ids
1944 def _timestamp_column_name(self, column: str) -> str:
1945 """Return column name before/after schema migration to MJD TAI."""
1946 return self._schema.timestamp_column_name(column)
1948 def _timestamp_column_value(self, time: astropy.time.Time) -> float | int:
1949 """Return column value before/after schema migration to MJD TAI."""
1950 if self._schema.has_mjd_timestamps:
1951 return float(time.tai.mjd)
1952 else:
1953 return int(time.datetime.astimezone(tz=datetime.UTC).timestamp() * 1000)
1955 def _get_diasource_data(self, source_ids: Iterable[DiaSourceId], *columns: str) -> list:
1956 """Select records from DiaSource table by diaSourceId and return all
1957 records as a list of named tuples.
1958 """
1959 context = self._context
1960 config = context.config
1961 partitioner = context.partitioner
1963 columns = ("diaSourceId",) + columns
1965 # Allow some uncertainty for coordinates and time when calculating
1966 # partitions.
1967 statements: list[tuple] = []
1968 pad_arcsec = 1.0
1969 pad_time_day = 10 / (24 * 3600)
1970 for source_id in source_ids:
1971 center = sphgeom.UnitVector3d(sphgeom.LonLat.fromDegrees(source_id.ra, source_id.dec))
1972 region = sphgeom.Circle(center, sphgeom.Angle.fromDegrees(pad_arcsec / 3600.0))
1973 spatial_where, _ = partitioner.spatial_where(region)
1975 tables, temporal_where = partitioner.temporal_where(
1976 ApdbTables.DiaSource,
1977 source_id.midpointMjdTai - pad_time_day,
1978 source_id.midpointMjdTai + pad_time_day,
1979 partitons_range=context.time_partitions_range,
1980 query_per_time_part=True,
1981 )
1983 id_where = QExpr('"diaSourceId" = {}', (source_id.diaSourceId,))
1985 for table in tables:
1986 query = Select(self._keyspace, table, columns)
1987 for clause in QExpr.combine(spatial_where, temporal_where, extra=id_where):
1988 statements.append(context.stmt_factory.with_params(query.where(clause), prepare=True))
1990 with self._timer(
1991 "select_time", tags={"table": "DiaSource", "method": "_get_diasource_data"}
1992 ) as timer:
1993 result = cast(
1994 list[tuple],
1995 select_concurrent(
1996 context.session,
1997 statements,
1998 "read_named_tuples",
1999 config.connection_config.read_concurrency,
2000 ),
2001 )
2002 timer.add_values(row_count=len(result), num_queries=len(statements))
2004 return result
2006 def _get_diaforcedsource_data(self, source_ids: Iterable[DiaForcedSourceId], *columns: str) -> list:
2007 """Select records from DiaForcedSource table by (diaObjectId, visit,
2008 detector) and return all records as a list of named tuples.
2009 """
2010 context = self._context
2011 config = context.config
2012 partitioner = context.partitioner
2014 columns = ("diaObjectId", "visit", "detector") + columns
2016 # Allow some uncertainty for coordinates and time when calculating
2017 # partitions.
2018 statements: list[tuple] = []
2019 pad_arcsec = 1.0
2020 pad_time_day = 10 / (24 * 3600)
2021 for source_id in source_ids:
2022 center = sphgeom.UnitVector3d(sphgeom.LonLat.fromDegrees(source_id.ra, source_id.dec))
2023 region = sphgeom.Circle(center, sphgeom.Angle.fromDegrees(pad_arcsec / 3600.0))
2024 spatial_where, _ = partitioner.spatial_where(region)
2026 tables, temporal_where = partitioner.temporal_where(
2027 ApdbTables.DiaForcedSource,
2028 source_id.midpointMjdTai - pad_time_day,
2029 source_id.midpointMjdTai + pad_time_day,
2030 partitons_range=context.time_partitions_range,
2031 query_per_time_part=True,
2032 )
2034 id_where = (
2035 (C("diaObjectId") == source_id.diaObjectId)
2036 & (C("visit") == source_id.visit)
2037 & (C("detector") == source_id.detector)
2038 )
2040 for table in tables:
2041 query = Select(self._keyspace, table, columns)
2042 for clause in QExpr.combine(spatial_where, temporal_where, extra=id_where):
2043 statements.append(context.stmt_factory.with_params(query.where(clause), prepare=True))
2045 with self._timer(
2046 "select_time", tags={"table": "DiaForcedSource", "method": "_get_diaforcedsource_data"}
2047 ) as timer:
2048 result = cast(
2049 list[tuple],
2050 select_concurrent(
2051 context.session,
2052 statements,
2053 "read_named_tuples",
2054 config.connection_config.read_concurrency,
2055 ),
2056 )
2057 timer.add_values(row_count=len(result), num_queries=len(statements))
2059 return result
2061 def _get_diaobject_data(self, object_ids: Iterable[DiaObjectId], *columns: str) -> list:
2062 """Select records from DiaObjectLast table by diaObjectId and return
2063 all records as a list of named tuples.
2064 """
2065 context = self._context
2066 config = context.config
2067 partitioner = context.partitioner
2069 table_name = context.schema.tableName(ApdbTables.DiaObjectLast)
2070 columns = ("diaObjectId",) + columns
2072 # Allow some uncertainty for coordinates when calculating partitions.
2073 pad_arcsec = 1.0
2074 ids_by_partition = defaultdict(list)
2075 for object_id in object_ids:
2076 pixels = partitioner.pixelization.circle_pixels(object_id.ra, object_id.dec, pad_arcsec)
2077 for pixel in pixels:
2078 ids_by_partition[pixel].append(object_id.diaObjectId)
2080 statements: list[tuple] = []
2081 for apdb_part, diaObjectIds in ids_by_partition.items():
2082 query = Select(self._keyspace, table_name, columns)
2083 query = query.where(C("apdb_part") == apdb_part)
2084 query = query.where(C("diaObjectId").in_(diaObjectIds))
2085 statements.append(context.stmt_factory.with_params(query, prepare=False))
2087 with self._timer(
2088 "select_time", tags={"table": "DiaObjectLast", "method": "_get_diaobject_data"}
2089 ) as timer:
2090 result = cast(
2091 list[tuple],
2092 select_concurrent(
2093 context.session,
2094 statements,
2095 "read_named_tuples",
2096 config.connection_config.read_concurrency,
2097 ),
2098 )
2099 timer.add_values(row_count=len(result), num_queries=len(statements))
2101 return result