Coverage for python/lsst/dax/apdb/cassandra/apdbCassandraAdmin.py: 13%

273 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-25 22:15 +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__ = ["ApdbCassandraAdmin"] 

25 

26import dataclasses 

27import itertools 

28import logging 

29import warnings 

30from collections import defaultdict 

31from collections.abc import Iterable, Mapping 

32from typing import TYPE_CHECKING, Protocol 

33 

34import astropy.time 

35 

36from lsst.utils.iteration import chunk_iterable 

37 

38try: 

39 import cassandra 

40except ImportError: 

41 pass 

42 

43from ..apdbAdmin import ApdbAdmin, DiaForcedSourceLocator, DiaObjectLocator, DiaSourceLocator 

44from ..apdbSchema import ApdbTables 

45from ..monitor import MonAgent 

46from ..timer import Timer 

47from .cassandra_utils import StatementFactory, execute_concurrent, quote_id 

48from .config import ApdbCassandraConfig, ApdbCassandraTimePartitionRange 

49from .queries import Column as C # noqa: N817 

50from .queries import Delete, Select 

51from .sessionFactory import SessionContext 

52 

53if TYPE_CHECKING: 

54 from .apdbCassandra import ApdbCassandra 

55 from .partitioner import Partitioner 

56 

57_LOG = logging.getLogger(__name__) 

58 

59_MON = MonAgent(__name__) 

60 

61 

62class ConfirmDeletePartitions(Protocol): 

63 """Protocol for callable which confirms deletion of partitions.""" 

64 

65 def __call__(self, *, partitions: list[int], tables: list[str], partitioner: Partitioner) -> bool: ... 65 ↛ exitline 65 didn't return from function '__call__' because

66 

67 

68@dataclasses.dataclass 

69class DatabaseInfo: 

70 """Collection of information about a specific database.""" 

71 

72 name: str 

73 """Keyspace name.""" 

74 

75 permissions: dict[str, set[str]] | None = None 

76 """Roles that can access the database and their permissions. 

77 

78 `None` means that authentication information is not accessible due to 

79 system table permissions. If anonymous access is enabled then dictionary 

80 will be empty but not `None`. 

81 """ 

82 

83 

84class ApdbCassandraAdmin(ApdbAdmin): 

85 """Implementation of `ApdbAdmin` for Cassandra backend. 

86 

87 Parameters 

88 ---------- 

89 apdb : `ApdbCassandra` 

90 APDB implementation. 

91 """ 

92 

93 def __init__(self, apdb: ApdbCassandra): 

94 self._apdb = apdb 

95 

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

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

98 return Timer(name, _MON, _LOG, tags=tags) 

99 

100 @classmethod 

101 def list_databases(cls, host: str) -> Iterable[DatabaseInfo]: 

102 """Return the list of keyspaces with APDB databases. 

103 

104 Parameters 

105 ---------- 

106 host : `str` 

107 Name of one of the hosts in Cassandra cluster. 

108 

109 Returns 

110 ------- 

111 databases : `~collections.abc.Iterable` [`DatabaseInfo`] 

112 Information about databases that contain APDB instance. 

113 """ 

114 # For DbAuth we need to use database name "*" to try to match any 

115 # database. 

116 config = ApdbCassandraConfig(contact_points=(host,), keyspace="*") 

117 with SessionContext(config) as session: 

118 stmt_factory = StatementFactory(session) 

119 

120 # Get names of all keyspaces containing DiaSource table 

121 table_name = ApdbTables.DiaSource.table_name() 

122 query = Select("system_schema", "tables", ["keyspace_name"], extra_clause="ALLOW FILTERING") 

123 query = query.where(C("table_name") == table_name) 

124 stmt, params = stmt_factory.with_params(query) 

125 

126 result = session.execute(stmt, params) 

127 keyspaces = [row[0] for row in result.all()] 

128 

129 if not keyspaces: 

130 return [] 

131 

132 # Retrieve roles for each keyspace. 

133 resources = [f"data/{keyspace}" for keyspace in keyspaces] 

