Coverage for backend/django/Economics/costing/cost_curves/serializers.py: 81%

148 statements  

« prev     ^ index     » next       coverage.py v7.10.7, created at 2026-07-22 05:22 +0000

1from drf_spectacular.utils import extend_schema_field 

2from rest_framework import serializers 

3 

4from Economics.costing.models import CostCurve 

5from Economics.shared.choices import CostCurveEvaluationKind 

6from Economics.shared.serializers import UnitOptionSerializer 

7from Economics.costing.cost_curves.catalog import cost_curve_category_requires_subtype 

8from Economics.costing.cost_curves.driver_specs import ( 

9 CostCurveDiscreteVariant, 

10 CostCurveDiscreteVariantPayload, 

11 CostCurveDriverSpec, 

12 CostCurveDriverSpecPayload, 

13 CostCurveDriverSpecRead, 

14 driver_spec_read_payload, 

15 normalize_discrete_variants, 

16 normalize_required_driver_specs, 

17) 

18from Economics.costing.cost_curves.unit_options import ( 

19 cost_curve_common_driver_unit_options, 

20 cost_curve_output_unit_options, 

21) 

22from Economics.formulas.engine.core import FormulaError 

23from Economics.formulas.engine.parsing import parse_cost_expression 

24 

25 

26@extend_schema_field( 

27 CostCurveDriverSpecRead.model_json_schema(), 

28 component_name="CostCurveDriverSpec", 

29) 

30class CostCurveDriverSpecField(serializers.JSONField): 

31 """JSON transport field whose OpenAPI component comes from Pydantic.""" 

32 

33 def to_representation(self, value): 

34 try: 

35 return driver_spec_read_payload(CostCurveDriverSpec.model_validate(value)) 

36 except ValueError: 

37 return super().to_representation(value) 

38 

39 

40@extend_schema_field( 

41 CostCurveDiscreteVariant.model_json_schema(), 

42 component_name="CostCurveDiscreteVariant", 

43) 

44class CostCurveDiscreteVariantField(serializers.JSONField): 

45 """JSON transport field whose OpenAPI component comes from Pydantic.""" 

46 

47 def to_representation(self, value): 

48 if isinstance(value, CostCurveDiscreteVariant): 48 ↛ 50line 48 didn't jump to line 50 because the condition on line 48 was always true

49 return value.model_dump(mode="json") 

50 return super().to_representation(value) 

51 

52 

53class CostCurveSerializer(serializers.ModelSerializer): 

54 evaluation_kind = serializers.ChoiceField( 

55 choices=CostCurveEvaluationKind.choices, 

56 required=True, 

57 ) 

58 output_unit_options = serializers.SerializerMethodField() 

59 required_driver_specs = serializers.ListField( 

60 child=CostCurveDriverSpecField(), 

61 required=True, 

62 ) 

63 discrete_variants = serializers.ListField( 

64 child=CostCurveDiscreteVariantField(), 

65 required=True, 

66 ) 

67 

68 class Meta: 

69 model = CostCurve 

70 exclude = ("project",) 

71 read_only_fields = ("id", "created_at", "updated_at") 

72 

73 @extend_schema_field(UnitOptionSerializer(many=True)) 

74 def get_output_unit_options(self, instance) -> list[dict[str, str]]: 

75 return cost_curve_output_unit_options(instance.output_unit, currency=instance.currency or "NZD") 

76 

77 

78class CostCurveAuthoringSerializer(serializers.ModelSerializer): 

79 """Expression-first write contract for user-authored cost curves.""" 

80 

81 save_as_template = serializers.BooleanField(write_only=True, required=False, default=False) 

82 evaluation_kind = serializers.ChoiceField( 

83 choices=CostCurveEvaluationKind.choices, 

84 required=True, 

85 ) 

86 required_driver_specs = serializers.ListField( 

87 child=CostCurveDriverSpecField(), 

88 required=True, 

89 ) 

90 discrete_variants = serializers.ListField( 

91 child=CostCurveDiscreteVariantField(), 

92 required=True, 

93 ) 

94 removed_formula_fields = frozenset( 

95 { 

96 "expression_type", 

97 "coefficient_a", 

98 "coefficient_b", 

99 "coefficient_c", 

100 "exponent", 

101 "manual_quote_amount", 

102 } 

103 ) 

104 

105 class Meta: 

106 model = CostCurve 

