Coverage for tests/test_apdbReconstruct.py: 97%

114 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-16 10:41 +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/>. 

21 

22import unittest 

23 

24import lsst.utils.tests 

25import numpy as np 

26import pandas as pd 

27 

28from lsst.analysis.ap.apdbReconstruct import ( 

29 ApdbReconstructor, 

30 InMemoryDbQuery, 

31) 

32 

33 

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

51 

52 

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

64 

65 

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

83 

84 

85class TestFinalize(lsst.utils.tests.TestCase): 

86 """Tests for the staticmethod that does dedup + schema coercion.""" 

87 

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) 

98 

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) 

112 

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) 

122 

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) 

147 

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) 

168 

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) 

175 

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) 

193 

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) 

201 

202 

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

207 

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) 

215 

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) 

220 

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) 

228 

229 def test_load_source_raises_when_missing(self): 

230 with self.assertRaisesRegex(RuntimeError, "diaSourceId=999999"): 

231 self.query.load_source(999999) 

232 

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) 

237 

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) 

242 

243 def test_excluded_flag_validation(self): 

244 with self.assertRaisesRegex(ValueError, "not present"): 

245 self.query.set_excluded_diaSource_flags(["pixelFlags_bad"]) 

246 

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

263 

264 

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

269 

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

284 

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) 

294 

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

301 

302 

303class MemoryTester(lsst.utils.tests.MemoryTestCase): 

304 pass 

305 

306 

307def setup_module(module): 

308 lsst.utils.tests.init() 

309 

310 

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