Coverage for python/lsst/analysis/ap/apdb.py: 83%
194 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-30 04:48 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-30 04:48 -0700
1# This file is part of analysis_ap.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (https://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 <https://www.gnu.org/licenses/>.
22"""APDB connection management and data access tools.
23"""
25__all__ = ["DbQuery", "ApdbSqliteQuery", "ApdbPostgresQuery"]
27import abc
28import contextlib
29import os
30import warnings
31from importlib.resources import files
33import felis.datamodel
34import pandas as pd
35import sqlalchemy
37from lsst.pipe.tasks.schemaUtils import column_dtype, readSdmSchemaFile
40# Integer felis types whose default pandas read path silently coerces NULL
41# rows through float64. Drive the dtype from the SDM schema for these types
42# only and leave float / bool / string / timestamp columns to pandas' default
43# inference so we don't perturb unrelated behavior.
44_INT_FELIS_TYPES = frozenset({
45 felis.datamodel.DataType.long,
46 felis.datamodel.DataType.int,
47 felis.datamodel.DataType.short,
48 felis.datamodel.DataType.byte,
49})
52_apdb_schema_cache = None
55def _apdb_schema():
56 """Lazily load the APDB SDM schema (``sdm_schemas/apdb.yaml``) once."""
57 global _apdb_schema_cache
58 if _apdb_schema_cache is None:
59 path = os.fspath(files("lsst.sdm.schemas").joinpath("apdb.yaml"))
60 _apdb_schema_cache = readSdmSchemaFile(path)
61 return _apdb_schema_cache
64def _schema_int_dtypes(table_name):
65 """Return a {column_name: pandas-dtype} mapping for integer columns of
66 one APDB table.
68 Returns an empty dict if the table is not in the SDM schema (e.g. an
69 older fixture with tables that have since been removed).
70 """
71 table_def = _apdb_schema().get(table_name)
72 if table_def is None: 72 ↛ 73line 72 didn't jump to line 73 because the condition on line 72 was never true
73 return {}
74 return {
75 cdef.name: column_dtype(cdef.datatype, nullable=cdef.nullable)
76 for cdef in table_def.columns
77 if cdef.datatype in _INT_FELIS_TYPES
78 }
81def _read_query(connection, query, table_name=None):
82 """Run ``query`` on ``connection`` and return a DataFrame with correct
83 integer-column dtypes.
85 ``pd.read_sql_query`` represents SQL NULLs as ``NaN`` and so forces any
86 nullable BIGINT column through ``float64``.
87 Here we fetch rows directly from the cursor and build each integer
88 column with the dtype declared by the SDM schema (looked up via
89 ``schemaUtils.column_dtype``); other columns fall through to pandas'
90 default inference.
91 """
92 cursor = connection.execute(query)
93 columns = list(cursor.keys())
94 rows = cursor.fetchall()
95 dtype_map = _schema_int_dtypes(table_name) if table_name else {}
96 data = {}
97 for i, col in enumerate(columns):
98 values = [row[i] for row in rows]
99 dtype = dtype_map.get(col)
100 if dtype is not None:
101 try:
102 data[col] = pd.array(values, dtype=dtype)
103 continue
104 except (TypeError, ValueError):
105 # Cast impossible (driver returned an unexpected type);
106 # fall through to pandas inference.
107 pass
108 data[col] = values
109 return pd.DataFrame(data)
112class DbQuery(abc.ABC):
113 """Abstract interface for APDB queries.
115 Notes
116 -----
117 APDB interface used by AP pipeline is defined by `lsst.dax.apdb.Apdb`
118 class. Methods in this class are for non-pipeline tools that can analyse
119 data produced by pipeline. APDB schema is not designed for analysis queries
120 and performance of these methods can be non-optimal, especially for
121 Cassandra backend. It is expected that these analysis queries should not be
122 executed on production Cassandra service.
123 """
125 def set_excluded_diaSource_flags(self, flag_list: list[str]) -> None:
126 """Set flags of diaSources to exclude when loading diaSources.
128 Any diaSources with configured flags are not returned
129 when calling `load_sources_for_object` or `load_sources`
130 with `exclude_flagged = True`.
132 Parameters
133 ----------
134 flag_list : `list` [`str`]
135 Flag names to exclude.
136 """
137 raise NotImplementedError()
139 def load_sources_for_object(
140 self, dia_object_id: int, exclude_flagged: bool = False, limit: int = 100000
141 ) -> pd.DataFrame:
142 """Load diaSources for a single diaObject.
144 Parameters
145 ----------
146 dia_object_id : `int`
147 Id of object to load sources for.
148 exclude_flagged : `bool`, optional
149 Exclude sources that have selected flags set.
150 Use `set_excluded_diaSource_flags` to configure which flags
151 are excluded.
152 limit : `int`
153 Maximum number of rows to return.
155 Returns
156 -------
157 data : `pandas.DataFrame`
158 A data frame of diaSources for the specified diaObject.
159 """
160 raise NotImplementedError()
162 def load_forced_sources_for_object(
163 self, dia_object_id: int, exclude_flagged: bool = False, limit: int = 100000
164 ) -> pd.DataFrame:
165 """Load diaForcedSources for a single diaObject.
167 Parameters
168 ----------
169 dia_object_id : `int`
170 Id of object to load sources for.
171 exclude_flagged : `bool`, optional
172 Exclude sources that have selected flags set.
173 Use `set_excluded_diaSource_flags` to configure which flags
174 are excluded.
175 limit : `int`
176 Maximum number of rows to return.
178 Returns
179 -------
180 data : `pandas.DataFrame`
181 A data frame of diaSources for the specified diaObject.
182 """
183 raise NotImplementedError()
185 def load_source(self, id: int) -> pd.Series:
186 """Load one diaSource.
188 Parameters
189 ----------
190 id : `int`
191 The diaSourceId to load data for.
193 Returns
194 -------
195 data : `pandas.Series`
196 The requested diaSource.
197 """
198 raise NotImplementedError()
200 def load_sources(self, exclude_flagged: bool = False, limit: int = 100000) -> pd.DataFrame:
201 """Load diaSources.
203 Parameters
204 ----------
205 exclude_flagged : `bool`, optional
206 Exclude sources that have selected flags set.
207 Use `set_excluded_diaSource_flags` to configure which flags
208 are excluded.
209 limit : `int`
210 Maximum number of rows to return.
212 Returns
213 -------
214 data : `pandas.DataFrame`
215 All available diaSources.
216 """
217 raise NotImplementedError()
219 def load_object(self, id: int) -> pd.Series:
220 """Load the most-recently updated version of one diaObject.
222 Parameters
223 ----------
224 id : `int`
225 The diaObjectId to load data for.
227 Returns
228 -------
229 data : `pandas.Series`
230 The requested object.
231 """
232 raise NotImplementedError()
234 def load_objects(self, limit: int = 100000, latest: bool = True) -> pd.DataFrame:
235 """Load all diaObjects.
237 Parameters
238 ----------
239 limit : `int`
240 Maximum number of rows to return.
241 latest : `bool`
242 Only load diaObjects where validityEnd is None.
243 These are the most-recently updated diaObjects.
245 Returns
246 -------
247 data : `pandas.DataFrame`
248 All available diaObjects.
249 """
250 raise NotImplementedError()
252 def load_forced_source(self, id: int) -> pd.Series:
253 """Load one diaForcedSource.
255 Parameters
256 ----------
257 id : `int`
258 The diaForcedSourceId to load data for.
260 Returns
261 -------
262 data : `pandas.Series`
263 The requested forced source.
264 """
265 raise NotImplementedError()
267 def load_forced_sources(self, limit: int = 100000) -> pd.DataFrame:
268 """Load all diaForcedSources.
270 Parameters
271 ----------
272 limit : `int`
273 Maximum number of rows to return.
275 Returns
276 -------
277 data : `pandas.DataFrame`
278 All available diaForcedSources.
279 """
280 raise NotImplementedError()
283class DbSqlQuery(DbQuery):
284 """Base class for APDB connection and query management for SQL backends.
286 Subclasses must specify a ``connection`` property to use as a context-
287 manager for queries.
289 Parameters
290 ----------
291 instrument : `str`
292 Short name (e.g. "DECam") of instrument to make a dataId unpacker
293 and to add to the table columns; supports any gen3 instrument.
294 To be deprecated once this information is in the database.
295 """
297 def __init__(self, instrument=None):
298 if instrument is not None: 298 ↛ 299line 298 didn't jump to line 299 because the condition on line 298 was never true
299 warnings.warn("The instrument name is now pulled from the APDB; "
300 "this kwarg is ignored and will be removed after v29",
301 FutureWarning,
302 stacklevel=2)
304 self.set_excluded_diaSource_flags(['pixelFlags_bad',
305 'pixelFlags_suspect',
306 'pixelFlags_saturatedCenter',
307 'pixelFlags_interpolated',
308 'pixelFlags_interpolatedCenter',
309 'pixelFlags_edge',
310 ])
312 key = "instrument"
313 table = self._tables["metadata"]
314 sql = sqlalchemy.sql.select(table.columns.value).where(table.columns.name == key)
315 with self.connection as conn:
316 result = conn.execute(sql)
317 self._instrument = result.scalar()
319 @property
320 @contextlib.contextmanager
321 @abc.abstractmethod
322 def connection(self):
323 """Context manager for database connections.
325 Yields
326 ------
327 connection : `sqlalchemy.engine.Connection`
328 Connection to the database that will be queried. Whether the
329 connection is closed after the context manager is closed is
330 implementation dependent.
331 """
332 pass
334 def set_excluded_diaSource_flags(self, flag_list):
335 # Docstring is inherited.
336 for flag in flag_list:
337 if flag not in self._tables["DiaSource"].columns:
338 raise ValueError(f"flag {flag} not included in DiaSource flags")
340 self.diaSource_flags_exclude = flag_list
342 def _make_flag_exclusion_query(self, query, table, flag_list):
343 """Attach a where clause excluding sources with any chosen flag set.
345 Parameters
346 ----------
347 query : `sqlalchemy.sql.Select`
348 Query to attach the where clause to.
349 table : `sqlalchemy.schema.Table`
350 Reflected table containing the flag columns.
351 flag_list : `list` [`str`]
352 Flag column names to exclude.
354 Returns
355 -------
356 query : `sqlalchemy.sql.Select`
357 Query with the flag exclusion clause attached.
358 """
359 return query.where(sqlalchemy.and_(table.columns[col] == False # noqa: E712
360 for col in flag_list))
362 def _load_table(self, table, *, where=None, exclude_flagged=False,
363 order_by=(), limit=None, fill_instrument=True):
364 """Run a parameterized SELECT and return the result as a DataFrame.
366 Parameters
367 ----------
368 table : `sqlalchemy.schema.Table`
369 Reflected table to query.
370 where : `sqlalchemy.sql.ClauseElement`, optional
371 Extra where clause to attach.
372 exclude_flagged : `bool`, optional
373 If True, attach the configured flag-exclusion clause.
374 order_by : `tuple` [`str`], optional
375 Column names to order by.
376 limit : `int`, optional
377 Maximum number of rows to return; None means no limit.
378 fill_instrument : `bool`, optional
379 If True, append an ``instrument`` column to the result.
381 Returns
382 -------
383 result : `pandas.DataFrame`
384 """
385 query = table.select()
386 if where is not None:
387 query = query.where(where)
388 if exclude_flagged:
389 query = self._make_flag_exclusion_query(query, table, self.diaSource_flags_exclude)
390 if order_by:
391 query = query.order_by(*[table.columns[c] for c in order_by])
392 if limit is not None:
393 query = query.limit(limit)
394 with self.connection as connection:
395 result = _read_query(connection, query, table_name=table.name)
396 if fill_instrument:
397 self._fill_from_instrument(result)
398 return result
400 def _load_one(self, table_name, id_column, id_value, fill_instrument=True):
401 """Load a single row from a table by id, raising if missing.
403 Parameters
404 ----------
405 table_name : `str`
406 Key into ``self._tables`` for the table to query.
407 id_column : `str`
408 Name of the id column to filter on.
409 id_value : `int`
410 Id value to match.
411 fill_instrument : `bool`, optional
412 If True, append an ``instrument`` column to the result.
414 Returns
415 -------
416 row : `pandas.Series`
418 Raises
419 ------
420 RuntimeError
421 If no row matches.
422 """
423 table = self._tables[table_name]
424 result = self._load_table(table,
425 where=table.columns[id_column] == id_value,
426 fill_instrument=fill_instrument)
427 if len(result) == 0:
428 raise RuntimeError(f"{id_column}={id_value} not found in {table_name} table")
429 return result.iloc[0]
431 def load_sources_for_object(self, dia_object_id, exclude_flagged=False, limit=100000):
432 # Docstring is inherited.
433 table = self._tables["DiaSource"]
434 return self._load_table(
435 table,
436 where=table.columns["diaObjectId"] == dia_object_id,
437 exclude_flagged=exclude_flagged,
438 order_by=("visit", "detector", "diaSourceId"),
439 limit=limit,
440 )
442 def load_forced_sources_for_object(self, dia_object_id, exclude_flagged=False, limit=100000):
443 # Docstring is inherited.
444 table = self._tables["DiaForcedSource"]
445 return self._load_table(
446 table,
447 where=table.columns["diaObjectId"] == dia_object_id,
448 exclude_flagged=exclude_flagged,
449 order_by=("visit", "detector", "diaForcedSourceId"),
450 limit=limit,
451 )
453 def load_source(self, id):
454 # Docstring is inherited.
455 return self._load_one("DiaSource", "diaSourceId", id)
457 def load_sources(self, exclude_flagged=False, limit=100000):
458 # Docstring is inherited.
459 return self._load_table(
460 self._tables["DiaSource"],
461 exclude_flagged=exclude_flagged,
462 order_by=("visit", "detector", "diaSourceId"),
463 limit=limit,
464 )
466 @staticmethod
467 def _validity_end_column(table):
468 """Return the DiaObject "validity end" column.
470 sdm_schemas renamed this to ``validityEndMjdTai`` (it's also nullable
471 double-precision MJD now, not a TIMESTAMP). Older fixtures still have
472 the original ``validityEnd`` name; this helper prefers the current
473 name and falls back to the legacy one.
474 """
475 for name in ("validityEndMjdTai", "validityEnd"): 475 ↛ 478line 475 didn't jump to line 478 because the loop on line 475 didn't complete
476 if name in table.columns:
477 return table.columns[name]
478 raise KeyError("DiaObject has neither validityEndMjdTai nor validityEnd")
480 def load_object(self, id):
481 # Docstring is inherited.
482 table = self._tables["DiaObject"]
483 result = self._load_table(
484 table,
485 where=sqlalchemy.and_(
486 self._validity_end_column(table) == None, # noqa: E711
487 table.columns["diaObjectId"] == id,
488 ),
489 fill_instrument=False,
490 )
491 if len(result) == 0:
492 raise RuntimeError(f"diaObjectId={id} not found in DiaObject table")
493 return result.iloc[0]
495 def load_objects(self, limit=100000, latest=True):
496 # Docstring is inherited.
497 table = self._tables["DiaObject"]
498 where = self._validity_end_column(table) == None if latest else None # noqa: E711
499 return self._load_table(
500 table,
501 where=where,
502 order_by=("diaObjectId",),
503 limit=limit,
504 fill_instrument=False,
505 )
507 def load_forced_source(self, id):
508 # Docstring is inherited.
509 return self._load_one("DiaForcedSource", "diaForcedSourceId", id)
511 def load_forced_sources(self, limit=100000):
512 # Docstring is inherited.
513 return self._load_table(
514 self._tables["DiaForcedSource"],
515 order_by=("visit", "detector", "diaForcedSourceId"),
516 limit=limit,
517 )
519 def iter_sources(self, page_size=100000, reliability_min=None, reliability_max=None):
520 """Yield DiaSources in pages of ``page_size`` rows.
522 Parameters
523 ----------
524 page_size : `int`
525 Number of rows per page.
526 reliability_min, reliability_max : `float`, optional
527 Inclusive bounds on the reliability column.
529 Yields
530 ------
531 page : `pandas.DataFrame`
532 One page of DiaSources, with the ``instrument`` column attached.
533 """
534 table = self._tables["DiaSource"]
535 clauses = []
536 if reliability_min is not None: 536 ↛ 537line 536 didn't jump to line 537 because the condition on line 536 was never true
537 clauses.append(table.columns["reliability"] >= reliability_min)
538 if reliability_max is not None: 538 ↛ 539line 538 didn't jump to line 539 because the condition on line 538 was never true
539 clauses.append(table.columns["reliability"] <= reliability_max)
540 where = sqlalchemy.and_(*clauses) if clauses else None
542 offset = 0
543 while True:
544 query = table.select()
545 if where is not None: 545 ↛ 546line 545 didn't jump to line 546 because the condition on line 545 was never true
546 query = query.where(where)
547 query = query.order_by(table.columns["visit"],
548 table.columns["detector"],
549 table.columns["diaSourceId"])
550 query = query.limit(page_size).offset(offset)
551 with self.connection as connection:
552 page = _read_query(connection, query, table_name=table.name)
553 if len(page) == 0:
554 break
555 self._fill_from_instrument(page)
556 yield page
557 offset += page_size
559 def count_sources(self):
560 """Return the total number of DiaSources in the database.
562 Returns
563 -------
564 count : `int`
565 """
566 table = self._tables["DiaSource"]
567 query = sqlalchemy.select(sqlalchemy.func.count()).select_from(table)
568 with self.connection as connection:
569 return connection.execute(query).scalar()
571 def _fill_from_instrument(self, diaSources):
572 """Add an instrument column to a list of sources.
573 This method is temporary, until APDB has instrument in its metadata.
575 Parameters
576 ----------
577 diaSources : `pandas.core.frame.DataFrame`
578 Pandas dataframe with diaSources from an APDB; modified in-place.
579 """
580 # do nothing for an empty series
581 if len(diaSources) == 0:
582 return
584 diaSources['instrument'] = self._instrument
587class ApdbSqliteQuery(DbSqlQuery):
588 """Open an sqlite3 APDB file to load data from it.
590 This class keeps the sqlite connection open after initialization because
591 our sqlite usage is to load a local file. Closing and re-opening would
592 re-scan the whole file every time, and we don't need to worry about
593 multiple users when working with local sqlite files.
595 Parameters
596 ----------
597 filename : `str`
598 Path to the sqlite3 file containing the APDB to load.
599 instrument : `str`
600 Short name (e.g. "DECam") of instrument to make a dataId unpacker
601 and to add to the table columns; supports any gen3 instrument.
602 To be deprecated once this information is in the database.
603 """
605 def __init__(self, filename, instrument=None, **kwargs):
606 # For sqlite, use a larger pool and a faster timeout, to allow many
607 # repeat transactions with the same connection, as transactions on
608 # our sqlite DBs should be small and fast.
609 self._engine = sqlalchemy.create_engine(f"sqlite:///{filename}",
610 pool_timeout=5, pool_size=200)
612 with self.connection as connection:
613 metadata = sqlalchemy.MetaData()
614 metadata.reflect(bind=connection)
615 self._tables = metadata.tables
616 super().__init__(**kwargs)
618 @property
619 @contextlib.contextmanager
620 def connection(self):
621 yield self._engine.connect()
624class ApdbPostgresQuery(DbSqlQuery):
625 """Connect to a running postgres APDB instance and load data from it.
627 This class connects to the database only when the ``connection`` context
628 manager is entered, and closes the connection after it exits.
630 Parameters
631 ----------
632 namespace : `str`
633 Database namespace to load from. Called "schema" in postgres docs.
634 url : `str`
635 Complete url to connect to postgres database, without prepended
636 ``postgresql://``.
637 instrument : `str`
638 Short name (e.g. "DECam") of instrument to make a dataId unpacker
639 and to add to the table columns; supports any gen3 instrument.
640 To be deprecated once this information is in the database.
641 """
643 def __init__(self, namespace, url="rubin@usdf-prompt-processing-dev.slac.stanford.edu/lsst-devl",
644 instrument=None, **kwargs):
645 self._connection_string = f"postgresql://{url}"
646 self._namespace = namespace
647 self._engine = sqlalchemy.create_engine(self._connection_string, poolclass=sqlalchemy.pool.NullPool)
649 with self.connection as connection:
650 metadata = sqlalchemy.MetaData(schema=namespace)
651 metadata.reflect(bind=connection)
652 # ensure tables don't have schema prepended
653 self._tables = {}
654 for table in metadata.tables.values():
655 self._tables[table.name] = table
656 super().__init__(instrument=instrument, **kwargs)
658 @property
659 @contextlib.contextmanager
660 def connection(self):
661 _connection = self._engine.connect()
662 try:
663 yield _connection
664 finally:
665 _connection.close()