134 query = Select( 

135 "system_auth", 

136 "role_permissions", 

137 ("resource", "role", "permissions"), 

138 extra_clause="ALLOW FILTERING", 

139 ) 

140 query = query.where(C("resource").in_(resources)) 

141 stmt, params = stmt_factory.with_params(query) 

142 

143 try: 

144 result = session.execute(stmt, params) 

145 # If anonymous access is enabled then result will be empty, 

146 # set infos to have empty permissions dict in that case. 

147 infos = {keyspace: DatabaseInfo(name=keyspace, permissions={}) for keyspace in keyspaces} 

148 for row in result: 

149 _, _, keyspace = row[0].partition("/") 

150 role: str = row[1] 

151 role_permissions: set[str] = set(row[2]) 

152 infos[keyspace].permissions[role] = role_permissions # type: ignore[index] 

153 except cassandra.Unauthorized as exc: 

154 # Likely that access to role_permissions is not granted for 

155 # current user. 

156 warnings.warn( 

157 f"Authentication information is not accessible to current user - {exc}", stacklevel=2 

158 ) 

159 infos = {keyspace: DatabaseInfo(name=keyspace) for keyspace in keyspaces} 

160 

161 # Would be nice to get size estimate, but this is not available 

162 # via CQL queries. 

163 return infos.values() 

164 

165 @classmethod 

166 def delete_database(cls, host: str, keyspace: str, *, timeout: int = 3600) -> None: 

167 """Delete APDB database by dropping its keyspace. 

168 

169 Parameters 

170 ---------- 

171 host : `str` 

172 Name of one of the hosts in Cassandra cluster. 

173 keyspace : `str` 

174 Name of keyspace to delete. 

175 timeout : `int`, optional 

176 Timeout for delete operation in seconds. Dropping a large keyspace 

177 can be a long operation, but this default value of one hour should 

178 be sufficient for most or all cases. 

179 """ 

180 # For DbAuth we need to use database name "*" to try to match any 

181 # database. 

182 config = ApdbCassandraConfig(contact_points=(host,), keyspace="*") 

183 with SessionContext(config) as session: 

184 query = f"DROP KEYSPACE {quote_id(keyspace)}" 

185 session.execute(query, timeout=timeout) 

186 

187 @property 

188 def partitioner(self) -> Partitioner: 

189 """Partitoner used by this APDB instance (`Partitioner`).""" 

190 context = self._apdb._context 

191 return context.partitioner 

192 

193 def apdb_part(self, ra: float, dec: float) -> int: 

194 # docstring is inherited from a base class 

195 return self.partitioner.pixel(ra, dec) 

196 

197 def apdb_time_part(self, midpointMjdTai: float) -> int: 

198 # docstring is inherited from a base class 

199 return self.partitioner.time_partition(midpointMjdTai) 

200 

201 def delete_records( 

202 self, 

203 objects: Iterable[DiaObjectLocator], 

204 sources: Iterable[DiaSourceLocator], 

205 forced_sources: Iterable[DiaForcedSourceLocator], 

206 ) -> None: 

207 # docstring is inherited from a base class 

208 context = self._apdb._context 

209 config = context.config 

210 keyspace = self._apdb._keyspace 

211 has_dia_object_table = not (config.enable_replica and config.replica_skips_diaobjects) 

212 

213 # Group objects by partition. 

214 partitions = defaultdict(list) 

215 for object in objects: 

216 apdb_part = self.apdb_part(object.ra, object.dec) 

217 partitions[apdb_part].append(object.diaObjectId) 

218 object_ids = set(itertools.chain.from_iterable(partitions.values())) 

219 

220 # Group sources by associated object ID. 

221 source_groups = defaultdict(list) 

222 for source in sources: 

223 if source.diaObjectId in object_ids: 

224 source_groups[source.diaObjectId].append(source) 

225 

226 object_deletes = [] 

227 object_count = 0 

228 # Delete from DiaObjectLast table. 

229 for apdb_part, oids in partitions.items(): 

230 oids = sorted(oids) 

231 object_count += len(oids) 

