Coverage for python/lsst/dax/apdb/cassandra/apdbCassandra.py: 8%

912 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-26 09:10 +0000

1# This file is part of dax_apdb. 

2# 

3# Developed for the LSST Data Management System. 

4# This product includes software developed by the LSST Project 

5# (http://www.lsst.org). 

6# See the COPYRIGHT file at the top-level directory of this distribution 

7# for details of code ownership. 

8# 

9# This program is free software: you can redistribute it and/or modify 

10# it under the terms of the GNU General Public License as published by 

11# the Free Software Foundation, either version 3 of the License, or 

12# (at your option) any later version. 

13# 

14# This program is distributed in the hope that it will be useful, 

15# but WITHOUT ANY WARRANTY; without even the implied warranty of 

16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the 

17# GNU General Public License for more details. 

18# 

19# You should have received a copy of the GNU General Public License 

20# along with this program. If not, see <http://www.gnu.org/licenses/>. 

21 

22from __future__ import annotations 

23 

24__all__ = ["ApdbCassandra"] 

25 

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 

35 

36import numpy as np 

37import pandas 

38 

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 

45 

46 CASSANDRA_IMPORTED = True 

47except ImportError: 

48 CASSANDRA_IMPORTED = False 

49 

50import astropy.time 

51import felis.datamodel 

52 

53from lsst import sphgeom 

54from lsst.utils.iteration import chunk_iterable 

55 

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 

89 

90if TYPE_CHECKING: 

91 from ..apdbMetadata import ApdbMetadata 

92 from ..apdbUpdateRecord import ApdbUpdateRecord 

93 

94_LOG = logging.getLogger(__name__) 

95 

96_MON = MonAgent(__name__) 

97 

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

103 

104 

105class ApdbCassandra(Apdb): 

106 """Implementation of APDB database with Apache Cassandra backend. 

107 

108 Parameters 

109 ---------- 

110 config : `ApdbCassandraConfig` 

111 Configuration object. 

112 """ 

113 

114 def __init__(self, config: ApdbCassandraConfig): 

115 if not CASSANDRA_IMPORTED: 

116 raise CassandraMissingError() 

117 

118 self._config = config 

119 self._keyspace = config.keyspace 

120 self._schema = ApdbSchema(config.schema_file, config.ss_schema_file) 

121 

122 self._session_factory = SessionFactory(config) 

123 self._connection_context: ConnectionContext | None = None 

124 

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) 

135 

136 session = self._session_factory.session() 

137 self._connection_context = ConnectionContext( 

138 session, self._config, self.schema.tableSchemas, current_versions 

139 ) 

140 

141 if _LOG.isEnabledFor(logging.DEBUG): 

142 _LOG.debug("ApdbCassandra Configuration: %s", self._connection_context.config.model_dump()) 

143 

144 return self._connection_context 

145 

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) 

149 

150 @classmethod 

151 def apdbImplementationVersion(cls) -> VersionTuple: 

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

153 

154 Returns 

155 ------- 

156 version : `VersionTuple` 

157 Version of the code defined in implementation class. 

