Coverage for python/lsst/analysis/ap/apdb.py: 83%

194 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-09-06 02:34 -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/>. 

21 

22"""APDB connection management and data access tools. 

23""" 

24 

25__all__ = ["DbQuery", "ApdbSqliteQuery", "ApdbPostgresQuery"] 

26 

27import abc 

28import contextlib 

29import os 

30import warnings 

31from importlib.resources import files 

32 

33import felis.datamodel 

34import pandas as pd 

35import sqlalchemy 

36 

37from lsst.pipe.tasks.schemaUtils import column_dtype, readSdmSchemaFile 

38 

39 

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

50 

51 

52_apdb_schema_cache = None 

53 

54 

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 

62 

63 

64def _schema_int_dtypes(table_name): 

65 """Return a {column_name: pandas-dtype} mapping for integer columns of 

66 one APDB table. 

67 

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 } 

79 

80 

81def _read_query(connection, query, table_name=None): 

82 """Run ``query`` on ``connection`` and return a DataFrame with correct 

83 integer-column dtypes. 

84 

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) 

110 

111 

112class DbQuery(abc.ABC): 

113 """Abstract interface for APDB queries. 

114 

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

124 

125 def set_excluded_diaSource_flags(self, flag_list: list[str]) -> None: 

126 """Set flags of diaSources to exclude when loading diaSources. 

127 

128 Any diaSources with configured flags are not returned 

129 when calling `load_sources_for_object` or `load_sources` 

130 with `exclude_flagged = True`. 

131 

132 Parameters 

133 ---------- 

134 flag_list : `list` [`str`] 

135 Flag names to exclude. 

136 """ 

137 raise NotImplementedError() 

138 

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. 

143 

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. 

154 

155 Returns 

156 ------- 

157 data : `pandas.DataFrame` 

158 A data frame of diaSources for the specified diaObject. 

159 """ 

160 raise NotImplementedError() 

161 

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. 

166 

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. 

177 

178 Returns 

179 ------- 

180 data : `pandas.DataFrame` 

181 A data frame of diaSources for the specified diaObject. 

182 """ 

183 raise NotImplementedError() 

184 

185 def load_source(self, id: int) -> pd.Series: 

186 """Load one diaSource. 

187 

188 Parameters 

189 ---------- 

190 id : `int` 

191 The diaSourceId to load data for. 

192 

193 Returns 

194 ------- 

195 data : `pandas.Series` 

196 The requested diaSource. 

197 """ 

198 raise NotImplementedError() 

199 

200 def load_sources(self, exclude_flagged: bool = False, limit: int = 100000) -> pd.DataFrame: 

201 """Load diaSources. 

202 

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. 

211 

212 Returns 

213 ------- 

214 data : `pandas.DataFrame` 

215 All available diaSources. 

216 """ 

217 raise NotImplementedError() 

218 

219 def load_object(self, id: int) -> pd.Series: 

220 """Load the most-recently updated version of one diaObject. 

221 

222 Parameters 

223 ---------- 

224 id : `int` 

225 The diaObjectId to load data for. 

226 

227 Returns 

228 ------- 

229 data : `pandas.Series` 

230 The requested object. 

231 """ 

232 raise NotImplementedError() 

233 

234 def load_objects(self, limit: int = 100000, latest: bool = True) -> pd.DataFrame: 

235 """Load all diaObjects. 

236 

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. 

244 

245 Returns 

246 ------- 

247 data : `pandas.DataFrame` 

248 All available diaObjects. 

249 """ 

250 raise NotImplementedError() 

251 

252 def load_forced_source(self, id: int) -> pd.Series: 

253 """Load one diaForcedSource. 

254 

255 Parameters 

256 ---------- 

257 id : `int` 

258 The diaForcedSourceId to load data for. 

259 

260 Returns 

261 ------- 

262 data : `pandas.Series` 

263 The requested forced source. 

264 """ 

265 raise NotImplementedError() 

266 

267 def load_forced_sources(self, limit: int = 100000) -> pd.DataFrame: 

268 """Load all diaForcedSources. 

269 

270 Parameters 

271 ---------- 

272 limit : `int` 

273 Maximum number of rows to return. 

274 

275 Returns 

276 ------- 

277 data : `pandas.DataFrame` 

278 All available diaForcedSources. 

279 """ 

280 raise NotImplementedError() 

281 

282 

283class DbSqlQuery(DbQuery): 

284 """Base class for APDB connection and query management for SQL backends. 

285 

286 Subclasses must specify a ``connection`` property to use as a context- 

287 manager for queries. 

288 

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

296 

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) 

303 

304 self.set_excluded_diaSource_flags(['pixelFlags_bad', 

305 'pixelFlags_suspect', 

306 'pixelFlags_saturatedCenter', 

307 'pixelFlags_interpolated', 

308 'pixelFlags_interpolatedCenter', 

309 'pixelFlags_edge', 

310 ]) 

311 

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

318 

319 @property 

320 @contextlib.contextmanager 

321 @abc.abstractmethod 

322 def connection(self): 

323 """Context manager for database connections. 

324 

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 

333 

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

339 

340 self.diaSource_flags_exclude = flag_list 

341 

342 def _make_flag_exclusion_query(self, query, table, flag_list): 

343 """Attach a where clause excluding sources with any chosen flag set. 

344 

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. 

353 

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

361 

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. 

365 

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. 

380 

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 

399 

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. 

402 

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. 

413 

414 Returns 

415 ------- 

416 row : `pandas.Series` 

417 

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] 

430 

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 ) 

441 

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 ) 

452 

453 def load_source(self, id): 

454 # Docstring is inherited. 

455 return self._load_one("DiaSource", "diaSourceId", id) 

456 

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 ) 

465 

466 @staticmethod 

467 def _validity_end_column(table): 

468 """Return the DiaObject "validity end" column. 

469 

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

479 

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] 

494 

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 ) 

506 

507 def load_forced_source(self, id): 

508 # Docstring is inherited. 

509 return self._load_one("DiaForcedSource", "diaForcedSourceId", id) 

510 

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 ) 

518 

519 def iter_sources(self, page_size=100000, reliability_min=None, reliability_max=None): 

520 """Yield DiaSources in pages of ``page_size`` rows. 

521 

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. 

528 

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 

541 

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 

558 

559 def count_sources(self): 

560 """Return the total number of DiaSources in the database. 

561 

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

570 

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. 

574 

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 

583 

584 diaSources['instrument'] = self._instrument 

585 

586 

587class ApdbSqliteQuery(DbSqlQuery): 

588 """Open an sqlite3 APDB file to load data from it. 

589 

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. 

594 

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

604 

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) 

611 

612 with self.connection as connection: 

613 metadata = sqlalchemy.MetaData() 

614 metadata.reflect(bind=connection) 

615 self._tables = metadata.tables 

616 super().__init__(**kwargs) 

617 

618 @property 

619 @contextlib.contextmanager 

620 def connection(self): 

621 yield self._engine.connect() 

622 

623 

624class ApdbPostgresQuery(DbSqlQuery): 

625 """Connect to a running postgres APDB instance and load data from it. 

626 

627 This class connects to the database only when the ``connection`` context 

628 manager is entered, and closes the connection after it exits. 

629 

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

642 

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) 

648 

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) 

657 

658 @property 

659 @contextlib.contextmanager 

660 def connection(self): 

661 _connection = self._engine.connect() 

662 try: 

663 yield _connection 

664 finally: 

665 _connection.close()