232 for oid_chunk in chunk_iterable(oids, 1000): 

233 query = ( 

234 Delete(keyspace, "DiaObjectLast") 

235 .where(C("apdb_part") == apdb_part) 

236 .where(C("diaObjectId").in_(oid_chunk)) 

237 ) 

238 object_deletes.append(context.stmt_factory.with_params(query)) 

239 

240 # If DiaObject is in use then delete from that too. 

241 if has_dia_object_table: 

242 # Need temporal partitions for DiaObject, the only source for that 

243 # is the timestamp of the associated DiaSource. Problem here is 

244 # that DiaObject temporal partitioning is based on validityStart, 

245 # which is "visit_time"", but DiaSource does not record visit_time, 

246 # it is partitioned on midpointMjdTai. There is time_processed 

247 # defined for DiaSource but it does not match "visit_time" though 

248 # it is close. I use midpointMjdTai as approximation for 

249 # validityStart, this may skip some DiaObjects, but in production 

250 # we are not going to have DiaObjects table at all. There is also 

251 # a chance that DiaObject moves from one spatial partition to 

252 # another with the same consequences, which we also ignore. 

253 oids_by_partition: dict[tuple[int, int], list[int]] = defaultdict(list) 

254 for apdb_part, oids in partitions.items(): 

255 for oid in oids: 

256 temporal_partitions = { 

257 self.apdb_time_part(src.midpointMjdTai) for src in source_groups.get(oid, []) 

258 } 

259 for time_part in temporal_partitions: 

260 oids_by_partition[(apdb_part, time_part)].append(oid) 

261 for (apdb_part, time_part), oids in oids_by_partition.items(): 

262 for oid_chunk in chunk_iterable(oids, 1000): 

263 if config.partitioning.time_partition_tables: 

264 table_name = context.schema.tableName(ApdbTables.DiaObject, time_part) 

265 query = ( 

266 Delete(keyspace, table_name) 

267 .where(C("apdb_part") == apdb_part) 

268 .where(C("diaObjectId").in_(oid_chunk)) 

269 ) 

270 object_deletes.append(context.stmt_factory.with_params(query)) 

271 else: 

272 table_name = context.schema.tableName(ApdbTables.DiaObject) 

273 query = ( 

274 Delete(keyspace, table_name) 

275 .where(C("apdb_part") == apdb_part) 

276 .where(C("apdb_time_part") == time_part) 

277 .where(C("diaObjectId").in_(oid_chunk)) 

278 ) 

279 object_deletes.append(context.stmt_factory.with_params(query)) 

280 

281 # Delete from DiaObjectLastToPartition table. 

282 for oid_chunk in chunk_iterable(sorted(object_ids), 1000): 

283 query = Delete(keyspace, "DiaObjectLastToPartition").where(C("diaObjectId").in_(oid_chunk)) 

284 object_deletes.append(context.stmt_factory.with_params(query)) 

285 

286 # Group sources by partition. 

287 source_partitions = defaultdict(list) 

288 for source in itertools.chain.from_iterable(source_groups.values()): 

289 apdb_part = self.apdb_part(source.ra, source.dec) 

290 apdb_time_part = self.apdb_time_part(source.midpointMjdTai) 

291 source_partitions[(apdb_part, apdb_time_part)].append(source) 

292 

293 source_deletes = [] 

294 source_count = 0 

295 for (apdb_part, apdb_time_part), source_list in source_partitions.items(): 

296 source_ids = sorted(source.diaSourceId for source in source_list) 

297 source_count += len(source_ids) 

298 for id_chunk in chunk_iterable(source_ids, 1000): 

299 if config.partitioning.time_partition_tables: 

300 table_name = context.schema.tableName(ApdbTables.DiaSource, apdb_time_part) 

301 query = ( 

302 Delete(keyspace, table_name) 

303 .where(C("apdb_part") == apdb_part) 

304 .where(C("diaSourceId").in_(id_chunk)) 

305 ) 

306 source_deletes.append(context.stmt_factory.with_params(query)) 

307 else: 

308 table_name = context.schema.tableName(ApdbTables.DiaSource) 