158 """ 

159 return VERSION 

160 

161 def getConfig(self) -> ApdbCassandraConfig: 

162 # docstring is inherited from a base class 

163 return self._context.config 

164 

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

166 # docstring is inherited from a base class 

167 return self.schema.tableSchemas.get(table) 

168 

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. 

200 

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. 

265 

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 

325 

326 cls._makeSchema(config, drop=drop, replication_factor=replication_factor, table_options=table_options) 

327 

328 return config 

329 

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) 

335 

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 

346 

347 if not isinstance(config, ApdbCassandraConfig): 

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

349 

350 simple_schema = ApdbSchema(config.schema_file, config.ss_schema_file) 

351 

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 ) 

362 

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 ) 

387 

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

389 metadata = ApdbMetadataCassandra( 

390 session, meta_table_name, config.keyspace, "read_tuples", "write" 

391 ) 

392 

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 ) 

400 

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 ) 

408 

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) 

412 

413 # Store time partition range. 

414 if part_range_config: 

415 part_range_config.save_to_meta(metadata) 

416 

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 

421 

422 sp_where, num_sp_part = context.partitioner.spatial_where(region) 

423 _LOG.debug("getDiaObjects: #partitions: %s", len(sp_where)) 

424 

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

434 

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

447 

448 _LOG.debug("found %s DiaObjects", objects.shape[0]) 

449 return objects 

450 

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 

461 

462 months = config.read_sources_months 

463 if start_time is None and months == 0: 

464 return None 

465 

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) 

471 

472 return self._getSources(region, object_ids, mjd_start, mjd_end, ApdbTables.DiaSource) 

473 

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 

484 

485 months = config.read_forced_sources_months 

486 if start_time is None and months == 0: 

487 return None 

488 

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) 

494 

495 return self._getSources(region, object_ids, mjd_start, mjd_end, ApdbTables.DiaForcedSource) 

496 

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 

501 

502 if not context.has_dedup_table: 

503 raise TypeError("DiaObjectDedup table does not exist in this APDB instance.") 

504 

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

512 

513 column_names = context.schema.apdbColumnNames(ExtraTables.DiaObjectDedup) 

514 

515 validity_start_column = self._timestamp_column_name("validityStart") 

516 timestamp = None if since is None else self._timestamp_column_value(since) 

517 

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) 

523 

524 statement = context.stmt_factory(query, prepare=False) 

525 

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

531 

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) 

546 

547 _LOG.debug("found %s DiaObjectDedup records", objects.shape[0]) 

548 return objects 

549 

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 

556 

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 ) 

571 

572 # Group DiaObjects by partition. 

573 partitioned_object_ids = self._group_dia_objects_by_partition( 

574 context.partitioner, objects, max_dist_arcsec 

575 ) 

576 

577 # Columns to return. 

578 column_names = context.schema.apdbColumnNames(ApdbTables.DiaSource) 

579 

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 ) 

592 

593 _LOG.debug("getDiaSourcesForDiaObjects #queries: %s", len(statements)) 

594 

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

609 

610 # precise filtering on midpointMjdTai 

611 catalog = cast(pandas.DataFrame, catalog[catalog["midpointMjdTai"] >= start_time.tai.mjd]) 

612 

613 timer.add_values(row_count=len(catalog)) 

614 

615 _LOG.debug("found %d DiaSources", len(catalog)) 

616 return catalog 

617 

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 

627 

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) 

632 

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

636 

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 

647 

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

657 

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

666 

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) 

672 

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) 

677 

678 # fill region partition column for DiaObjects 

679 objects = self._add_apdb_part(objects) 

680 self._storeDiaObjects(objects, visit_time, replica_chunk) 

681 

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) 

687 

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) 

691 

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 

702 

703 source_ids = {source_id.diaSourceId for source_id in idMap} 

704 

705 # Find all DiaSources. 

706 found_sources = self._get_diasource_data( 

707 idMap, "apdb_part", "diaObjectId", "ra", "dec", "midpointMjdTai" 

708 ) 

709 

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

712 

713 found_sources_by_id = {row.diaSourceId: row for row in found_sources} 

714 

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

728 

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

733 

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) 

738 

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) 

745 

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

764 

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 

778 

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

784 

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 = [] 

798 

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) 

805 

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

810 

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 

824 

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

831 

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) 

835 

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 

842 

843 context = self._context 

844 config = context.config 

845 

846 pad_arcsec = 1.0 

847 partitioned_object_ids = self._group_dia_objects_by_partition( 

848 context.partitioner, objects, pad_arcsec 

849 ) 

850 

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

859 

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

871 

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

879 

880 # Filter existing records. 

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

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

883 

884 if not objects: 

885 return 0 

886 

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) 

891 

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

901 

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

908 

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

912 

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) 

918 

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 ] 

931 

932 self._storeUpdateRecords(update_records, replica_chunk, store_chunk=True) 

933 

934 return len(objects) 

935 

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

937 # docstring is inherited from a base class 

938 context = self._context 

939 

940 if not context.has_dedup_table: 

941 raise TypeError("DiaObjectDedup table does not exist in this APDB instance.") 

942 

943 if dedup_time is None: 

944 dedup_time = self._current_time() 

945 

946 validity_start_column = self._timestamp_column_name("validityStart") 

947 

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) 

952 

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

959 

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

969 

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) 

974 

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 

979 

980 now = self._current_time() 

981 reassign_time_column = self._timestamp_column_name("ssObjectReassocTime") 

982 reassignTime = self._timestamp_column_value(now) 

983 

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. 

987 

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

996 

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 ) 

1004 

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] 

1012 

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

1017 

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

1048 

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) 

1052 

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

1057 

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 

1067 

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

1072 

1073 diaSourceIds = list(diaSourceIds) 

1074 source_ids = {source_id.diaSourceId for source_id in diaSourceIds} 

1075 

1076 # Find all DiaSources. 

1077 found_sources = self._get_diasource_data( 

1078 diaSourceIds, "apdb_part", "diaObjectId", "ra", "dec", "midpointMjdTai", column_name 

1079 ) 

1080 

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

1083 

1084 found_sources_by_id = {row.diaSourceId: row for row in found_sources} 

1085 

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) 

1090 

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 

1098 

1099 apdb_part = source_row.apdb_part 

1100 time_part = context.partitioner.time_partition(source_row.midpointMjdTai) 

1101 

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

1120 

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 

1134 

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

1138 

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) 

1142 

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 

1152 

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

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

1155 

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

1160 

1161 diaForcedSourceIds = list(diaForcedSourceIds) 

1162 fsource_keys = {_fsrc_id(source) for source in diaForcedSourceIds} 

1163 

1164 found_fsources = self._get_diaforcedsource_data( 

1165 diaForcedSourceIds, "apdb_part", "ra", "dec", "midpointMjdTai", column_name 

1166 ) 

1167 

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

1171 

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) 

1177 

1178 for source_row in found_fsources: 

1179 # Ignore sources already withdrawn. 

1180 if getattr(source_row, column_name) is not None: 

1181 continue 

1182 

1183 apdb_part = source_row.apdb_part 

1184 time_part = context.partitioner.time_partition(source_row.midpointMjdTai) 

1185 

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

1208 

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 

1224 

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

1231 

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) 

1235 

1236 def countUnassociatedObjects(self) -> int: 

1237 # docstring is inherited from a base class 

1238 

1239 # It's too inefficient to implement it for Cassandra in current schema. 

1240 raise NotImplementedError() 

1241 

1242 @property 

1243 def schema(self) -> ApdbSchema: 

1244 # docstring is inherited from a base class 

1245 return self._schema 

1246 

1247 @property 

1248 def metadata(self) -> ApdbMetadata: 

1249 # docstring is inherited from a base class 

1250 context = self._context 

1251 return context.metadata 

1252 

1253 @property 

1254 def admin(self) -> ApdbCassandraAdmin: 

1255 # docstring is inherited from a base class 

1256 return ApdbCassandraAdmin(self) 

1257 

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. 

1267 

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. 

1280 

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 

1289 

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) 

1295 

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 ) 

1306 

1307 # We need to exclude extra partitioning columns from result. 

1308 column_names = context.schema.apdbColumnNames(table_name) 

1309 

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

1317 

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 ) 

1332 

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

1336 

1337 # precise filtering on midpointMjdTai 

1338 catalog = cast(pandas.DataFrame, catalog[catalog["midpointMjdTai"] > mjd_start]) 

1339 

1340 timer.add_values(row_count=len(catalog)) 

1341 

1342 _LOG.debug("found %d %ss", catalog.shape[0], table_name.name) 

1343 return catalog 

1344 

1345 def _storeReplicaChunk(self, replica_chunk: ReplicaChunk) -> None: 

1346 context = self._context 

1347 config = context.config 

1348 

1349 # Cassandra timestamp uses milliseconds since epoch 

1350 timestamp = int(replica_chunk.last_update_time.unix_tai * 1000) 

1351 

1352 # everything goes into a single partition 

1353 partition = 0 

1354 

1355 table_name = context.schema.tableName(ExtraTables.ApdbReplicaChunks) 

1356 

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) 

1362 

1363 query = Insert(self._keyspace, table_name, columns) 

1364 stmt = context.stmt_factory(query) 

1365 

1366 context.session.execute( 

1367 stmt, 

1368 values, 

1369 timeout=config.connection_config.write_timeout, 

1370 execution_profile="write", 

1371 ) 

1372 

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 

1377 

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) 

1387 

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

1399 

1400 if data.column_names() != ["diaObjectId", "apdb_part"]: 

1401 raise RuntimeError(f"Unexpected column names in query result: {data.column_names()}") 

1402 

1403 return {row[0]: row[1] for row in data.rows()} 

1404 

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 

1412 

1413 # Extract all object IDs. 

1414 new_partitions = dict(zip(objs["diaObjectId"], objs["apdb_part"])) 

1415 old_partitions = self._queryDiaObjectLastPartitions(objs["diaObjectId"]) 

1416 

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) 

1423 

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

1436 

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) 

1441 

1442 queries = [] 

1443 for oid, new_part in new_partitions.items(): 

1444 queries.append((statement, (oid, new_part))) 

1445 

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

1449 

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. 

1454 

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 

1467 

1468 context = self._context 

1469 config = context.config 

1470 

1471 self._deleteMovingObjects(objs) 

1472 

1473 validity_start_column = self._timestamp_column_name("validityStart") 

1474 timestamp = self._timestamp_column_value(visit_time) 

1475 

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 

1480 

1481 self._storeObjectsPandas(objs, ApdbTables.DiaObjectLast, extra_columns=extra_columns) 

1482 

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 

1491 

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 ) 

1498 

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) 

1509 

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) 

1518 

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. 

1526 

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. 

1537 

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 

1546 

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) 

1568 

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) 

1585 

1586 return subchunk 

1587 

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. 

1592 

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 ) 

1615 

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. 

1624 

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 

1638 

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 

1646 

1647 self._storeObjectsPandas( 

1648 id_map, ExtraTables.DiaSourceToPartition, extra_columns=extra_columns, time_part=None 

1649 ) 

1650 

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. 

1659 

1660 Takes Pandas catalog and stores a bunch of records in a table. 

1661 

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. 

1674 

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 

1683 

1684 # use extra columns if specified 

1685 if extra_columns is None: 

1686 extra_columns = {} 

1687 extra_fields = list(extra_columns.keys()) 

1688 

1689 # Fields that will come from dataframe. 

1690 df_fields = [column for column in records.columns if column not in extra_fields] 

1691 

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 

1696 

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

1703 

1704 batch_size = self._batch_size(table_name) 

1705 

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 

1735 

1736 key = tuple(partitioning_values[field] for field in partition_columns) 

1737 values_by_key[key].append(values) 

1738 

1739 table = context.schema.tableName(table_name, time_part) 

1740 

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 

1753 

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

1760 

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. 

1765 

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. 

1774 

1775 Raises 

1776 ------ 

1777 TypeError 

1778 Raised if replication is not enabled for this instance. 

1779 """ 

