Coverage for python/lsst/dax/apdb/cassandra/apdbCassandraAdmin.py: 14%
273 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 09:16 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 09:16 +0000
1# This file is part of dax_apdb.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (http://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <http://www.gnu.org/licenses/>.
22from __future__ import annotations
24__all__ = ["ApdbCassandraAdmin"]
26import dataclasses
27import itertools
28import logging
29import warnings
30from collections import defaultdict
31from collections.abc import Iterable, Mapping
32from typing import TYPE_CHECKING, Protocol
34import astropy.time
36from lsst.utils.iteration import chunk_iterable
38try:
39 import cassandra
40except ImportError:
41 pass
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
53if TYPE_CHECKING:
54 from .apdbCassandra import ApdbCassandra
55 from .partitioner import Partitioner
57_LOG = logging.getLogger(__name__)
59_MON = MonAgent(__name__)
62class ConfirmDeletePartitions(Protocol):
63 """Protocol for callable which confirms deletion of partitions."""
65 def __call__(self, *, partitions: list[int], tables: list[str], partitioner: Partitioner) -> bool: ... 65 ↛ exitline 65 didn't return from function '__call__' because
68@dataclasses.dataclass
69class DatabaseInfo:
70 """Collection of information about a specific database."""
72 name: str
73 """Keyspace name."""
75 permissions: dict[str, set[str]] | None = None
76 """Roles that can access the database and their permissions.
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 """
84class ApdbCassandraAdmin(ApdbAdmin):
85 """Implementation of `ApdbAdmin` for Cassandra backend.
87 Parameters
88 ----------
89 apdb : `ApdbCassandra`
90 APDB implementation.
91 """
93 def __init__(self, apdb: ApdbCassandra):
94 self._apdb = apdb
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)
100 @classmethod
101 def list_databases(cls, host: str) -> Iterable[DatabaseInfo]:
102 """Return the list of keyspaces with APDB databases.
104 Parameters
105 ----------
106 host : `str`
107 Name of one of the hosts in Cassandra cluster.
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)
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)
126 result = session.execute(stmt, params)
127 keyspaces = [row[0] for row in result.all()]
129 if not keyspaces:
130 return []
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)
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}
161 # Would be nice to get size estimate, but this is not available
162 # via CQL queries.
163 return infos.values()
165 @classmethod
166 def delete_database(cls, host: str, keyspace: str, *, timeout: int = 3600) -> None:
167 """Delete APDB database by dropping its keyspace.
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)
187 @property
188 def partitioner(self) -> Partitioner:
189 """Partitoner used by this APDB instance (`Partitioner`)."""
190 context = self._apdb._context
191 return context.partitioner
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)
197 def apdb_time_part(self, midpointMjdTai: float) -> int:
198 # docstring is inherited from a base class
199 return self.partitioner.time_partition(midpointMjdTai)
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)
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()))
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)
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))
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))
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))
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)
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))
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)
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))
352 _LOG.info(
353 "Deleting %d objects, %d sources, and %d forced sources",
354 object_count,
355 source_count,
356 forced_source_count,
357 )
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)
367 def time_partitions(self) -> ApdbCassandraTimePartitionRange:
368 """Return range of existing time partitions.
370 Returns
371 -------
372 range : `ApdbCassandraTimePartitionRange`
373 Time partition range.
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
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.
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.
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.
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")
422 context = self._apdb._context
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.")
429 # Partitions that we need to create.
430 partitions = self._partitions_to_add(time, forward, max_delta)
431 if not partitions:
432 return []
434 _LOG.debug("New partitions to create: %s", partitions)
436 # Tables that are time-partitioned.
437 keyspace = self._apdb._keyspace
438 tables = context.schema.time_partitioned_tables()
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
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}")
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)
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)
477 return partitions
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
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))
525 return partitions
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.
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:
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.
548 Partitions are deleted only if callable returns `True`.
550 Returns
551 -------
552 partitions : `list` [`int`]
553 List of partitons deleted from the database, empty list returned if
554 nothing is deleted.
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
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.")
570 partitions = self._partitions_to_delete(time, after)
571 if not partitions:
572 return []
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.")
578 # Tables that are time-partitioned.
579 keyspace = self._apdb._keyspace
580 tables = context.schema.time_partitioned_tables()
582 table_names = []
583 for table in tables:
584 for partition in partitions:
585 table_names.append(context.schema.tableName(table, partition))
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 []
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)
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)
607 return partitions
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
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)))