Coverage for tests/test_promptSourceSchema.py: 88%

50 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-02 09:58 +0000

1# This file is part of pipe_tasks. 

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 

22"""Check that the prompt_source pipeline produces exactly the columns defined 

23by the ``PromptSource`` table in ``sdm_schemas``. Only check column names, not 

24datatypes, since that requires running the detection and measurement task 

25which can't be imported here. 

26""" 

27 

28import os 

29import unittest 

30 

31import numpy as np 

32from astropy.table import Table 

33import importlib.resources 

34 

35import lsst.utils.tests 

36from lsst.utils import getPackageDir 

37 

38from lsst.pipe.tasks.postprocess import TransformSourceTableTask, TransformSourceTableConfig 

39from lsst.pipe.tasks.schemaUtils import readSdmSchemaFile 

40from lsst.pipe.tasks.split_primary import SplitPrimaryTask 

41 

42FUNCTOR_FILE = os.path.join(getPackageDir("pipe_tasks"), "schemas", "prompt_source.yaml") 

43 

44SCHEMA_FILE = importlib.resources.files("lsst.sdm.schemas") / "ap_extra.yaml" 

45TABLE_NAME = "PromptSource" 

46 

47 

48class PromptSourceSchemaTestCase(lsst.utils.tests.TestCase): 

49 """Check the persisted prompt_source columns against the PromptSource DDL. 

50 """ 

51 

52 @classmethod 

53 def setUpClass(cls): 

54 super().setUpClass() 

55 

56 schemaFile = os.path.expandvars(SCHEMA_FILE) 

57 if not os.path.exists(schemaFile): 57 ↛ 58line 57 didn't jump to line 58 because the condition on line 57 was never true

58 raise unittest.SkipTest(f"SDM schema file not available: {SCHEMA_FILE}") 

59 

60 schema = readSdmSchemaFile(schemaFile) 

61 if TABLE_NAME not in schema: 61 ↛ 62line 61 didn't jump to line 62 because the condition on line 61 was never true

62 raise ValueError(f"Table {TABLE_NAME!r} not in {schemaFile}.") 

63 

64 cls.schemaColumns = {column.name for column in schema[TABLE_NAME].columns} 

65 

66 config = TransformSourceTableConfig() 

67 config.functorFile = FUNCTOR_FILE 

68 transformTask = TransformSourceTableTask(config=config) 

69 producedColumns = set(transformTask.funcs.funcDict) | set(config.columnsFromDataId) 

70 

71 # The final prompt_source table is output by splitPromptSource, which is configured 

72 # in the pipeline to drop sky sources and remove the corresponding column. 

73 splitConfig = SplitPrimaryTask.ConfigClass() 

74 splitConfig.discard_primary_columns = ["sky_source"] 

75 splitTask = SplitPrimaryTask(config=splitConfig) 

76 

77 # SplitPrimaryTask only needs the boolean primary-flag column to mask 

78 # rows; the other columns are carried through (or dropped) by name, so 

79 # a minimal two-row table with placeholder values is sufficient. 

80 data = {} 

81 for name in producedColumns: 

82 if name == splitConfig.primary_flag_column: 

83 data[name] = np.array([True, False]) 

84 else: 

85 data[name] = np.zeros(2) 

86 full = Table(data) 

87 

88 cls.persistedColumns = set(splitTask.run(full=full).primary.colnames) 

89 

90 def testColumnsConform(self): 

91 """The persisted prompt_source columns must equal the schema exactly. 

92 """ 

93 missing = self.schemaColumns - self.persistedColumns 

94 extra = self.persistedColumns - self.schemaColumns 

95 

96 message = ( 

97 f"persisted prompt_source columns do not conform to the " 

98 f"{TABLE_NAME} schema.\n" 

99 f" In schema but not persisted ({len(missing)}): " 

100 f"{sorted(missing)}\n" 

101 f" Persisted but not in schema ({len(extra)}): " 

102 f"{sorted(extra)}" 

103 ) 

104 self.assertEqual(self.persistedColumns, self.schemaColumns, message) 

105 

106 

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

108 pass 

109 

110 

111def setup_module(module): 

112 lsst.utils.tests.init() 

113 

114 

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

116 lsst.utils.tests.init() 

117 unittest.main()