Coverage for tests/test_loadDiaCatalogs.py: 98%
120 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-30 04:30 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-30 04:30 -0700
1# This file is part of ap_association.
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 logging
23import os
24import astropy.units
25import numpy as np
26import pandas as pd
27import tempfile
28import unittest
29import yaml
31from lsst.ap.association import LoadDiaCatalogsTask
32from lsst.ap.association.loadDiaCatalogs import dropDuplicateRows
33from lsst.ap.association.utils import getMidpointFromTimespan, readSchemaFromApdb
34from lsst.dax.apdb import Apdb, ApdbSql, ApdbTables
35from lsst.resources import ResourcePath
36import lsst.utils.tests
37from utils_tests import makeExposure, makeDiaObjects, makeDiaSources, makeDiaForcedSources, makeRegionTime, \
38 getRegion
41class TestLoadDiaCatalogs(unittest.TestCase):
43 def setUp(self):
44 # Create an instance of random generator with fixed seed.
45 rng = np.random.default_rng(1234)
47 self.db_file_fd, self.db_file = tempfile.mkstemp(
48 dir=os.path.dirname(__file__))
49 self.addCleanup(os.remove, self.db_file)
50 self.addCleanup(os.close, self.db_file_fd)
52 self.apdbConfig = ApdbSql.init_database(db_url="sqlite:///" + self.db_file)
53 self.config_file = tempfile.NamedTemporaryFile()
54 self.addCleanup(self.config_file.close)
55 self.apdbConfig.save(self.config_file.name)
56 self.apdb = Apdb.from_config(self.apdbConfig)
57 self.schema = readSchemaFromApdb(self.apdb)
59 self.exposure = makeExposure(False, False)
60 self.regionTime = makeRegionTime(exposure=self.exposure)
61 self.dateTime = getMidpointFromTimespan(self.regionTime.timespan)
63 self.diaObjects = makeDiaObjects(20, self.exposure, rng)
64 self.diaSources = makeDiaSources(
65 100, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, rng)
66 self.diaForcedSources = makeDiaForcedSources(
67 200, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, rng)
69 # Store the test diaSources as though they were observed a month before
70 # the current exposure.
71 dateTime = self.regionTime.timespan.begin.tai - 30 * astropy.units.day
72 self.apdb.store(dateTime,
73 self.diaObjects,
74 self.diaSources,
75 self.diaForcedSources)
77 # These columns are not in the DPDD, yet do appear in DiaSource.yaml.
78 # We don't need to check them against the default APDB schema.
79 self.ignoreColumns = ["band", "bboxSize", "isDipole", "flags"]
81 def _makeConfig(self, **kwargs):
82 config = LoadDiaCatalogsTask.ConfigClass()
83 config.apdb_config_url = self.config_file.name
84 config.update(**kwargs)
85 return config
87 def testRun(self):
88 """Test the full run method for the loader.
89 """
90 diaConfig = self._makeConfig()
91 diaLoader = LoadDiaCatalogsTask(config=diaConfig)
92 result = diaLoader.run(self.regionTime)
94 self.assertEqual(len(result.diaObjects), len(self.diaObjects))
95 self.assertEqual(len(result.diaSources), len(self.diaSources))
96 self.assertEqual(len(result.diaForcedSources),
97 len(self.diaForcedSources))
99 def testLoadDiaObjects(self):
100 """Test that the correct number of diaObjects are loaded.
101 """
102 diaConfig = self._makeConfig()
103 diaLoader = LoadDiaCatalogsTask(config=diaConfig)
104 region = getRegion(self.exposure)
105 diaObjects = diaLoader.loadDiaObjects(region,
106 self.schema)
107 self.assertEqual(len(diaObjects), len(self.diaObjects))
109 def testLoadDiaForcedSources(self):
110 """Test that the correct number of diaForcedSources are loaded.
111 """
112 diaConfig = self._makeConfig()
113 diaLoader = LoadDiaCatalogsTask(config=diaConfig)
114 region = getRegion(self.exposure)
115 diaForcedSources = diaLoader.loadDiaForcedSources(
116 self.diaObjects,
117 region,
118 self.dateTime,
119 self.schema)
120 self.assertEqual(len(diaForcedSources), len(self.diaForcedSources))
122 def testLoadDiaSources(self):
123 """Test that the correct number of diaSources are loaded.
125 Also check that they can be properly loaded both by location and
126 ``diaObjectId``.
127 """
128 diaConfig = self._makeConfig()
129 diaLoader = LoadDiaCatalogsTask(config=diaConfig)
131 region = getRegion(self.exposure)
132 diaSources = diaLoader.loadDiaSources(self.diaObjects,
133 region,
134 self.dateTime,
135 self.schema)
136 self.assertEqual(len(diaSources), len(self.diaSources))
138 def test_apdbSchema(self):
139 """Test that the default DiaSource schema from dax_apdb agrees with the
140 column names defined here in ap_association/data/DiaSource.yaml.
141 """
142 tableDef = self.apdb.tableDef(ApdbTables.DiaSource)
143 apdbSchemaColumns = [column.name for column in tableDef.columns]
145 functorFile = ResourcePath("resource://lsst.ap.association/resources/data/DiaSource.yaml")
146 with functorFile.open("r") as yaml_stream:
147 diaSourceFunctor = yaml.safe_load_all(yaml_stream)
148 for functor in diaSourceFunctor:
149 diaSourceColumns = [column for column in list(functor['funcs'].keys())
150 if column not in self.ignoreColumns]
151 self.assertLess(set(diaSourceColumns), set(apdbSchemaColumns))
154class TestDropDuplicateRows(unittest.TestCase):
155 """Tests of the deduplication shared by LoadDiaCatalogsTask and
156 DiaPipelineTask.
157 """
159 def setUp(self):
160 self.log = logging.getLogger("TestDropDuplicateRows")
162 def test_noDuplicates(self):
163 """Test that an already-unique catalog is only indexed.
164 """
165 catalog = pd.DataFrame({"diaObjectId": [3, 1, 2], "value": [30, 10, 20]})
167 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log)
169 self.assertEqual(result.index.name, "diaObjectId")
170 # Order must follow the Apdb, not the sorted index.
171 self.assertEqual(list(result["diaObjectId"]), [3, 1, 2])
173 def test_keepsFirstRowWhole(self):
174 """Test that deduplication keeps whole rows.
176 Combining values across duplicates, as `groupby().first()` does,
177 would produce a row that the Apdb never held.
178 """
179 catalog = pd.DataFrame({"diaObjectId": [2, 1, 1],
180 "a": [9.0, np.nan, 5.0],
181 "b": [1, 2, 3]})
183 with self.assertLogs(self.log.name, level="WARNING"):
184 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log)
186 self.assertFalse(result.index.has_duplicates)
187 self.assertEqual(list(result["diaObjectId"]), [2, 1])
188 # The row kept for diaObjectId=1 is the first one that is whole
189 # (including the NaN) rather than "a" from one duplicate and "b" from
190 # the other.
191 self.assertTrue(np.isnan(result.loc[1, "a"]))
192 self.assertEqual(result.loc[1, "b"], 2)
194 def test_multiIndex(self):
195 """Test deduplication on the compound DiaSource index.
196 """
197 index = ["diaObjectId", "band", "diaSourceId"]
198 catalog = pd.DataFrame({"diaObjectId": [1, 1, 1],
199 "band": ["g", "g", "r"],
200 "diaSourceId": [10, 10, 11],
201 "value": [1, 2, 3]})
203 with self.assertLogs(self.log.name, level="WARNING"):
204 result = dropDuplicateRows(catalog, index, "DiaSources", self.log)
206 self.assertEqual(list(result.index.names), index)
207 self.assertFalse(result.index.has_duplicates)
208 self.assertEqual(list(result["value"]), [1, 3])
210 def test_inputLeftUnchanged(self):
211 """Test that the input catalog is not modified.
213 The result is a new catalog that owns its data, so indexing it or
214 writing to it must leave the caller's catalog alone.
215 """
216 for label, ids in (("without duplicates", [3, 1, 2]), ("with duplicates", [1, 1, 2])):
217 with self.subTest(label):
218 catalog = pd.DataFrame({"diaObjectId": ids, "value": [10, 20, 30]})
219 hasDuplicates = len(set(ids)) < len(ids)
221 if hasDuplicates:
222 with self.assertLogs(self.log.name, level="WARNING"):
223 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log)
224 else:
225 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log)
227 self.assertIsNot(result, catalog)
228 self.assertIsNone(catalog.index.name)
229 self.assertEqual(list(catalog["diaObjectId"]), ids)
231 # The result owns its data, so this write stays local to it.
232 result.loc[result.index[0], "value"] = 999
233 self.assertEqual(list(catalog["value"]), [10, 20, 30])
236class MemoryTester(lsst.utils.tests.MemoryTestCase):
237 pass
240def setup_module(module):
241 lsst.utils.tests.init()
244if __name__ == "__main__": 244 ↛ 245line 244 didn't jump to line 245 because the condition on line 244 was never true
245 lsst.utils.tests.init()
246 unittest.main()