309 query = ( 

310 Delete(keyspace, table_name) 

311 .where(C("apdb_part") == apdb_part) 

312 .where(C("apdb_time_part") == apdb_time_part) 

313 .where(C("diaSourceId").in_(id_chunk)) 

314 ) 

315 source_deletes.append(context.stmt_factory.with_params(query)) 

316 

317 # Group forced sources by partition. 

318 forced_source_partitions = defaultdict(list) 

319 for forced_source in forced_sources: 

320 if forced_source.diaObjectId in object_ids: 

321 apdb_part = self.apdb_part(forced_source.ra, forced_source.dec) 

322 apdb_time_part = self.apdb_time_part(forced_source.midpointMjdTai) 

323 forced_source_partitions[(apdb_part, apdb_time_part)].append(forced_source) 

324 

325 forced_source_deletes = [] 

326 forced_source_count = 0 

327 for (apdb_part, apdb_time_part), forced_source_list in forced_source_partitions.items(): 

328 clustering_keys = sorted( 

329 (fsource.diaObjectId, fsource.visit, fsource.detector) for fsource in forced_source_list 

330 ) 

331 forced_source_count += len(clustering_keys) 

332 for key_chunk in chunk_iterable(clustering_keys, 1000): 

333 cl_str = ",".join(f"({oid}, {v}, {d})" for oid, v, d in key_chunk) 

334 if config.partitioning.time_partition_tables: 

335 table_name = context.schema.tableName(ApdbTables.DiaForcedSource, apdb_time_part) 

336 query = ( 

337 Delete(keyspace, table_name) 

338 .where(C("apdb_part") == apdb_part) 

339 .where(f'("diaObjectId", visit, detector) IN ({cl_str})') 

340 ) 

341 forced_source_deletes.append(context.stmt_factory.with_params(query)) 

342 else: 

343 table_name = context.schema.tableName(ApdbTables.DiaForcedSource) 

344 query = ( 

345 Delete(keyspace, table_name) 

346 .where(C("apdb_part") == apdb_part) 

347 .where(C("apdb_time_part") == apdb_time_part) 

348 .where(f'("diaObjectId", visit, detector) IN ({cl_str})') 

349 ) 

350 forced_source_deletes.append(context.stmt_factory.with_params(query)) 

351 

352 _LOG.info( 

353 "Deleting %d objects, %d sources, and %d forced sources", 

354 object_count, 

355 source_count, 

356 forced_source_count, 

357 ) 

358 

359 # Now run all queries. 

360 with self._timer("delete_forced_sources"): 

361 execute_concurrent(context.session, forced_source_deletes) 

362 with self._timer("delete_sources"): 

363 execute_concurrent(context.session, source_deletes) 

364 with self._timer("delete_objects"): 

365 execute_concurrent(context.session, object_deletes) 

366 

367 def time_partitions(self) -> ApdbCassandraTimePartitionRange: 

368 """Return range of existing time partitions. 

369 

370 Returns 

371 ------- 

372 range : `ApdbCassandraTimePartitionRange` 

373 Time partition range. 

374 

375 Raises 

376 ------ 

377 TypeError 

378 Raised if APDB instance does not use time-partition tables. 

379 """ 

380 context = self._apdb._context 

381 part_range = context.time_partitions_range 

382 if not part_range: 

383 raise TypeError("This APDB instance does not use time-partitioned tables.") 

384 return part_range 

385 

386 def extend_time_partitions( 

387 self, 

388 time: astropy.time.Time, 

389 forward: bool = True, 

390 max_delta: astropy.time.TimeDelta | None = None, 

391 ) -> list[int]: 

