Coverage for tests/test_plotDiaSourceLightcurve.py: 97%

148 statements  

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

22import unittest 

23 

24import lsst.afw.image 

25import lsst.geom 

26import lsst.meas.base.tests 

27import lsst.utils.tests 

28import numpy as np 

29import pandas as pd 

30import PIL 

31 

32from lsst.analysis.ap import plotDiaSourceLightcurve 

33 

34# Pull in the DATA fixture from the cutouts test so we get the full set of 

35# flag columns expected by _annotate_image. 

36from test_plotImageSubtractionCutouts import DATA, skyCenter 

37 

38 

39DIA_OBJECT_ID = 999999999999000001 

40 

41 

42# Add the extra columns required by the lightcurve task. 

43def _augment_with_object_columns(data): 

44 data = data.copy() 

45 data["diaObjectId"] = [DIA_OBJECT_ID, DIA_OBJECT_ID] 

46 data["midpointMjdTai"] = [60100.5, 60110.5] 

47 return data 

48 

49 

50class _StubApdbQuery: 

51 """Minimal DbQuery stub returning canned sources/forced DataFrames. 

52 

53 Tracks the number of calls so tests can verify the per-diaObject cache. 

54 """ 

55 

56 def __init__(self, sources, forced): 

57 self._sources = sources 

58 self._forced = forced 

59 self.source_calls = 0 

60 self.forced_calls = 0 

61 

62 def load_sources_for_object(self, dia_object_id, exclude_flagged=False, limit=100000): 

63 self.source_calls += 1 

64 self.last_source_kwargs = {"exclude_flagged": exclude_flagged, "limit": limit} 

65 return self._sources.copy() 

66 

67 def load_forced_sources_for_object(self, dia_object_id, exclude_flagged=False, limit=100000): 

68 self.forced_calls += 1 

69 self.last_forced_kwargs = {"exclude_flagged": exclude_flagged, "limit": limit} 

70 return self._forced.copy() 

71 

72 

73def _make_sources_frame(visits, bands, mjds, fluxes, errs, with_err=True): 

74 frame = pd.DataFrame({ 

75 "visit": visits, 

76 "band": bands, 

77 "midpointMjdTai": mjds, 

78 "psfFlux": fluxes, 

79 }) 

80 if with_err: 

81 frame["psfFluxErr"] = errs 

82 return frame 

83 

84 

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

86 """Tests for PlotDiaSourceLightcurveTask.""" 

87 

88 def setUp(self): 

89 bbox = lsst.geom.Box2I(lsst.geom.Point2I(0, 0), lsst.geom.Point2I(100, 100)) 

90 self.centroid = lsst.geom.Point2D(50, 50) 

91 dataset = lsst.meas.base.tests.TestDataset(bbox, crval=skyCenter) 

92 self.scale = 0.3 

93 dataset.addSource(instFlux=1e5, centroid=self.centroid) 

94 self.science, _ = dataset.realize(noise=1000.0, schema=dataset.makeMinimalSchema()) 

95 self.template, _ = dataset.realize(noise=5.0, schema=dataset.makeMinimalSchema()) 

96 self.difference = lsst.afw.image.ExposureF(self.science, deep=True) 

97 self.difference.image -= self.template.image 

98 self.data = _augment_with_object_columns(DATA) 

99 

100 def _make_task(self, sources=None, forced=None, **config_kwargs): 

101 config = plotDiaSourceLightcurve.PlotDiaSourceLightcurveConfig() 

102 for key, value in config_kwargs.items(): 

103 setattr(config, key, value) 

104 query = None 

105 if sources is not None or forced is not None: 

106 query = _StubApdbQuery( 

107 sources if sources is not None else pd.DataFrame(), 

108 forced if forced is not None else pd.DataFrame(), 

109 ) 

110 task = plotDiaSourceLightcurve.PlotDiaSourceLightcurveTask( 

111 config=config, output_path="", apdb_query=query) 

112 return task, query 

113 

114 def _record(self, row): 

115 """Return a single-row numpy.record matching DataFrame ``iloc``.""" 

116 return self.data.iloc[[row]].to_records(index=False)[0] 

117 

118 def test_generate_image_no_apdb(self): 

119 """Without an APDB handle, the lightcurve panel renders a placeholder 

120 but the figure still produces a valid PNG. 

121 """ 

122 task, _ = self._make_task() 

123 cutout = task.generate_image(self.science, self.template, self.difference, 

124 skyCenter, self.scale, 

125 source=self._record(0)) 

126 with PIL.Image.open(cutout) as im: 

127 self.assertGreater(im.width, 0) 

128 self.assertGreater(im.height, 0) 

129 

130 def test_generate_image_with_lightcurve(self): 

131 """With matching sources and forced rows, both groups render.""" 