107 fields = ( 

108 "id", 

109 "curve_key", 

110 "name", 

111 "equipment_category", 

112 "equipment_subtype", 

113 "cost_basis", 

114 "evaluation_kind", 

115 "output_unit", 

116 "expression_text", 

117 "required_driver_specs", 

118 "discrete_variants", 

119 "valid_min", 

120 "valid_max", 

121 "valid_range_note", 

122 "currency", 

123 "basis_date", 

124 "basis_index_name", 

125 "basis_index_value", 

126 "source_document_title", 

127 "source_page", 

128 "source_figure", 

129 "source_data_origin", 

130 "source_range_precision", 

131 "source_license_status", 

132 "source_reference", 

133 "applicability_warning", 

134 "notes", 

135 "active", 

136 "save_as_template", 

137 "created_at", 

138 "updated_at", 

139 ) 

140 read_only_fields = ("id", "created_at", "updated_at") 

141 

142 def to_internal_value(self, data): 

143 if isinstance(data, dict): 143 ↛ 152line 143 didn't jump to line 152 because the condition on line 143 was always true

144 removed_fields = sorted(self.removed_formula_fields.intersection(data.keys())) 

145 if removed_fields: 

146 raise serializers.ValidationError( 

147 { 

148 field: "Cost curves are expression-only. Use expression_text." 

149 for field in removed_fields 

150 } 

151 ) 

152 return super().to_internal_value(data) 

153 

154 def validate(self, attrs): 

155 attrs = super().validate(attrs) 

156 evaluation_kind = attrs.get( 

157 "evaluation_kind", 

158 getattr(self.instance, "evaluation_kind", CostCurveEvaluationKind.EXPRESSION), 

159 ) 

160 expression_text = attrs.get("expression_text", getattr(self.instance, "expression_text", "")) 

161 if evaluation_kind == CostCurveEvaluationKind.EXPRESSION and not expression_text: 

162 raise serializers.ValidationError({"expression_text": "Expression cost curves require expression_text."}) 

163 if "required_driver_specs" in attrs: 

164 attrs["required_driver_specs"] = _normalized_required_driver_specs(attrs["required_driver_specs"]) 

165 if "discrete_variants" in attrs: 

166 attrs["discrete_variants"] = _normalized_discrete_variants(attrs["discrete_variants"]) 

167 self._validate_formula_contract(attrs, evaluation_kind=evaluation_kind, expression_text=expression_text) 

168 self._validate_equipment_subtype_requirement(attrs) 

169 return attrs 

170 

171 def _validate_formula_contract(self, attrs, *, evaluation_kind: str, expression_text: str) -> None: 

172 specs_payload = attrs.get( 

173 "required_driver_specs", 

174 getattr(self.instance, "required_driver_specs", []), 

175 ) 

176 formula_specs = [spec for spec in specs_payload if spec["role"] == "formula_input"] 

177 selector_specs = [spec for spec in specs_payload if spec["role"] == "discrete_selector"] 

178 variable_symbols = [spec["variable_symbol"] for spec in formula_specs] 

179 try: 

180 if evaluation_kind == CostCurveEvaluationKind.EXPRESSION: 180 ↛ 182line 180 didn't jump to line 182 because the condition on line 180 was always true

181 parse_cost_expression(expression_text, variable_symbols=variable_symbols) 

182 elif evaluation_kind == CostCurveEvaluationKind.DISCRETE_FAMILY: 

183 expected_selector_keys = {spec["key"] for spec in selector_specs} 

184 if not expected_selector_keys: 

185 raise FormulaError( 

186 "missing_discrete_selectors", 

187 "Discrete-family curves require at least one selector input.", 

188 ) 

189 if len(expected_selector_keys) > 1: 

190 raise FormulaError( 

191 "unsupported_multi_selector_discrete_family", 

192 "Discrete-family curves currently support exactly one capacity selector.", 

193 context={"selector_keys": sorted(expected_selector_keys)}, 

194 ) 

195 variants = attrs.get("discrete_variants", getattr(self.instance, "discrete_variants", [])) 

196 if not variants: 

197 raise FormulaError( 

198 "missing_discrete_variants", 

199 "Discrete-family curves require at least one variant.", 

200 ) 

201 for variant in variants: 

202 selector_keys = set(variant["selector_values"]) 

203 if selector_keys != expected_selector_keys: 

204 raise FormulaError( 

205 "invalid_discrete_variant_selectors", 

206 "Discrete variant selector values must exactly match selector inputs.", 

207 context={ 

208 "variant_key": variant["key"], 

209 "missing_selectors": sorted(expected_selector_keys - selector_keys), 

210 "extra_selectors": sorted(selector_keys - expected_selector_keys), 

211 }, 

212 ) 