392 """Extend set of time-partitioned tables to include specified time. 

393 

394 Parameters 

395 ---------- 

396 time : `astropy.time.Time` 

397 Time to which to extend partitions. 

398 forward : `bool`, optional 

399 If `True` then extend partitions into the future, time should be 

400 later than the end time of the last existing partition. If `False` 

401 then extend partitions into the past, time should be earlier than 

402 the start time of the first existing partition. 

403 max_delta : `astropy.time.TimeDelta`, optional 

404 Maximum possible extension of the aprtitions, default is 365 days. 

405 

406 Returns 

407 ------- 

408 partitions : `list` [`int`] 

409 List of partitons added to the database, empty list returned if 

410 ``time`` is already in the existing partition range. 

411 

412 Raises 

413 ------ 

414 TypeError 

415 Raised if APDB instance does not use time-partition tables. 

416 ValueError 

417 Raised if extension request exceeds time limit of ``max_delta``. 

418 """ 

419 if max_delta is None: 

420 max_delta = astropy.time.TimeDelta(365, format="jd") 

421 

422 context = self._apdb._context 

423 

424 # Get current partitions. 

425 part_range = context.time_partitions_range 

426 if not part_range: 

427 raise TypeError("This APDB instance does not use time-partitioned tables.") 

428 

429 # Partitions that we need to create. 

430 partitions = self._partitions_to_add(time, forward, max_delta) 

431 if not partitions: 

432 return [] 

433 

434 _LOG.debug("New partitions to create: %s", partitions) 

435 

436 # Tables that are time-partitioned. 

437 keyspace = self._apdb._keyspace 

438 tables = context.schema.time_partitioned_tables() 

439 

440 # Easiest way to create new tables is to take DDL from existing one 

441 # and update table name. 

442 table_name_token = "%TABLE_NAME%" 

443 table_schemas = {} 

444 for table in tables: 

445 existing_table_name = context.schema.tableName(table, part_range.end) 

446 query = f'DESCRIBE TABLE "{keyspace}"."{existing_table_name}"' 

447 result = context.session.execute(query).one() 

448 if not result: 

449 raise LookupError(f'Failed to read schema for table "{keyspace}"."{existing_table_name}"') 

450 schema: str = result.create_statement 

451 schema = schema.replace(existing_table_name, table_name_token) 

452 table_schemas[table] = schema 

453 

454 # Be paranoid and check that none of the new tables exist. 

455 exsisting_tables = context.schema.existing_tables(*tables) 

456 for table in tables: 

457 new_tables = {context.schema.tableName(table, partition) for partition in partitions} 

458 old_tables = new_tables.intersection(exsisting_tables[table]) 

459 if old_tables: 

460 raise ValueError(f"Some to be created tables already exist: {old_tables}") 

461 

462 # Now can create all of them. 

463 for table, schema in table_schemas.items(): 

464 for partition in partitions: 

465 new_table_name = context.schema.tableName(table, partition) 

466 _LOG.debug("Creating table %s", new_table_name) 

467 new_ddl = schema.replace(table_name_token, new_table_name) 

468 context.session.execute(new_ddl) 

469 

470 # Update metadata. 

471 if forward: 

472 part_range.end = max(partitions) 

473 else: 

474 part_range.start = min(partitions) 

475 part_range.save_to_meta(context.metadata) 

476 

477 return partitions 

478 

479 def _partitions_to_add( 

480 self, 

481 time: astropy.time.Time, 

482 forward: bool, 

483 max_delta: astropy.time.TimeDelta, 

484 ) -> list[int]: 

485 """Make the list of time partitions to add to current range.""" 

486 context = self._apdb._context 

487 part_range = context.time_partitions_range 

488 assert part_range is not None 

489 

490 new_partition = context.partitioner.time_partition(time) 

491 if forward: 

492 if new_partition <= part_range.end: 

493 _LOG.debug( 

494 "Partition for time=%s (%d) is below existing end (%d)", 

495 time, 

496 new_partition, 

497 part_range.end, 

498 ) 

499 return [] 

500 _, end = context.partitioner.partition_period(part_range.end) 

501 if time - end > max_delta: 

502 raise ValueError( 

503 f"Extension exceeds limit: current end time = {end.isot}, new end time = {time.isot}, " 

504 f"limit = {max_delta.jd} days" 

505 ) 

506 partitions = list(range(part_range.end + 1, new_partition + 1)) 

507 else: 

508 if new_partition >= part_range.start: 