1780 context = self._context 

1781 config = context.config 

1782 

1783 if not context.schema.replication_enabled: 

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

1785 

1786 if store_chunk: 

1787 self._storeReplicaChunk(chunk) 

1788 

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

1793 

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

1817 

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] 

1822 

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

1826 

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. 

1830 

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. 

1836 

1837 Returns 

1838 ------- 

1839 df : `pandas.DataFrame` 

1840 DataFrame with ``apdb_part`` column which contains pixel index 

1841 for ra/dec coordinates. 

1842 

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 

1851 

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 

1861 

1862 def _make_empty_catalog(self, table_name: ApdbTables) -> pandas.DataFrame: 

1863 """Make an empty catalog for a table with a given name. 

1864 

1865 Parameters 

1866 ---------- 

1867 table_name : `ApdbTables` 

1868 Name of the table. 

1869 

1870 Returns 

1871 ------- 

1872 catalog : `pandas.DataFrame` 

1873 An empty catalog. 

1874 """ 

1875 table = self.schema.tableSchemas[table_name] 

1876 

1877 data = {columnDef.name: pandas.Series(dtype=columnDef.pandas_type) for columnDef in table.columns} 

1878 return pandas.DataFrame(data) 

1879 

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

1881 """Update timestamp columns in input DataFrame to be naive datetime 

1882 type. 

1883 

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 

1894 

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 

1899 

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 

1916 

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. 

1921 

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. 

1930 

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 

1943 

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) 

1947 

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) 

1954 

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 

1962 

1963 columns = ("diaSourceId",) + columns 

1964 

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) 

1974 

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 ) 

1982 

1983 id_where = QExpr('"diaSourceId" = {}', (source_id.diaSourceId,)) 

1984 

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

1989 

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

2003 

2004 return result 

2005 

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 

2013 

2014 columns = ("diaObjectId", "visit", "detector") + columns 

2015 

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) 

2025 

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 ) 

2033 

2034 id_where = ( 

2035 (C("diaObjectId") == source_id.diaObjectId) 

2036 & (C("visit") == source_id.visit) 

2037 & (C("detector") == source_id.detector) 

2038 ) 

2039 

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

2044 

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

2058 

2059 return result 

2060 

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 

2068 

2069 table_name = context.schema.tableName(ApdbTables.DiaObjectLast) 

2070 columns = ("diaObjectId",) + columns 

2071 

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) 

2079 

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

2086 

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

2100 

2101 return result