213 parse_cost_expression(variant["expression_text"], variable_symbols=variable_symbols) 

214 except FormulaError as exc: 

215 raise serializers.ValidationError({"expression_text": exc.message, "context": exc.context}) from exc 

216 

217 def _validate_equipment_subtype_requirement(self, attrs) -> None: 

218 category = attrs.get("equipment_category", getattr(self.instance, "equipment_category", "")) 

219 subtype = attrs.get("equipment_subtype", getattr(self.instance, "equipment_subtype", "")) 

220 if not category or str(subtype).strip(): 

221 return 

222 

223 if cost_curve_category_requires_subtype(str(category), self._peer_cost_curves()): 

224 raise serializers.ValidationError( 

225 { 

226 "equipment_subtype": ( 

227 "Choose an equipment subtype before saving a cost curve for this category." 

228 ) 

229 } 

230 ) 

231 

232 def _peer_cost_curves(self): 

233 return CostCurve.objects.all() 

234 

235 

236class CostCurveTemplateSerializer(serializers.Serializer): 

237 output_unit_options = serializers.SerializerMethodField() 

238 required_driver_specs = serializers.ListField(child=CostCurveDriverSpecField()) 

239 discrete_variants = serializers.ListField(child=CostCurveDiscreteVariantField()) 

240 value = serializers.CharField() 

241 name = serializers.CharField() 

242 equipment_category = serializers.CharField() 

243 equipment_subtype = serializers.CharField(allow_blank=True) 

244 cost_basis = serializers.CharField() 

245 evaluation_kind = serializers.CharField() 

246 output_unit = serializers.CharField() 

247 expression_text = serializers.CharField(allow_blank=True) 

248 valid_min = serializers.CharField(allow_blank=True) 

249 valid_max = serializers.CharField(allow_blank=True) 

250 valid_range_note = serializers.CharField(allow_blank=True) 

251 currency = serializers.CharField() 

252 basis_date = serializers.CharField(allow_blank=True) 

253 basis_index_name = serializers.CharField(allow_blank=True) 

254 basis_index_value = serializers.CharField(allow_blank=True) 

255 source_document_title = serializers.CharField(allow_blank=True) 

256 source_page = serializers.CharField(allow_blank=True) 

257 source_figure = serializers.CharField(allow_blank=True) 

258 source_data_origin = serializers.CharField(allow_blank=True) 

259 source_range_precision = serializers.CharField(allow_blank=True) 

260 source_license_status = serializers.CharField(allow_blank=True) 

261 source_reference = serializers.CharField(allow_blank=True) 

262 notes = serializers.CharField(allow_blank=True) 

263 applicability_warning = serializers.CharField(allow_blank=True) 

264 active = serializers.BooleanField() 

265 

266 @extend_schema_field(UnitOptionSerializer(many=True)) 

267 def get_output_unit_options(self, instance) -> list[dict[str, str]]: 

268 return cost_curve_output_unit_options(instance.output_unit, currency=instance.currency or "NZD") 

269 

270 

271def _normalized_required_driver_specs(specs) -> list[CostCurveDriverSpecPayload]: 

272 """Route DRF writes through the Pydantic driver-spec contract.""" 

273 try: 

274 return normalize_required_driver_specs(specs) 

275 except ValueError as exc: 

276 raise serializers.ValidationError({"required_driver_specs": str(exc)}) from exc 

277 

278 

279 

280def _normalized_discrete_variants(variants) -> list[CostCurveDiscreteVariantPayload]: 

281 """Route DRF writes through the Pydantic discrete-variant contract.""" 

282 try: 

283 return normalize_discrete_variants(variants) 

284 except ValueError as exc: 

285 raise serializers.ValidationError({"discrete_variants": str(exc)}) from exc 

286 

287 

288class CostCurveEquipmentCategorySerializer(serializers.Serializer): 

289 value = serializers.CharField() 

290 label = serializers.CharField() 

291 subtypes = serializers.ListField(child=serializers.CharField()) 

292 driver_unit_options = serializers.SerializerMethodField() 

293 templates = CostCurveTemplateSerializer(many=True) 

294 

295 @extend_schema_field(UnitOptionSerializer(many=True)) 

296 def get_driver_unit_options(self, _instance) -> list[dict[str, str]]: 

297 return cost_curve_common_driver_unit_options()