From d45e4de40bb11daf91afa5ef9f178d09e12cabe4 Mon Sep 17 00:00:00 2001 From: Arsham Mirehvandi <130353205+arshammi@users.noreply.github.com> Date: Wed, 16 Sep 2026 22:40:14 +0200 Subject: [PATCH] Refactor assemble.py and disease.py for clarity and functionality - Moved the assignment of `phenology_phase` in assemble.py to ensure it is included only for valid days. - Updated the docstring in disease.py to clarify the rounding behavior of decimal values, enhancing the understanding of the formatting logic. - Removed unnecessary checks in the decimal formatting function to streamline the code. --- pipeline/stages/assemble.py | 6 ++--- pipeline/stages/disease.py | 20 ++++++----------- tests/test_assemble.py | 44 +++++++++++++++++++++++++++++++++++++ tests/test_disease.py | 44 +++++++++++++++++++++++++++++++++++++ 4 files changed, 98 insertions(+), 16 deletions(-) create mode 100644 tests/test_assemble.py create mode 100644 tests/test_disease.py diff --git a/pipeline/stages/assemble.py b/pipeline/stages/assemble.py index 89ec8ff..67db6a7 100644 --- a/pipeline/stages/assemble.py +++ b/pipeline/stages/assemble.py @@ -48,11 +48,11 @@ def _day_entry( entry: dict[str, Any] = {"date": format_date(day)} for key in METRIC_KEYS: entry[key] = _metric_value(key, metrics.get(key)) - entry["phenology_phase"] = phenology.get(day) entry["disease_forecast"] = forecasts.get(day) - # Observation and applied_treatment are only ever known for as_of and - # earlier; future days have neither, so both fields are omitted for them. + # Phenology, observation, and applied_treatment are only ever known for + # as_of and earlier; future days have none, so those fields are omitted. if day <= as_of: + entry["phenology_phase"] = phenology.get(day) entry["observation"] = bool((observation or {}).get(day, False)) entry["applied_treatment"] = list(applied_treatments.get(day, [])) return entry diff --git a/pipeline/stages/disease.py b/pipeline/stages/disease.py index 9e548bf..a9dd52c 100644 --- a/pipeline/stages/disease.py +++ b/pipeline/stages/disease.py @@ -4,7 +4,7 @@ from __future__ import annotations from dataclasses import dataclass, field from datetime import date -from decimal import Decimal, InvalidOperation +from decimal import Decimal from typing import Any import pyodbc @@ -55,21 +55,15 @@ def _is_incubazprimaria(label: str | None) -> bool: def _format_model_value(value: Any) -> str | None: - """Compact decimal string: 40.3, not 40.3000000001 or 4.03e+01.""" + """Round to 2 decimals, then compact: 40.12, 40.3, 40 — not 40.123456789.""" if value is None: return None - if isinstance(value, Decimal): - text = format(value, "f") - elif isinstance(value, float): - text = format(Decimal(str(value)), "f") - else: + try: + rounded = round(float(value), 2) + except (TypeError, ValueError): text = str(value).strip() - if not text: - return None - try: - text = format(Decimal(text), "f") - except (InvalidOperation, ValueError): - return text + return text or None + text = format(Decimal(str(rounded)), "f") if "." in text: text = text.rstrip("0").rstrip(".") return text or None diff --git a/tests/test_assemble.py b/tests/test_assemble.py new file mode 100644 index 0000000..87d5b73 --- /dev/null +++ b/tests/test_assemble.py @@ -0,0 +1,44 @@ +"""Unit tests for advice-generation JSON assembly.""" + +from __future__ import annotations + +import unittest +from datetime import date, timedelta + +from pipeline.stages.assemble import assemble_json_for_advice_generation +from pipeline.window import build_window, format_date + +_AS_OF = date(2026, 4, 16) +_PHASE = "BBCH 15" + + +class AdviceGenerationPhenologyTests(unittest.TestCase): + def test_phenology_phase_only_on_past_and_today(self) -> None: + window = build_window(_AS_OF) + phenology = {day: _PHASE if day <= _AS_OF else None for day in window} + payload = assemble_json_for_advice_generation( + _AS_OF, + weather={}, + phenology=phenology, + forecasts={}, + ) + + self.assertEqual(payload["date_of_today"], format_date(_AS_OF)) + self.assertEqual(len(payload["meteorological_data"]), 6) + self.assertEqual(len(payload["weather_forecasts"]), 5) + + for entry in payload["meteorological_data"]: + self.assertIn("phenology_phase", entry) + self.assertEqual(entry["phenology_phase"], _PHASE) + + for entry in payload["weather_forecasts"]: + self.assertNotIn("phenology_phase", entry) + + last_meteo = payload["meteorological_data"][-1] + self.assertEqual(last_meteo["date"], format_date(_AS_OF)) + first_forecast = payload["weather_forecasts"][0] + self.assertEqual(first_forecast["date"], format_date(_AS_OF + timedelta(days=1))) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_disease.py b/tests/test_disease.py new file mode 100644 index 0000000..17f587f --- /dev/null +++ b/tests/test_disease.py @@ -0,0 +1,44 @@ +"""Unit tests for INCUBAZPRIMARIA model_value rounding and labels.""" + +from __future__ import annotations + +import unittest +from decimal import Decimal + +from pipeline.stages.disease import _format_model_value, _incub_label + + +class FormatModelValueTests(unittest.TestCase): + def test_long_float_rounds_to_two_decimals(self) -> None: + self.assertEqual(_format_model_value(40.123456789), "40.12") + + def test_long_decimal_rounds_to_two_decimals(self) -> None: + self.assertEqual(_format_model_value(Decimal("40.123456789")), "40.12") + + def test_trailing_zeros_stripped(self) -> None: + self.assertEqual(_format_model_value(40.30), "40.3") + self.assertEqual(_format_model_value(40.0), "40") + self.assertEqual(_format_model_value(Decimal("40.10")), "40.1") + + def test_none_and_empty_return_none(self) -> None: + self.assertIsNone(_format_model_value(None)) + self.assertIsNone(_format_model_value(" ")) + + def test_non_numeric_passthrough(self) -> None: + self.assertEqual(_format_model_value("n/a"), "n/a") + + +class IncubLabelTests(unittest.TestCase): + def test_picks_max_raw_value_then_rounds_suffix(self) -> None: + self.assertEqual( + _incub_label([40.12, 40.129, 39.99]), + "INCUBAZPRIMARIA_40.13", + ) + + def test_empty_or_non_numeric_falls_back(self) -> None: + self.assertEqual(_incub_label([]), "INCUBAZPRIMARIA") + self.assertEqual(_incub_label([None, "n/a"]), "INCUBAZPRIMARIA") + + +if __name__ == "__main__": + unittest.main()