Coverage for tests/test_apdbReconstruct.py: 97%
114 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-04 09:49 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-04 09:49 +0000
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/>.
22import unittest
24import lsst.utils.tests
25import numpy as np
26import pandas as pd
28from lsst.analysis.ap.apdbReconstruct import (
29 ApdbReconstructor,
30 InMemoryDbQuery,
31)
34def _diaSources():
35 """Two distinct diaSourceIds, plus one duplicate (later visit wins)."""
36 return pd.DataFrame({
37 "diaSourceId": [100, 200, 100],
38 "diaObjectId": [10, 20, 10],
39 "visit": [1, 2, 3],
40 "detector": [50, 50, 50],
41 "ra": [1.0, 2.0, 1.5],
42 "dec": [-1.0, -2.0, -1.5],
43 "midpointMjdTai": [60100.0, 60101.0, 60110.0],
44 "psfFlux": [100.0, 200.0, 150.0],
45 "psfFluxErr": [10.0, 20.0, 15.0],
46 "band": ["g", "r", "g"],
47 "x": [1.0, 2.0, 1.0],
48 "y": [1.0, 2.0, 1.0],
49 "psfNdata": [10, 10, 10],
50 })
53def _diaObjects():
54 """diaObjectId=10 appears twice with different validityStart; the later
55 snapshot should win after dedup. diaObjectId=20 appears once.
56 """
57 return pd.DataFrame({
58 "diaObjectId": [10, 20, 10],
59 "validityStartMjdTai": [60100.0, 60101.0, 60110.0],
60 "ra": [1.0, 2.0, 1.5],
61 "dec": [-1.0, -2.0, -1.5],
62 "nDiaSources": [1, 1, 2],
63 })
66def _diaForcedSources():
67 """Three rows with the second duplicating the first on PK (latest wins)."""
68 return pd.DataFrame({
69 "diaForcedSourceId": [1000, 1001, 1002],
70 "diaObjectId": [10, 10, 20],
71 "visit": [1, 1, 2], # row 0 and row 1 share PK
72 "detector": [50, 50, 50],
73 "midpointMjdTai": [60100.0, 60100.0, 60101.0],
74 "psfFlux": [100.0, 150.0, 200.0], # later (150) should win
75 "psfFluxErr": [10.0, 15.0, 20.0],
76 "band": ["g", "g", "r"],
77 "ra": [1.0, 1.0, 2.0],
78 "dec": [-1.0, -1.0, -2.0],
79 "scienceFlux": [50.0, 75.0, 100.0],
80 "scienceFluxErr": [5.0, 7.0, 10.0],
81 "timeProcessedMjdTai": [60101.0, 60102.0, 60102.0],
82 })
85class TestFinalize(lsst.utils.tests.TestCase):
86 """Tests for the staticmethod that does dedup + schema coercion."""
88 def test_diaSource_dedup_by_id(self):
89 result = ApdbReconstructor.finalize(
90 _diaSources(), _diaObjects(), _diaForcedSources(),
91 coerce_to_schema=False)
92 # 3 input rows -> 2 unique diaSourceIds.
93 self.assertEqual(len(result.diaSources), 2)
94 self.assertEqual(sorted(result.diaSources["diaSourceId"]), [100, 200])
95 # The duplicate-keep="last" entry should win: visit 3 wins over visit 1
96 row100 = result.diaSources.set_index("diaSourceId").loc[100]
97 self.assertEqual(int(row100["visit"]), 3)
99 def test_diaForcedSource_dedup_by_pk(self):
100 result = ApdbReconstructor.finalize(
101 _diaSources(), _diaObjects(), _diaForcedSources(),
102 coerce_to_schema=False)
103 # 3 input rows -> 2 unique (diaObjectId, visit, detector).
104 self.assertEqual(len(result.diaForcedSources), 2)
105 # The later row for the dup PK should win (psfFlux=150, not 100).
106 dup = result.diaForcedSources[
107 (result.diaForcedSources["diaObjectId"] == 10)
108 & (result.diaForcedSources["visit"] == 1)
109 & (result.diaForcedSources["detector"] == 50)]
110 self.assertEqual(len(dup), 1)
111 self.assertAlmostEqual(float(dup["psfFlux"].iloc[0]), 150.0)
113 def test_diaObject_keeps_latest_by_validity(self):
114 result = ApdbReconstructor.finalize(
115 _diaSources(), _diaObjects(), _diaForcedSources(),
116 coerce_to_schema=False)
117 self.assertEqual(len(result.diaObjects), 2)
118 # diaObject 10 had two snapshots; the later (validityStart=60110)
119 # should win — nDiaSources=2, not 1.
120 row10 = result.diaObjects.set_index("diaObjectId").loc[10]
121 self.assertEqual(int(row10["nDiaSources"]), 2)
123 def test_diaObject_dedup_skips_nan_nDiaSources(self):
124 """Passthrough snapshots from quanta that touch a diaObject but
125 don't actually update it leave ``nDiaSources`` as NaN. The dedup
126 must skip those in favor of an older snapshot that carries the
127 real count — otherwise the survivor's NaN gets fillna(0)'d during
128 schema coercion and the user sees ``nDiaSources=0`` even though
129 the diaObject has real diaSources.
130 """
131 diaObjects = pd.DataFrame({
132 "diaObjectId": [10, 10, 10, 20], # noqa: E241
133 "validityStartMjdTai": [0.0, 61167.0, 61168.0, 61167.0], # noqa: E241
134 "nDiaSources": [1.0, np.nan, np.nan, 2.0], # noqa: E241
135 "ra": [1.0, 1.0, 1.0, 2.0], # noqa: E241
136 "dec": [-1.0, -1.0, -1.0, -2.0], # noqa: E241
137 })
138 empty = pd.DataFrame()
139 result = ApdbReconstructor.finalize(
140 empty, diaObjects, empty, coerce_to_schema=False)
141 # diaObject 10: must survive with nDiaSources=1 (NOT NaN, NOT 0).
142 row10 = result.diaObjects.set_index("diaObjectId").loc[10]
143 self.assertEqual(int(row10["nDiaSources"]), 1)
144 # diaObject 20: untouched (only one row), should round-trip.
145 row20 = result.diaObjects.set_index("diaObjectId").loc[20]
146 self.assertEqual(int(row20["nDiaSources"]), 2)
148 def test_diaObject_dedup_prefers_higher_nDiaSources(self):
149 """When a diaObject has multiple informative snapshots, dedup
150 picks the one with the most diaSources (which is also typically
151 the latest update; this ordering is robust against snapshot
152 re-ordering during concatenation).
153 """
154 diaObjects = pd.DataFrame({
155 "diaObjectId": [10, 10, 10], # noqa: E241
156 "validityStartMjdTai": [60100.0, 60110.0, 60120.0], # noqa: E241
157 "nDiaSources": [1.0, 3.0, 2.0], # noqa: E241
158 "ra": [1.0, 1.0, 1.0], # noqa: E241
159 "dec": [-1.0, -1.0, -1.0], # noqa: E241
160 })
161 empty = pd.DataFrame()
162 result = ApdbReconstructor.finalize(
163 empty, diaObjects, empty, coerce_to_schema=False)
164 row10 = result.diaObjects.set_index("diaObjectId").loc[10]
165 # The snapshot with nDiaSources=3 wins (highest count), even
166 # though a later validity is available.
167 self.assertEqual(int(row10["nDiaSources"]), 3)
169 def test_diaObject_history_keeps_all(self):
170 result = ApdbReconstructor.finalize(
171 _diaSources(), _diaObjects(), _diaForcedSources(),
172 coerce_to_schema=False, history=True)
173 # Full update trail preserved.
174 self.assertEqual(len(result.diaObjects), 3)
176 def test_schema_coercion_dtypes(self):
177 result = ApdbReconstructor.finalize(
178 _diaSources(), _diaObjects(), _diaForcedSources(),
179 coerce_to_schema=True)
180 # Integer IDs come out as Int64 (nullable long) or int64.
181 # diaSourceId is non-nullable -> int64;
182 # diaObjectId is nullable -> Int64.
183 self.assertEqual(str(result.diaSources["diaSourceId"].dtype), "int64")
184 self.assertEqual(str(result.diaSources["diaObjectId"].dtype), "Int64")
185 # Coercion also fills in schema columns that were missing — these
186 # ought to appear with sensible defaults.
187 self.assertIn("snr", result.diaSources.columns)
188 # Extra columns NOT in the schema are dropped by
189 # convertDataFrameToSdmSchema.
190 # (We didn't introduce any in the fixtures, but verify the call
191 # didn't add a bogus index column.)
192 self.assertNotIn("index", result.diaSources.columns)
194 def test_empty_inputs(self):
195 empty = pd.DataFrame()
196 result = ApdbReconstructor.finalize(empty, empty, empty,
197 coerce_to_schema=False)
198 self.assertEqual(len(result.diaSources), 0)
199 self.assertEqual(len(result.diaObjects), 0)
200 self.assertEqual(len(result.diaForcedSources), 0)
203class TestInMemoryDbQuery(lsst.utils.tests.TestCase):
204 """Tests that the DbQuery adapter routes queries against the underlying
205 DataFrames the way the SQL backends do.
206 """
208 def setUp(self):
209 recon = ApdbReconstructor.finalize(
210 _diaSources(), _diaObjects(), _diaForcedSources(),
211 coerce_to_schema=False)
212 self.query = InMemoryDbQuery(recon.diaSources,
213 recon.diaObjects,
214 recon.diaForcedSources)
216 def test_load_sources_for_object(self):
217 result = self.query.load_sources_for_object(10)
218 self.assertEqual(len(result), 1)
219 self.assertEqual(int(result["diaSourceId"].iloc[0]), 100)
221 def test_load_forced_sources_for_object_ignores_exclude_flagged(self):
222 # DiaForcedSource has no flag columns; exclude_flagged is a no-op
223 # for parity with the abstract interface.
224 result = self.query.load_forced_sources_for_object(
225 10, exclude_flagged=True)
226 self.assertEqual(len(result), 1)
227 self.assertEqual(int(result["visit"].iloc[0]), 1)
229 def test_load_source_raises_when_missing(self):
230 with self.assertRaisesRegex(RuntimeError, "diaSourceId=999999"):
231 self.query.load_source(999999)
233 def test_load_object_round_trip(self):
234 obj = self.query.load_object(10)
235 self.assertEqual(int(obj["diaObjectId"]), 10)
236 self.assertEqual(int(obj["nDiaSources"]), 2)
238 def test_load_forced_source_round_trip(self):
239 result = self.query.load_forced_source(1002)
240 self.assertEqual(int(result["diaForcedSourceId"]), 1002)
241 self.assertEqual(int(result["visit"]), 2)
243 def test_excluded_flag_validation(self):
244 with self.assertRaisesRegex(ValueError, "not present"):
245 self.query.set_excluded_diaSource_flags(["pixelFlags_bad"])
247 def test_load_sources_with_exclude_flagged(self):
248 # Add a flag column to one row and verify it's excluded.
249 diaSrc = _diaSources()
250 diaSrc["pixelFlags_bad"] = [False, True, False]
251 recon = ApdbReconstructor.finalize(
252 diaSrc, _diaObjects(), _diaForcedSources(),
253 coerce_to_schema=False)
254 q = InMemoryDbQuery(recon.diaSources, recon.diaObjects,
255 recon.diaForcedSources)
256 q.set_excluded_diaSource_flags(["pixelFlags_bad"])
257 # No exclusion requested: all 2 deduped rows returned.
258 self.assertEqual(len(q.load_sources()), 2)
259 # Exclusion requested: the flagged row is dropped.
260 flagged = q.load_sources(exclude_flagged=True)
261 self.assertEqual(len(flagged), 1)
262 self.assertNotIn(200, flagged["diaSourceId"].tolist())
265class TestDatasetNameDefaults(lsst.utils.tests.TestCase):
266 """Pin the dataset-name defaults to the ApPipe.yaml `associateApdb`
267 config, since downstream production tooling relies on these names.
268 """
270 def test_default_names_match_ap_pipe(self):
271 # The "apdb" entries come from ApPipe.yaml's `associateApdb` task
272 # config; the "preloaded_*" entries come from
273 # `lsst.ap.association.LoadDiaCatalogsTask` output names. Both
274 # are loaded so the reconstruction includes both prior history
275 # and the current run's new rows.
276 recon = ApdbReconstructor(butler=None)
277 self.assertEqual(recon.dataset_names["diaSource"],
278 ["dia_source_apdb", "preloaded_dia_source"])
279 self.assertEqual(recon.dataset_names["diaObject"],
280 ["dia_object_apdb", "preloaded_dia_object"])
281 self.assertEqual(recon.dataset_names["diaForcedSource"],
282 ["dia_forced_source_apdb",
283 "preloaded_dia_forced_source"])
285 def test_dataset_names_override(self):
286 # A full dict replaces DEFAULT_DATASET_NAMES wholesale.
287 override = {
288 "diaSource": ["goodSeeingDiff_assocDiaSrc"],
289 "diaObject": ["goodSeeingDiff_diaObject"],
290 "diaForcedSource": ["goodSeeingDiff_diaForcedSrc"],
291 }
292 recon = ApdbReconstructor(butler=None, dataset_names=override)
293 self.assertEqual(recon.dataset_names, override)
295 def test_dataset_names_override_list(self):
296 # A list override replaces the default list entirely.
297 recon = ApdbReconstructor(
298 butler=None,
299 dataset_names={"diaSource": ["a", "b", "c"]})
300 self.assertEqual(recon.dataset_names["diaSource"], ["a", "b", "c"])
303class MemoryTester(lsst.utils.tests.MemoryTestCase):
304 pass
307def setup_module(module):
308 lsst.utils.tests.init()
311if __name__ == "__main__": 311 ↛ 312line 311 didn't jump to line 312 because the condition on line 311 was never true
312 lsst.utils.tests.init()
313 unittest.main()