132 sources = _make_sources_frame( 

133 visits=[1234, 5678, 9999], 

134 bands=["r", "g", "r"], 

135 mjds=[60100.5, 60110.5, 60120.5], 

136 fluxes=[1234.5, 2345.6, 3456.7], 

137 errs=[123.5, 234.5, 345.6], 

138 ) 

139 # Forced has a visit (8888) that does NOT appear in sources — that 

140 # one should be drawn with the forced-only marker. Visit 1234 IS in 

141 # sources, so the forced entry for it should be suppressed. 

142 forced = _make_sources_frame( 

143 visits=[1234, 8888], 

144 bands=["r", "g"], 

145 mjds=[60100.5, 60115.5], 

146 fluxes=[1200.0, 800.0], 

147 errs=[120.0, 80.0], 

148 ) 

149 task, query = self._make_task(sources=sources, forced=forced) 

150 cutout = task.generate_image(self.science, self.template, self.difference, 

151 skyCenter, self.scale, 

152 source=self._record(0)) 

153 self.assertEqual(query.source_calls, 1) 

154 self.assertEqual(query.forced_calls, 1) 

155 with PIL.Image.open(cutout) as im: 

156 self.assertGreater(im.width, 0) 

157 self.assertGreater(im.height, 0) 

158 

159 def test_lightcurve_cache_reuses_query(self): 

160 """Adjacent diaSources on the same diaObject should hit the cache.""" 

161 sources = _make_sources_frame( 

162 visits=[1234, 5678], 

163 bands=["r", "g"], 

164 mjds=[60100.5, 60110.5], 

165 fluxes=[1234.5, 2345.6], 

166 errs=[123.5, 234.5], 

167 ) 

168 forced = pd.DataFrame(columns=["visit", "band", "midpointMjdTai", 

169 "psfFlux", "psfFluxErr"]) 

170 task, query = self._make_task(sources=sources, forced=forced) 

171 for i in range(2): 

172 task.generate_image(self.science, self.template, self.difference, 

173 skyCenter, self.scale, 

174 source=self._record(i)) 

175 # Both sources share the same diaObjectId; the second call must come 

176 # from the cache. 

177 self.assertEqual(query.source_calls, 1) 

178 self.assertEqual(query.forced_calls, 1) 

179 

180 def test_forced_only_dedup_by_visit(self): 

181 """Forced rows whose ``visit`` matches a diaSource are suppressed.""" 

182 sources = _make_sources_frame( 

183 visits=[1234, 5678], 

184 bands=["r", "g"], 

185 mjds=[60100.5, 60110.5], 

186 fluxes=[1234.5, 2345.6], 

187 errs=[123.5, 234.5], 

188 ) 

189 forced = _make_sources_frame( 

190 visits=[1234, 5678, 7777, 8888], 

191 bands=["r", "g", "r", "g"], 

192 mjds=[60100.5, 60110.5, 60112.5, 60115.5], 

193 fluxes=[1200.0, 2400.0, 500.0, 800.0], 

194 errs=[120.0, 240.0, 50.0, 80.0], 

195 ) 

196 task, _ = self._make_task(sources=sources, forced=forced) 

197 forced_only = forced[~forced["visit"].isin(sources["visit"])] 

198 self.assertEqual(set(forced_only["visit"]), {7777, 8888}) 

199 # Also exercise the rendering path to confirm it does not raise. 

200 task.generate_image(self.science, self.template, self.difference, 

201 skyCenter, self.scale, source=self._record(0)) 

202 

203 def test_no_psfFluxErr(self): 

204 """Missing or NaN psfFluxErr should fall back to no error bars.""" 

205 sources = _make_sources_frame( 

206 visits=[1234, 5678], 

207 bands=["r", "g"], 

208 mjds=[60100.5, 60110.5], 

209 fluxes=[1234.5, 2345.6], 

210 errs=None, with_err=False, 

211 ) 

212 forced = _make_sources_frame( 

213 visits=[7777], 

214 bands=["r"], 

215 mjds=[60112.5], 

216 fluxes=[500.0], 

217 errs=[np.nan], with_err=True, 

218 ) 

219 task, _ = self._make_task(sources=sources, forced=forced) 

220 cutout = task.generate_image(self.science, self.template, self.difference, 

221 skyCenter, self.scale, 

222 source=self._record(0)) 

223 with PIL.Image.open(cutout) as im: 

224 self.assertGreater(im.width, 0) 

225 

226 def test_forced_query_skips_exclude_flagged(self): 

227 """DiaForcedSource lacks the diaSource flag columns; the forced 

228 query must always be called with exclude_flagged=False, even when 

229 the config asks to exclude flagged diaSources. 

230 """ 

231 sources = _make_sources_frame( 

232 visits=[1234], bands=["r"], mjds=[60100.5], 

233 fluxes=[1234.5], errs=[123.5]) 

