Coverage for tests/test_loadDiaCatalogs.py: 98%

120 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-25 15:44 -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/>. 

21 

22import logging 

23import os 

24import astropy.units 

25import numpy as np 

26import pandas as pd 

27import tempfile 

28import unittest 

29import yaml 

30 

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 

39 

40 

41class TestLoadDiaCatalogs(unittest.TestCase): 

42 

43 def setUp(self): 

44 # Create an instance of random generator with fixed seed. 

45 rng = np.random.default_rng(1234) 

46 

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) 

51 

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) 

58 

59 self.exposure = makeExposure(False, False) 

60 self.regionTime = makeRegionTime(exposure=self.exposure) 

61 self.dateTime = getMidpointFromTimespan(self.regionTime.timespan) 

62 

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) 

68 

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) 

76 

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

80 

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 

86 

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) 

93 

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

98 

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

108 

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

121 

122 def testLoadDiaSources(self): 

123 """Test that the correct number of diaSources are loaded. 

124 

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) 

130 

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

137 

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] 

144 

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

152 

153 

154class TestDropDuplicateRows(unittest.TestCase): 

155 """Tests of the deduplication shared by LoadDiaCatalogsTask and 

156 DiaPipelineTask. 

157 """ 

158 

159 def setUp(self): 

160 self.log = logging.getLogger("TestDropDuplicateRows") 

161 

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

166 

167 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log) 

168 

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

172 

173 def test_keepsFirstRowWhole(self): 

174 """Test that deduplication keeps whole rows. 

175 

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

182 

183 with self.assertLogs(self.log.name, level="WARNING"): 

184 result = dropDuplicateRows(catalog, "diaObjectId", "DiaObjects", self.log) 

185 

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) 

193 

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

202 

203 with self.assertLogs(self.log.name, level="WARNING"): 

204 result = dropDuplicateRows(catalog, index, "DiaSources", self.log) 

205 

206 self.assertEqual(list(result.index.names), index) 

207 self.assertFalse(result.index.has_duplicates) 

208 self.assertEqual(list(result["value"]), [1, 3]) 

209 

210 def test_inputLeftUnchanged(self): 

211 """Test that the input catalog is not modified. 

212 

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) 

220 

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) 

226 

227 self.assertIsNot(result, catalog) 

228 self.assertIsNone(catalog.index.name) 

229 self.assertEqual(list(catalog["diaObjectId"]), ids) 

230 

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

234 

235 

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

237 pass 

238 

239 

240def setup_module(module): 

241 lsst.utils.tests.init() 

242 

243 

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