509 _LOG.debug( 

510 "Partition for time=%s (%d) is above existing start (%d)", 

511 time, 

512 new_partition, 

513 part_range.start, 

514 ) 

515 return [] 

516 start, _ = context.partitioner.partition_period(part_range.start) 

517 if start - time > max_delta: 

518 raise ValueError( 

519 f"Extension exceeds limit: current start time = {start.isot}, " 

520 f"new start time = {time.isot}, " 

521 f"limit = {max_delta.jd} days" 

522 ) 

523 partitions = list(range(new_partition, part_range.start)) 

524 

525 return partitions 

526 

527 def delete_time_partitions( 

528 self, time: astropy.time.Time, after: bool = False, *, confirm: ConfirmDeletePartitions | None = None 

529 ) -> list[int]: 

530 """Delete time-partitioned tables before or after specified time. 

531 

532 Parameters 

533 ---------- 

534 time : `astropy.time.Time` 

535 Time before or after which to remove partitions. Partition that 

536 includes this time is not deleted. 

537 after : `bool`, optional 

538 If `True` then delete partitions after the specified time. Default 

539 is to delete partitions before this time. 

540 confirm : `~collections.abc.Callable`, optional 

541 A callable that will be called to confirm deletion of the 

542 partitions. The callable needs to accept three keyword arguments: 

543 

544 - `partitions` - a list of partition numbers to be deleted, 

545 - `tables` - a list of table names to be deleted, 

546 - `partitioner` - a `Partitioner` instance. 

547 

548 Partitions are deleted only if callable returns `True`. 

549 

550 Returns 

551 ------- 

552 partitions : `list` [`int`] 

553 List of partitons deleted from the database, empty list returned if 

554 nothing is deleted. 

555 

556 Raises 

557 ------ 

558 TypeError 

559 Raised if APDB instance does not use time-partition tables. 

560 ValueError 

561 Raised if requested to delete all partitions. 

562 """ 

563 context = self._apdb._context 

564 

565 # Get current partitions. 

566 part_range = context.time_partitions_range 

567 if not part_range: 

568 raise TypeError("This APDB instance does not use time-partitioned tables.") 

569 

570 partitions = self._partitions_to_delete(time, after) 

571 if not partitions: 

572 return [] 

573 

574 # Cannot delete all partitions. 

575 if min(partitions) == part_range.start and max(partitions) == part_range.end: 

576 raise ValueError("Cannot delete all partitions.") 

577 

578 # Tables that are time-partitioned. 

579 keyspace = self._apdb._keyspace 

580 tables = context.schema.time_partitioned_tables() 

581 

582 table_names = [] 

583 for table in tables: 

584 for partition in partitions: 

585 table_names.append(context.schema.tableName(table, partition)) 

586 

587 if confirm is not None: 

588 # It can raise an exception, but at this point it's completely 

589 # harmless. 

590 answer = confirm(partitions=partitions, tables=table_names, partitioner=context.partitioner) 

591 if not answer: 

592 return [] 

593 

594 for table_name in table_names: 

595 _LOG.debug("Dropping table %s", table_name) 

596 # Use IF EXISTS just in case. 

597 query = f'DROP TABLE IF EXISTS "{keyspace}"."{table_name}"' 

598 context.session.execute(query) 

599 

600 # Update metadata. 

601 if after: 

602 part_range.end = min(partitions) - 1 

603 else: 

604 part_range.start = max(partitions) + 1 

605 part_range.save_to_meta(context.metadata) 

606 

607 return partitions 

608 

609 def _partitions_to_delete( 

610 self, 

611 time: astropy.time.Time, 

612 after: bool = False, 

613 ) -> list[int]: 

614 """Make the list of time partitions to delete.""" 

615 context = self._apdb._context 

616 part_range = context.time_partitions_range 

617 assert part_range is not None 

618 

619 partition = context.partitioner.time_partition(time) 

620 if after: 

621 return list(range(max(partition + 1, part_range.start), part_range.end + 1)) 

622 else: 

623 return list(range(part_range.start, min(partition, part_range.end + 1)))