234 forced = pd.DataFrame(columns=["visit", "band", "midpointMjdTai", 

235 "psfFlux", "psfFluxErr"]) 

236 task, query = self._make_task(sources=sources, forced=forced, 

237 lightcurve_exclude_flagged=True) 

238 task.generate_image(self.science, self.template, self.difference, 

239 skyCenter, self.scale, 

240 source=self._record(0)) 

241 self.assertTrue(query.last_source_kwargs["exclude_flagged"]) 

242 self.assertFalse(query.last_forced_kwargs["exclude_flagged"]) 

243 

244 def test_nan_dia_object_id(self): 

245 """A NaN diaObjectId (unassociated diaSource) must not crash; the 

246 lightcurve panel renders its empty placeholder. 

247 """ 

248 data = self.data.copy() 

249 data["diaObjectId"] = data["diaObjectId"].astype(float) 

250 data.loc[0, "diaObjectId"] = np.nan 

251 sources = _make_sources_frame( 

252 visits=[1234], bands=["r"], mjds=[60100.5], 

253 fluxes=[1234.5], errs=[123.5]) 

254 forced = pd.DataFrame(columns=["visit", "band", "midpointMjdTai", 

255 "psfFlux", "psfFluxErr"]) 

256 task, query = self._make_task(sources=sources, forced=forced) 

257 record = data.iloc[[0]].to_records(index=False)[0] 

258 cutout = task.generate_image(self.science, self.template, self.difference, 

259 skyCenter, self.scale, source=record) 

260 # APDB was never queried — diaObjectId is NaN. 

261 self.assertEqual(query.source_calls, 0) 

262 self.assertEqual(query.forced_calls, 0) 

263 with PIL.Image.open(cutout) as im: 

264 self.assertGreater(im.width, 0) 

265 

266 def test_forced_legend_single_black_entry(self): 

267 """Forced points are colored per-band, but contribute a single 

268 black legend entry regardless of band count. 

269 """ 

270 sources = _make_sources_frame( 

271 visits=[1234, 5678], 

272 bands=["r", "g"], 

273 mjds=[60100.5, 60110.5], 

274 fluxes=[1234.5, 2345.6], 

275 errs=[123.5, 234.5], 

276 ) 

277 # Forced-only visits span two bands: should still produce one entry. 

278 forced = _make_sources_frame( 

279 visits=[7777, 8888], 

280 bands=["r", "g"], 

281 mjds=[60112.5, 60115.5], 

282 fluxes=[500.0, 800.0], 

283 errs=[50.0, 80.0], 

284 ) 

285 task, _ = self._make_task(sources=sources, forced=forced) 

286 import matplotlib.pyplot as plt 

287 fig, ax = plt.subplots() 

288 try: 

289 task._draw_lightcurve(ax, sources, forced, 

290 current_source=self._record(0)) 

291 handles, labels = ax.get_legend_handles_labels() 

292 forced_labels = [lbl for lbl in labels if "forced" in lbl] 

293 self.assertEqual(len(forced_labels), 1) 

294 self.assertIn("(n=2)", forced_labels[0]) 

295 # The single forced legend handle should be drawn in black. 

296 forced_handle = handles[labels.index(forced_labels[0])] 

297 self.assertEqual(forced_handle.get_color(), "black") 

298 finally: 

299 plt.close(fig) 

300 

301 def test_njobs_downgraded(self): 

302 """Requesting multiprocessing should be downgraded with a warning.""" 

303 sources = _make_sources_frame( 

304 visits=[1234], bands=["r"], mjds=[60100.5], 

305 fluxes=[1234.5], errs=[123.5]) 

306 forced = pd.DataFrame(columns=["visit", "band", "midpointMjdTai", 

307 "psfFlux", "psfFluxErr"]) 

308 task, _ = self._make_task(sources=sources, forced=forced) 

309 # write_images would normally attempt multiprocessing; the override 

310 # should silently drop njobs to 0 and not raise. 

311 with self.assertLogs(task.log.name, level="WARNING") as ctx: 

312 # Use just one row so we don't need real files on disk; the 

313 # cutouts task will try to look them up via butler_cache and 

314 # fail, but the warning should fire first. 

315 try: 

316 task.write_images(self.data.head(0), butler=None, njobs=4) 

317 except Exception: 

318 pass 

319 self.assertTrue(any("njobs" in msg for msg in ctx.output)) 

320 

321 

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

323 pass 

324 

325 

326def setup_module(module): 

327 lsst.utils.tests.init() 

328 

329 

330if __name__ == "__main__": 330 ↛ 331line 330 didn't jump to line 331 because the condition on line 330 was never true

331 lsst.utils.tests.init() 

332 unittest.main()