Coverage for tests/test_plotDiaSourceLightcurve.py: 97%
148 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-11 11:12 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-11 11:12 +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.afw.image
25import lsst.geom
26import lsst.meas.base.tests
27import lsst.utils.tests
28import numpy as np
29import pandas as pd
30import PIL
32from lsst.analysis.ap import plotDiaSourceLightcurve
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
39DIA_OBJECT_ID = 999999999999000001
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
50class _StubApdbQuery:
51 """Minimal DbQuery stub returning canned sources/forced DataFrames.
53 Tracks the number of calls so tests can verify the per-diaObject cache.
54 """
56 def __init__(self, sources, forced):
57 self._sources = sources
58 self._forced = forced
59 self.source_calls = 0
60 self.forced_calls = 0
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()
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()
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
85class TestPlotDiaSourceLightcurve(lsst.utils.tests.TestCase):
86 """Tests for PlotDiaSourceLightcurveTask."""
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)
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
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]
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)
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)
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)
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))
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)
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"])
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)
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)
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))
322class MemoryTester(lsst.utils.tests.MemoryTestCase):
323 pass
326def setup_module(module):
327 lsst.utils.tests.init()
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()