""" era5_data_pipeline.py ===================== """ from __future__ import annotations import hashlib import json import logging import math import os import random from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any, Dict, List, Optional, Tuple import zone_observation as _zo assert _zo.SCHEMA_VERSION == 3, ( f"Schema mismatch: expected 3, got {_zo.SCHEMA_VERSION}" ) from zone_observation import ( BasinContext, DataSource, EpisodeContext, ForecastConfig, GeoPolygon, ZoneObs, derive_helio_regime, make_synthetic_basin_context, make_synthetic_episode_context, make_synthetic_zone_obs, _stable_seed, ) logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Optional dependencies # --------------------------------------------------------------------------- try: import requests REQUESTS_AVAILABLE = True except ImportError: REQUESTS_AVAILABLE = False logger.warning("requests not installed — Open-Meteo unavailable") try: import cdsapi CDSAPI_AVAILABLE = True except ImportError: CDSAPI_AVAILABLE = False logger.info("cdsapi not installed — ERA5 will fall back to synthetic") try: import numpy as np NUMPY_AVAILABLE = True except ImportError: NUMPY_AVAILABLE = False try: import ee # Google Earth Engine — IMERG / CHIRPS / SMAP access EE_AVAILABLE = True except ImportError: EE_AVAILABLE = False logger.info("earthengine-api not installed — satellite sources unavailable") _EE_INITIALIZED = False def _ensure_ee_initialized() -> None: """Lazily call ee.Initialize() once per process. NOTE: Earth Engine requires a registered Google Cloud project and prior `earthengine authenticate` (or a service-account key) — this cannot be exercised in an offline/sandboxed environment. This function is defensive on purpose: any failure here propagates to the caller's try/except in fetch_zone_obs(), which falls back through the same ERA5 -> Open-Meteo -> synthetic chain as every other source. It has been validated for correct control flow, but the actual GEE calls in _fetch_imerg/_fetch_smap below have NOT been exercised against a live Earth Engine backend — verify against a real authenticated project before relying on them in production. """ global _EE_INITIALIZED if _EE_INITIALIZED: return if not EE_AVAILABLE: raise RuntimeError("earthengine-api not installed") project = os.environ.get("EARTHENGINE_PROJECT") if project: ee.Initialize(project=project) else: ee.Initialize() _EE_INITIALIZED = True # --------------------------------------------------------------------------- # Zone registry # --------------------------------------------------------------------------- _ZONE_REGISTRY: Dict[str, GeoPolygon] = {} def register_zone(polygon: GeoPolygon) -> None: """Register a sourcing zone polygon for lat/lon resolution.""" _ZONE_REGISTRY[polygon.zone_id] = polygon logger.info( f"Registered zone {polygon.zone_id} centroid={polygon.centroid}" ) def _resolve_latlon(zone_id: str) -> Tuple[float, float]: if zone_id not in _ZONE_REGISTRY: raise KeyError( f"Zone '{zone_id}' not registered. " f"Call register_zone() before fetching data." ) return _ZONE_REGISTRY[zone_id].centroid # --------------------------------------------------------------------------- # Config / cache # --------------------------------------------------------------------------- _CACHE_DIR = Path(os.environ.get("WEATHER_CACHE_DIR", ".cache/era5")) _CACHE_DIR.mkdir(parents=True, exist_ok=True) _ERA5_CACHE_DIR = _CACHE_DIR / "era5_nc" _ERA5_CACHE_DIR.mkdir(parents=True, exist_ok=True) _OPENMETEO_URL = "https://api.open-meteo.com/v1/forecast" _OPENMETEO_ARCHIVE_URL = "https://archive-api.open-meteo.com/v1/archive" _TIMEOUT_S = int(os.environ.get("WEATHER_HTTP_TIMEOUT", "30")) _CACHE_TTL_DAYS = int(os.environ.get("WEATHER_CACHE_TTL_DAYS", "7")) _ERA5_TTL_DAYS = int(os.environ.get("WEATHER_ERA5_TTL_DAYS", "30")) # ERA5 bounding box padding in degrees around zone centroid _ERA5_BOX_PAD = float(os.environ.get("WEATHER_ERA5_BOX_PAD", "0.5")) # --------------------------------------------------------------------------- # Cache utilities # --------------------------------------------------------------------------- def _cache_key(url: str, params: Dict[str, Any]) -> str: payload = f"{url}{json.dumps(params, sort_keys=True)}" return hashlib.sha256(payload.encode()).hexdigest() def _cached_get(url: str, params: Dict[str, Any]) -> Dict[str, Any]: """HTTP GET with file-based JSON cache. Uses context managers (no leaks).""" if not REQUESTS_AVAILABLE: raise RuntimeError("requests not installed — cannot fetch HTTP data") path = _CACHE_DIR / f"{_cache_key(url, params)}.json" if path.exists(): age = ( datetime.now(timezone.utc) - datetime.fromtimestamp(path.stat().st_mtime, tz=timezone.utc) ).days if age < _CACHE_TTL_DAYS: with open(path) as f: return json.load(f) resp = requests.get(url, params=params, timeout=_TIMEOUT_S) resp.raise_for_status() data = resp.json() with open(path, "w") as f: json.dump(data, f) return data def _era5_cache_path(zone_id: str, date_range: Tuple[datetime, datetime]) -> Path: """Deterministic cache file path for an ERA5 download.""" key = _stable_seed( zone_id + date_range[0].date().isoformat() + date_range[1].date().isoformat() ) return _ERA5_CACHE_DIR / f"{zone_id}_{key}.nc" # --------------------------------------------------------------------------- # Timezone helper # --------------------------------------------------------------------------- def _ensure_utc(dt: datetime) -> datetime: """Normalise a datetime to UTC at the pipeline boundary. Called on all externally-supplied datetimes before they reach ZoneObs, so ZoneObs.__post_init__ never sees a naive datetime in production. """ if dt.tzinfo is None: return dt.replace(tzinfo=timezone.utc) return dt.astimezone(timezone.utc) # --------------------------------------------------------------------------- # ERA5 variable mapping # --------------------------------------------------------------------------- # CDS short names → ZoneObs fields + unit conversion factors. # All ERA5 variables are in SI units; conversions applied in _build_era5_obs. # # ERA5 hourly variables we request: # total_precipitation m → mm (×1000) # 2m_temperature K → °C (−273.15) # 2m_dewpoint_temperature K → °C (−273.15) used for RH # 10m_u_component_of_wind m/s # 10m_v_component_of_wind m/s → speed = sqrt(u²+v²) # surface_pressure Pa (for ET0 calculation) # volumetric_soil_water_layer_1 m³/m³ → % (×100) # potential_evaporation m → mm (×1000, sign convention varies) _ERA5_VARIABLES: List[str] = [ "total_precipitation", "2m_temperature", "2m_dewpoint_temperature", "10m_u_component_of_wind", "10m_v_component_of_wind", "surface_pressure", "volumetric_soil_water_layer_1", "potential_evaporation", ] def _dewpoint_to_rh(temp_c: float, dewpoint_c: float) -> float: """Magnus formula: relative humidity from temperature and dewpoint (%). Accurate to ±0.4% RH for temperatures in the range 0–60°C. """ a, b = 17.625, 243.04 # Magnus coefficients rh = 100.0 * math.exp( (a * dewpoint_c / (b + dewpoint_c)) - (a * temp_c / (b + temp_c)) ) return max(0.0, min(100.0, rh)) def _build_era5_obs( zone_id: str, start: datetime, nc_path: Path, ) -> ZoneObs: """Map a downloaded ERA5 NetCDF file to a ZoneObs. Aggregates hourly ERA5 data to daily statistics over the window starting at `start`. All unit conversions are explicit and documented. Requires numpy and netCDF4 (or xarray). Falls through to synthetic if neither is available — caller handles the exception. """ try: import netCDF4 as nc # type: ignore except ImportError: try: import xarray as xr # type: ignore return _build_era5_obs_xarray(zone_id, start, nc_path) except ImportError: raise ImportError( "netCDF4 or xarray required for ERA5 ingestion. " "Install with: pip install netCDF4 or pip install xarray" ) import numpy as np ds = nc.Dataset(str(nc_path), "r") try: # --- Helper: read first spatial point (nearest to centroid) --- def _daily_mean(var: str) -> float: """Mean of all 24 hourly values for variable `var`.""" if var not in ds.variables: return 0.0 data = ds.variables[var][:24].flatten() data = np.ma.filled(data, np.nan) valid = data[np.isfinite(data)] return float(np.mean(valid)) if len(valid) > 0 else 0.0 def _daily_max(var: str) -> float: if var not in ds.variables: return 0.0 data = ds.variables[var][:24].flatten() data = np.ma.filled(data, np.nan) valid = data[np.isfinite(data)] return float(np.max(valid)) if len(valid) > 0 else 0.0 def _daily_sum(var: str, scale: float = 1.0) -> float: if var not in ds.variables: return 0.0 data = ds.variables[var][:24].flatten() data = np.ma.filled(data, np.nan) valid = data[np.isfinite(data)] return float(np.sum(valid) * scale) if len(valid) > 0 else 0.0 # --- Temperature (K → °C) --- temp_mean_c = _daily_mean("t2m") - 273.15 temp_max_c = _daily_max("t2m") - 273.15 # ERA5 doesn't have a daily min variable — approximate from hourly if "t2m" in ds.variables: t_data = np.ma.filled(ds.variables["t2m"][:24].flatten(), np.nan) valid = t_data[np.isfinite(t_data)] temp_min_c = float(np.min(valid)) - 273.15 if len(valid) > 0 else temp_mean_c else: temp_min_c = temp_mean_c # --- Dewpoint → RH --- dewpoint_mean_c = _daily_mean("d2m") - 273.15 rh_mean = _dewpoint_to_rh(temp_mean_c, dewpoint_mean_c) dewpoint_max_c = _daily_max("d2m") - 273.15 rh_max = _dewpoint_to_rh(temp_min_c, dewpoint_max_c) # RH highest at min temp # --- Precipitation: ERA5 total_precipitation in m, convert to mm --- # ERA5 accumulates per hour; sum 24h for daily total precip_24h = _daily_sum("tp", scale=1000.0) # m → mm # For multi-day aggregates we read from the nc file's full time span. # If the file only covers 1 day we fall back to scaling. total_hours = len(ds.variables.get("tp", [])) if "tp" in ds.variables else 24 if total_hours >= 24 * 7: def _window_sum(var: str, hours: int, scale: float = 1.0) -> float: if var not in ds.variables: return 0.0 data = np.ma.filled( ds.variables[var][:hours].flatten(), np.nan ) valid = data[np.isfinite(data)] return float(np.sum(valid) * scale) if len(valid) > 0 else 0.0 precip_7d = _window_sum("tp", 24*7, scale=1000.0) precip_14d = _window_sum("tp", 24*14, scale=1000.0) precip_30d = _window_sum("tp", 24*30, scale=1000.0) else: # Single-day file: extrapolate (approximate — pipeline should # request a 30-day window for full aggregates) precip_7d = precip_24h * 7.0 precip_14d = precip_24h * 14.0 precip_30d = precip_24h * 30.0 # --- Wind (m/s) --- u = _daily_mean("u10") v = _daily_mean("v10") wind_mean = math.sqrt(u**2 + v**2) u_max = _daily_max("u10") v_max = _daily_max("v10") wind_max = math.sqrt(u_max**2 + v_max**2) # --- Soil moisture (m³/m³ → %) --- soil_pct = _daily_mean("swvl1") * 100.0 # --- Potential evaporation (m → mm, ERA5 PE is negative convention) --- pe_m = _daily_sum("pev", scale=1.0) et0_mm = abs(pe_m) * 1000.0 # ERA5 PE is negative (energy leaving surface) return ZoneObs( zone_id=zone_id, valid_time=start, source=DataSource.ERA5_REANALYSIS, precip_24h_mm=max(0.0, precip_24h), precip_7d_mm=max(0.0, precip_7d), precip_14d_mm=max(0.0, precip_14d), precip_30d_mm=max(0.0, precip_30d), temp_mean_c=temp_mean_c, temp_max_c=max(temp_mean_c, temp_max_c), temp_min_c=min(temp_mean_c, temp_min_c), temp_anomaly_idx=0.0, # requires climatology — set in scorer precip_anomaly_idx=0.0, # requires climatology — set in scorer evapotranspiration_mm=max(0.0, et0_mm), wind_speed_mean_ms=max(0.0, wind_mean), wind_speed_max_ms=max(0.0, wind_max), rh_mean_pct=rh_mean, rh_max_pct=max(rh_mean, rh_max), soil_moisture_pct=max(0.0, min(100.0, soil_pct)), soil_moisture_anom=0.0, # requires climatology — set in scorer quality_flag=0, ) finally: ds.close() def _build_era5_obs_xarray( zone_id: str, start: datetime, nc_path: Path, ) -> ZoneObs: """xarray fallback for _build_era5_obs when netCDF4 is unavailable.""" import xarray as xr import numpy as np ds = xr.open_dataset(str(nc_path)) try: def _mean(var: str) -> float: if var not in ds: return 0.0 return float(ds[var].values.flatten()[np.isfinite( ds[var].values.flatten() )].mean()) if len(ds[var].values.flatten()) > 0 else 0.0 def _max(var: str) -> float: if var not in ds: return 0.0 vals = ds[var].values.flatten() valid = vals[np.isfinite(vals)] return float(valid.max()) if len(valid) > 0 else 0.0 def _sum(var: str, scale: float = 1.0) -> float: if var not in ds: return 0.0 vals = ds[var].values.flatten() valid = vals[np.isfinite(vals)] return float(valid.sum() * scale) if len(valid) > 0 else 0.0 temp_mean_c = _mean("t2m") - 273.15 temp_max_c = _max("t2m") - 273.15 if "t2m" in ds: _t2m_flat = ds["t2m"].values.flatten() _t2m_valid = _t2m_flat[np.isfinite(_t2m_flat)] temp_min_c = float(np.min(_t2m_valid)) - 273.15 if len(_t2m_valid) > 0 else temp_mean_c else: temp_min_c = temp_mean_c dewpoint_mean_c = _mean("d2m") - 273.15 dewpoint_max_c = _max("d2m") - 273.15 rh_mean = _dewpoint_to_rh(temp_mean_c, dewpoint_mean_c) rh_max = _dewpoint_to_rh(temp_min_c, dewpoint_max_c) precip_24h = max(0.0, _sum("tp", scale=1000.0)) # Use window sums if the file covers enough hours, otherwise extrapolate. # Matches the same logic used in the netCDF4 path (_build_era5_obs). if "tp" in ds: total_hours = len(ds["tp"].values.flatten()) else: total_hours = 0 if total_hours >= 24 * 7: def _window_sum_xr(var: str, hours: int, scale: float = 1.0) -> float: if var not in ds: return 0.0 vals = ds[var].values.flatten()[:hours] valid = vals[np.isfinite(vals)] return float(valid.sum() * scale) if len(valid) > 0 else 0.0 precip_7d = max(0.0, _window_sum_xr("tp", 24 * 7, scale=1000.0)) precip_14d = max(0.0, _window_sum_xr("tp", 24 * 14, scale=1000.0)) precip_30d = max(0.0, _window_sum_xr("tp", 24 * 30, scale=1000.0)) else: # Single-day file: extrapolate (approximate — pipeline should # request a 30-day window for full aggregates) precip_7d = precip_24h * 7.0 precip_14d = precip_24h * 14.0 precip_30d = precip_24h * 30.0 u = _mean("u10") v = _mean("v10") wind_mean = math.sqrt(u**2 + v**2) wind_max = math.sqrt(_max("u10")**2 + _max("v10")**2) soil_pct = _mean("swvl1") * 100.0 et0_mm = abs(_sum("pev")) * 1000.0 return ZoneObs( zone_id=zone_id, valid_time=start, source=DataSource.ERA5_REANALYSIS, precip_24h_mm=precip_24h, precip_7d_mm=precip_7d, precip_14d_mm=precip_14d, precip_30d_mm=precip_30d, temp_mean_c=temp_mean_c, temp_max_c=max(temp_mean_c, temp_max_c), temp_min_c=min(temp_mean_c, temp_min_c), evapotranspiration_mm=max(0.0, et0_mm), wind_speed_mean_ms=max(0.0, wind_mean), wind_speed_max_ms=max(0.0, wind_max), rh_mean_pct=rh_mean, rh_max_pct=max(rh_mean, rh_max), soil_moisture_pct=max(0.0, min(100.0, soil_pct)), quality_flag=0, ) finally: ds.close() # --------------------------------------------------------------------------- # Source fetchers # --------------------------------------------------------------------------- def _fetch_openmeteo(zone_id: str, date_range: Tuple[datetime, datetime]) -> ZoneObs: """Fetch from Open-Meteo API with full variable coverage. FIX from previous version: - Added temperature_2m_max, temperature_2m_min (were both set to mean) - Added relative_humidity_2m_mean, relative_humidity_2m_max (were both 0.0 — fungi risk signal was always zero) - Added et0_fao_evapotranspiration (was always 0.0) - Added precipitation_probability_max for forecast quality signal - wind_speed_10m_max added (previously only mean available) FIX (this pass): added soil_moisture_0_to_7cm_mean. Previously soil_moisture_pct was never populated on this path (always 0.0), which disabled the 0.4-weighted soil term of ZoneObs.drought_signal() for every Open-Meteo observation -- including all historical/archive replay. VERIFIED LIVE against both endpoints (archive + forecast API accept the variable as a daily mean, m3/m3); converted x100 to match the field's percent units (same convention as ERA5 swvl1 x 100 in _build_era5_obs). Uses archive API for historical dates, forecast API for future dates. """ lat, lon = _resolve_latlon(zone_id) start = _ensure_utc(date_range[0]) end = _ensure_utc(date_range[1]) # Open-Meteo archive API for historical; forecast API for recent/future today = datetime.now(timezone.utc).date() use_archive = start.date() < today - timedelta(days=5) url = _OPENMETEO_ARCHIVE_URL if use_archive else _OPENMETEO_URL daily_vars = ",".join([ "temperature_2m_mean", "temperature_2m_max", "temperature_2m_min", "precipitation_sum", "precipitation_hours", "wind_speed_10m_mean", "wind_speed_10m_max", "relative_humidity_2m_mean", "relative_humidity_2m_max", "et0_fao_evapotranspiration", "shortwave_radiation_sum", "soil_moisture_0_to_7cm_mean", ]) params = { "latitude": lat, "longitude": lon, "daily": daily_vars, "start_date": start.date().isoformat(), "end_date": end.date().isoformat(), "timezone": "UTC", } data = _cached_get(url, params) daily = data.get("daily", {}) def _safe(key: str, i: int = 0) -> float: vals = daily.get(key, []) return float(vals[i]) if i < len(vals) and vals[i] is not None else 0.0 def _safe_sum(key: str, n: int) -> float: vals = daily.get(key, []) return float(sum( v for v in vals[:n] if v is not None )) precip_all = daily.get("precipitation_sum", []) temp_mean = _safe("temperature_2m_mean") temp_max = _safe("temperature_2m_max") temp_min = _safe("temperature_2m_min") # Guard physical plausibility (API occasionally returns bad values) temp_max = max(temp_mean, temp_max) temp_min = min(temp_mean, temp_min) rh_mean = _safe("relative_humidity_2m_mean") rh_max = _safe("relative_humidity_2m_max") rh_max = max(rh_mean, rh_max) wind_mean = _safe("wind_speed_10m_mean") wind_max = _safe("wind_speed_10m_max") wind_max = max(wind_mean, wind_max) return ZoneObs( zone_id=zone_id, valid_time=start, source=DataSource.OPENMETEO_LIVE, precip_24h_mm=_safe("precipitation_sum"), precip_7d_mm=_safe_sum("precipitation_sum", 7), precip_14d_mm=_safe_sum("precipitation_sum", 14), precip_30d_mm=_safe_sum("precipitation_sum", 30), temp_mean_c=temp_mean, temp_max_c=temp_max, temp_min_c=temp_min, evapotranspiration_mm=_safe("et0_fao_evapotranspiration"), wind_speed_mean_ms=wind_mean / 3.6, # km/h → m/s wind_speed_max_ms=wind_max / 3.6, rh_mean_pct=rh_mean, rh_max_pct=rh_max, # m3/m3 -> % (verified live; same convention as ERA5 swvl1 x 100) soil_moisture_pct=max(0.0, min(100.0, _safe("soil_moisture_0_to_7cm_mean") * 100.0)), quality_flag=0, ) def _fetch_era5(zone_id: str, date_range: Tuple[datetime, datetime]) -> ZoneObs: """Fetch ERA5 reanalysis data via CDS API. Downloads hourly ERA5 data for a bounding box around the zone centroid, caches the NetCDF file locally, then aggregates to a single ZoneObs. CDS API credentials must be configured at ~/.cdsapirc: url: https://cds.climate.copernicus.eu/api/v2 key: : Falls back to synthetic data if: - cdsapi not installed - CDS request fails (quota, network, invalid dates) - NetCDF parsing fails (netCDF4 and xarray both unavailable) Bug 2.1 fix: the CDS request now correctly handles date ranges that span multiple months or years. Previously only start.month and start.year were passed, causing empty day ranges and missing data for any window crossing a month boundary (e.g. a 30-day precip window). The request now enumerates all (year, month, days) tuples that fall within [start, end]. """ if not CDSAPI_AVAILABLE: logger.debug("cdsapi unavailable — falling back to synthetic for %s", zone_id) return _fetch_synthetic(zone_id, date_range) lat, lon = _resolve_latlon(zone_id) start = _ensure_utc(date_range[0]) end = _ensure_utc(date_range[1]) # CDS bounding box: [north, west, south, east] bbox = [ round(lat + _ERA5_BOX_PAD, 2), round(lon - _ERA5_BOX_PAD, 2), round(lat - _ERA5_BOX_PAD, 2), round(lon + _ERA5_BOX_PAD, 2), ] nc_path = _era5_cache_path(zone_id, date_range) # Use cached file if within TTL if nc_path.exists(): age_days = ( datetime.now(timezone.utc) - datetime.fromtimestamp(nc_path.stat().st_mtime, tz=timezone.utc) ).days if age_days < _ERA5_TTL_DAYS: logger.debug("ERA5 cache hit for %s", zone_id) else: nc_path.unlink() # expired — delete and re-fetch if not nc_path.exists(): # --- Bug 2.1 fix: enumerate all (year, month, day) tuples in range --- # Build a set of unique years, months, and days that appear in [start, end]. # The CDS API accepts arrays for year/month/day and returns the union of all # matching hours; we pass all years, months, and days that appear in the # window so that month-boundary-crossing ranges are fully covered. years: set = set() months: set = set() days: set = set() cursor = start.date() end_date = end.date() while cursor <= end_date: years.add(cursor.year) months.add(cursor.month) days.add(cursor.day) cursor += timedelta(days=1) logger.info( "Requesting ERA5 data for %s bbox=%s dates=%s to %s", zone_id, bbox, start.date().isoformat(), end.date().isoformat(), ) try: client = cdsapi.Client(quiet=True) client.retrieve( "reanalysis-era5-single-levels", { "product_type": "reanalysis", "variable": _ERA5_VARIABLES, "year": sorted(str(y) for y in years), "month": sorted(f"{m:02d}" for m in months), "day": sorted(f"{d:02d}" for d in days), "time": [f"{h:02d}:00" for h in range(24)], "area": bbox, "format": "netcdf", }, str(nc_path), ) logger.info("ERA5 download complete: %s", nc_path) except Exception as e: logger.warning( "ERA5 CDS request failed for %s: %s — falling back to synthetic", zone_id, e ) if nc_path.exists(): nc_path.unlink() return _fetch_synthetic(zone_id, date_range) try: return _build_era5_obs(zone_id, start, nc_path) except Exception as e: logger.warning( "ERA5 NetCDF parse failed for %s: %s — falling back to synthetic", zone_id, e ) return _fetch_synthetic(zone_id, date_range) def _fetch_imerg(zone_id: str, date_range: Tuple[datetime, datetime]) -> ZoneObs: """Fetch satellite-retrieved precipitation from GPM IMERG via Earth Engine. IMERG is a genuine multi-satellite retrieval (passive microwave + IR merged, gauge-calibrated in the Final product), not short-forecast model output the way ERA5's total_precipitation is -- it does not carry ERA5's known double-ITCZ bias / underestimated convective peaks over island terrain. Other fields (temp, wind, RH, soil) are NOT covered by IMERG; this fetcher only replaces the precipitation fields and leaves everything else at ZoneObs defaults, matching the "satellite source overrides one variable family" design in ForecastConfig.use_satellite_precip. Window convention matches _fetch_era5 exactly: date_range[0] is the window start, and 24h/7d/14d/30d are sums of the *first* N days from that start (not trailing windows ending at start). VERIFY BEFORE PRODUCTION USE: the collection ID and band name below ("NASA/GPM_L3/IMERG_V07", band "precipitation", mm/hr) reflect the GEE data catalog as documented at the time this was written. This has not been exercised against a live, authenticated Earth Engine project in this environment (no network egress to Earth Engine here, and EE requires a registered Google Cloud project + `earthengine authenticate`) -- confirm the band name at https://developers.google.com/earth-engine/datasets/catalog/NASA_GPM_L3_IMERG_V07 before trusting this in a real training run. Falls back to whatever fetch_zone_obs()'s caller chain does next (ERA5 -> Open-Meteo -> synthetic) on any failure — this function itself just raises; it does not catch. """ _ensure_ee_initialized() lat, lon = _resolve_latlon(zone_id) start = _ensure_utc(date_range[0]) region = ee.Geometry.Point([lon, lat]).buffer(_ERA5_BOX_PAD * 111_000) # deg -> m, rough def _window_sum_mm(n_days: int) -> float: window_end = start + timedelta(days=n_days) coll = ( ee.ImageCollection("NASA/GPM_L3/IMERG_V07") .filterDate(start.isoformat(), window_end.isoformat()) .filterBounds(region) .select("precipitation") # mm/hr, calibrated -- VERIFY band name, see docstring ) # Each image is a half-hourly rate (mm/hr); sum(rate) * 0.5h/image = mm total total_mm_image = coll.sum().multiply(0.5) stats = total_mm_image.reduceRegion( reducer=ee.Reducer.mean(), geometry=region, scale=11_000, bestEffort=True, ).getInfo() return float(stats.get("precipitation", 0.0) or 0.0) precip_24h = max(0.0, _window_sum_mm(1)) precip_7d = max(precip_24h, _window_sum_mm(7)) precip_14d = max(precip_7d, _window_sum_mm(14)) precip_30d = max(precip_14d, _window_sum_mm(30)) return ZoneObs( zone_id=zone_id, valid_time=start, source=DataSource.SATELLITE_PRECIP, precip_24h_mm=precip_24h, precip_7d_mm=precip_7d, precip_14d_mm=precip_14d, precip_30d_mm=precip_30d, precip_satellite_mm=precip_24h, quality_flag=0, ) def _fetch_smap(zone_id: str, date_range: Tuple[datetime, datetime]) -> ZoneObs: """Fetch satellite-retrieved surface soil moisture from SMAP via Earth Engine. Uses SMAP L4 (3-hourly, gap-filled surface + root-zone product) rather than L3 (native 9km retrieval, ~2-3 day revisit gaps) for continuous coverage -- note L4 blends the raw retrieval with a land model to fill those gaps, so it is less "purely observational" than L3. Populates both the raw soil_moisture_satellite_pct provenance field and the canonical soil_moisture_pct field consumed by crop_risk_scorer.py. VERIFY BEFORE PRODUCTION USE: same caveat as _fetch_imerg -- collection ID "NASA/SMAP/SPL4SMGP/007" and band "sm_surface" (m3/m3) reflect the documented GEE catalog at write time and have not been exercised against a live authenticated Earth Engine backend in this environment. """ _ensure_ee_initialized() lat, lon = _resolve_latlon(zone_id) start = _ensure_utc(date_range[0]) end = start + timedelta(days=1) region = ee.Geometry.Point([lon, lat]).buffer(_ERA5_BOX_PAD * 111_000) coll = ( ee.ImageCollection("NASA/SMAP/SPL4SMGP/007") .filterDate(start.isoformat(), end.isoformat()) .filterBounds(region) .select("sm_surface") # m3/m3 volumetric water content -- VERIFY band name, see docstring ) stats = coll.mean().reduceRegion( reducer=ee.Reducer.mean(), geometry=region, scale=9_000, bestEffort=True, ).getInfo() vwc = stats.get("sm_surface") if vwc is None: raise RuntimeError(f"SMAP: no data returned for zone={zone_id} date={start.date()}") soil_pct = max(0.0, min(100.0, float(vwc) * 100.0)) return ZoneObs( zone_id=zone_id, valid_time=start, source=DataSource.SATELLITE_SOIL, soil_moisture_pct=soil_pct, soil_moisture_satellite_pct=soil_pct, quality_flag=0, ) def _fetch_synthetic(zone_id: str, date_range: Tuple[datetime, datetime]) -> ZoneObs: """Deterministic synthetic ZoneObs. Last-resort fallback in production, primary source during early training before real data is available. """ seed = _stable_seed(zone_id + date_range[0].isoformat()) # Keyword args required — positional args silently pass seed into crop_stage return make_synthetic_zone_obs(zone_id=zone_id, seed=seed) # --------------------------------------------------------------------------- # Noise injection # --------------------------------------------------------------------------- def _inject_noise(obs: ZoneObs, rng: random.Random, scale: float) -> ZoneObs: """Return a new ZoneObs with small perturbations on key meteorological fields. Uses to_dict()/from_dict() — safe under slots=True and future field additions. Never mutates the input obs. Perturbs precip, temperature, RH, and wind independently to simulate inter-station measurement variability. Scale=0.05 ≈ ±5% noise. """ def perturb(x: float) -> float: return x * (1.0 + rng.uniform(-scale, scale)) base = obs.to_dict() base.pop("_schema_version", None) base["precip_24h_mm"] = max(0.0, perturb(obs.precip_24h_mm)) base["precip_7d_mm"] = max(0.0, perturb(obs.precip_7d_mm)) base["precip_14d_mm"] = max(0.0, perturb(obs.precip_14d_mm)) base["precip_30d_mm"] = max(0.0, perturb(obs.precip_30d_mm)) base["temp_mean_c"] = perturb(obs.temp_mean_c) base["temp_max_c"] = perturb(obs.temp_max_c) base["temp_min_c"] = perturb(obs.temp_min_c) base["rh_mean_pct"] = min(100.0, max(0.0, perturb(obs.rh_mean_pct))) base["rh_max_pct"] = min(100.0, max(0.0, perturb(obs.rh_max_pct))) base["wind_speed_mean_ms"] = max(0.0, perturb(obs.wind_speed_mean_ms)) base["wind_speed_max_ms"] = max(0.0, perturb(obs.wind_speed_max_ms)) # Preserve monotonic constraint: perturbed aggregates must stay ordered p24 = base["precip_24h_mm"] p7 = max(base["precip_7d_mm"], p24) p14 = max(base["precip_14d_mm"], p7) p30 = max(base["precip_30d_mm"], p14) base["precip_7d_mm"] = p7 base["precip_14d_mm"] = p14 base["precip_30d_mm"] = p30 # Preserve temp ordering t_mean = base["temp_mean_c"] base["temp_max_c"] = max(t_mean, base["temp_max_c"]) base["temp_min_c"] = min(t_mean, base["temp_min_c"]) # Preserve RH ordering base["rh_max_pct"] = max(base["rh_mean_pct"], base["rh_max_pct"]) # Preserve wind ordering base["wind_speed_max_ms"] = max( base["wind_speed_mean_ms"], base["wind_speed_max_ms"] ) return ZoneObs.from_dict(base) # --------------------------------------------------------------------------- # Basin-scale context (ENSO / IOD / helio) # --------------------------------------------------------------------------- # Verified live during development (see chat record): this exact URL and # ASCII format ("SEAS YR ANOM", one row per 3-month season) were fetched and # confirmed. NOAA now uses RONI (Relative ONI) rather than the legacy ONI # for official ENSO monitoring, per NWS Public Information Statement 26-05. _NOAA_RONI_URL = "https://www.cpc.ncep.noaa.gov/data/indices/RONI.ascii.txt" # NOT independently verified live in this environment (no network egress to # psl.noaa.gov here). This follows PSL's long-standing "data/correlation/" # raw-index convention (first line: start/end year; each following line: # "year jan feb ... dec"; missing values sentinel ~ -99.9). Confirm this # resolves and parses correctly against # https://psl.noaa.gov/data/timeseries/month/DS/DMI/ before production use — # if the exact filename differs, update _PSL_DMI_URL and, if the format # differs, _parse_psl_monthly_ascii below. # Primary + fallback DMI sources. The old correlation/dmi.data path has been # intermittently 502; the HadISST long series is the stable PSL product. _PSL_DMI_URLS: Tuple[str, ...] = ( "https://psl.noaa.gov/gcos_wgsp/Timeseries/Data/dmi.had.long.data", "https://psl.noaa.gov/data/timeseries/month/data/dmi.had.long.data", "https://psl.noaa.gov/data/correlation/dmi.data", ) # NOAA SWPC real-time JSON endpoints (public, no API key). # Plasma: the legacy json/solar-wind/plasma-*.json paths 404'd after the # 2026 RTSW migration (SCN 26-21). Prefer the new rtsw_wind product; # field name is proton_speed (was speed). Keep a products/summary fallback. _SWPC_PLASMA_URLS: Tuple[str, ...] = ( "https://services.swpc.noaa.gov/json/rtsw/rtsw_wind_1m.json", "https://services.swpc.noaa.gov/products/summary/solar-wind-speed.json", "https://services.swpc.noaa.gov/products/solar-wind/plasma-1-day.json", ) _SWPC_KP_URL = "https://services.swpc.noaa.gov/json/planetary_k_index_1m.json" _SWPC_XRAY_URL = "https://services.swpc.noaa.gov/json/goes/primary/xrays-1-day.json" _SEASON_FOR_MONTH: Dict[int, str] = { 1: "DJF", 2: "JFM", 3: "FMA", 4: "MAM", 5: "AMJ", 6: "MJJ", 7: "JJA", 8: "JAS", 9: "ASO", 10: "SON", 11: "OND", 12: "NDJ", } # Chronological season order for walk-back when the current (incomplete) # season has not been published yet on CPC's RONI table. _SEASON_ORDER: Tuple[str, ...] = ( "DJF", "JFM", "FMA", "MAM", "AMJ", "MJJ", "JJA", "JAS", "ASO", "SON", "OND", "NDJ", ) def _parse_cpc_seasonal_ascii(text: str) -> Dict[Tuple[int, str], float]: """Parse NOAA CPC's seasonal index ASCII format: 'SEAS YR [TOTAL] ANOM'. Works for both RONI (3 columns: SEAS YR ANOM) and legacy ONI (4 columns: SEAS YR TOTAL ANOM) since it always reads the last column as the value. Returns {(year, season_code): anomaly}. """ out: Dict[Tuple[int, str], float] = {} lines = text.strip().splitlines() for line in lines[1:]: # skip header row parts = line.split() if len(parts) < 3: continue season = parts[0] try: year = int(parts[1]) value = float(parts[-1]) except (ValueError, IndexError): continue out[(year, season)] = value return out def _parse_psl_monthly_ascii( text: str, year: int, month: int, missing_below: float = -90.0 ) -> Optional[float]: """Parse PSL's standard 'year v1 v2 ... v12' monthly index format. Returns None if the year/month isn't found or the value is a missing sentinel (PSL commonly uses -99.9 / -999.9 / -9999 style sentinels, all comfortably below missing_below). """ for line in text.strip().splitlines(): parts = line.split() if len(parts) != 13: continue try: row_year = int(parts[0]) values = [float(v) for v in parts[1:]] except ValueError: continue if row_year == year: v = values[month - 1] return None if v <= missing_below else v return None def _lookup_roni_with_lag( table: Dict[Tuple[int, str], float], year: int, season: str ) -> Tuple[float, str]: """Return (value, label) for the requested season, or the latest prior. CPC RONI is a 3-month running index published with a lag: the current incomplete season (e.g. JJA while still in July) is often absent. Walking back keeps the value *real* instead of falling through to synthetic. """ if (year, season) in table: return table[(year, season)], f"{season} {year}" try: idx = _SEASON_ORDER.index(season) except ValueError: idx = 0 y, i = year, idx for _ in range(24): # at most 2 years of lag i -= 1 if i < 0: i = len(_SEASON_ORDER) - 1 y -= 1 key = (y, _SEASON_ORDER[i]) if key in table: return table[key], f"{_SEASON_ORDER[i]} {y} (lagged from {season} {year})" raise ValueError(f"No RONI value for {season} {year} or any prior season in table") def _lookup_dmi_with_lag( text: str, year: int, month: int ) -> Tuple[float, str]: """Return (value, label) for year/month, or the latest prior published month.""" y, m = year, month for _ in range(24): v = _parse_psl_monthly_ascii(text, y, m) if v is not None: label = f"{y}-{m:02d}" if (y, m) != (year, month): label += f" (lagged from {year}-{month:02d})" return v, label m -= 1 if m < 1: m = 12 y -= 1 raise ValueError(f"No DMI value for {year}-{month:02d} or any prior month") def _extract_solar_wind_speeds(payload: Any) -> List[float]: """Pull bulk/proton speed samples from heterogeneous SWPC JSON shapes.""" speeds: List[float] = [] def _from_row(row: Any) -> None: if not isinstance(row, dict): return for key in ("proton_speed", "speed", "wind_speed", "value"): sp = row.get(key) if sp is None: continue try: f = float(sp) except (TypeError, ValueError): continue if f > 0: speeds.append(f) return if isinstance(payload, list): # Common case: list of dicts. Some products put a header row first. for row in payload: if isinstance(row, dict): _from_row(row) elif isinstance(row, (list, tuple)) and len(row) >= 2: # legacy [time_tag, speed, ...] rows try: f = float(row[1]) if f > 0: speeds.append(f) except (TypeError, ValueError): continue elif isinstance(payload, dict): _from_row(payload) for v in payload.values(): if isinstance(v, list): speeds.extend(_extract_solar_wind_speeds(v)) return speeds def _fetch_swpc_helio(valid_date: datetime) -> Dict[str, Any]: """Pull latest NOAA SWPC solar-wind / Kp / GOES X-ray snapshots. Returns a partial dict of whatever endpoints succeed. Callers must supply quiet-Sun defaults for any missing keys — never invent storm values on failure (anti-saturation: missing data must not overweight risk the way the old absolute humidity term did). Endpoints are public JSON; no API key required. Each is fetched independently so one outage does not zero the whole block. """ out: Dict[str, Any] = {} if not REQUESTS_AVAILABLE: return out # --- Solar wind speed (RTSW proton_speed / legacy speed, km/s) --- plasma_err: Optional[Exception] = None for url in _SWPC_PLASMA_URLS: try: resp = requests.get(url, timeout=_TIMEOUT_S) resp.raise_for_status() speeds = _extract_solar_wind_speeds(resp.json()) if speeds: out["solar_wind_speed_kms"] = speeds[-1] plasma_err = None break plasma_err = ValueError(f"no speed samples in {url}") except Exception as e: plasma_err = e continue if "solar_wind_speed_kms" not in out and plasma_err is not None: logger.warning( "_fetch_swpc_helio: plasma fetch failed (%s) for %s", plasma_err, valid_date.date(), ) # --- Planetary K-index (most recent 1-min / 3-hour estimate) --- try: resp = requests.get(_SWPC_KP_URL, timeout=_TIMEOUT_S) resp.raise_for_status() rows = resp.json() kps = [] for row in rows: try: kp = row.get("kp_index", row.get("kp")) if kp is not None: kps.append(float(kp)) except (TypeError, ValueError, AttributeError): continue if kps: out["kp_index"] = kps[-1] except Exception as e: logger.warning( "_fetch_swpc_helio: Kp fetch failed (%s) for %s", e, valid_date.date(), ) # --- GOES primary X-ray long-channel flux (W/m²) --- try: resp = requests.get(_SWPC_XRAY_URL, timeout=_TIMEOUT_S) resp.raise_for_status() rows = resp.json() fluxes = [] for row in rows: try: flux = row.get("flux") energy = str(row.get("energy", "")).lower() if flux is not None and float(flux) > 0: if "0.1-0.8" in energy or "long" in energy or not energy: fluxes.append(float(flux)) except (TypeError, ValueError, AttributeError): continue if fluxes: out["goes_xray_flux"] = fluxes[-1] except Exception as e: logger.warning( "_fetch_swpc_helio: GOES X-ray fetch failed (%s) for %s", e, valid_date.date(), ) if out: kp = float(out.get("kp_index", 2.0)) xray = float(out.get("goes_xray_flux", 1e-7)) out["helio_regime"] = derive_helio_regime(kp, xray) return out def fetch_basin_context( valid_date: datetime, config: Optional[ForecastConfig] = None, ) -> BasinContext: """Fetch published basin-scale climate + helio indices for one date. Deliberately does NOT re-derive ENSO/IOD from ERA5 SST — ERA5's SST field is itself an interpolated boundary condition (from HadISST2/OSTIA), not an independently assimilated variable, so re-deriving ENSO/IOD from it would add a layer of local interpretation on top of an already-derived product. This pulls NOAA CPC's and NOAA PSL's own published index values directly instead — the same series the field actually uses. Helio fields (solar wind, Kp, GOES X-ray) come from NOAA SWPC real-time JSON endpoints. Soft mode (default, ``require_real_basin_context=False``) -------------------------------------------------------- - ENSO / IOD: per-index synthetic fallback on network/parse failure. - Helio: missing endpoints fall back to quiet-Sun dataclass defaults (not random synthetic draws) so offline / partial outages never inject artificial "storm" context. - ``requests`` missing → full synthetic BasinContext. Strict mode (``require_real_basin_context=True``) ------------------------------------------------- Refuses every synthetic / quiet-default fallback. Raises ``RuntimeError`` if any of the following cannot be obtained from live published sources: * RONI (ENSO) * DMI (IOD) * SWPC solar wind speed, Kp, and GOES X-ray flux (all three required) Use this for production scoring / backtests where synthetic basin values would invalidate the result. Training should leave the flag False. itcz_latitude_deg and mslp_regional_hpa remain neutral defaults in both modes (no published single-index feed yet). """ cfg = config or ForecastConfig() valid_date = _ensure_utc(valid_date) season = _SEASON_FOR_MONTH[valid_date.month] strict = bool(getattr(cfg, "require_real_basin_context", False)) if not REQUESTS_AVAILABLE: if strict: raise RuntimeError( "require_real_basin_context=True but the 'requests' package " "is not installed — cannot fetch live basin/helio indices." ) logger.info("requests not installed — basin context falls back to synthetic") return make_synthetic_basin_context(valid_date=valid_date) # --- ENSO (RONI) — exact season, else latest published prior season --- try: resp = requests.get(_NOAA_RONI_URL, timeout=_TIMEOUT_S) resp.raise_for_status() table = _parse_cpc_seasonal_ascii(resp.text) enso_oni, roni_label = _lookup_roni_with_lag(table, valid_date.year, season) if "lagged" in roni_label: logger.info( "fetch_basin_context: RONI using %s for %s", roni_label, valid_date.date(), ) except Exception as e: if strict: raise RuntimeError( f"require_real_basin_context=True: RONI fetch failed for " f"{valid_date.date()}: {e}" ) from e logger.warning( "fetch_basin_context: RONI fetch/parse failed (%s) — synthetic ENSO for %s", e, valid_date.date(), ) enso_oni = make_synthetic_basin_context( valid_date=valid_date, seed=_stable_seed(f"enso_{valid_date.date().isoformat()}") ).enso_oni # --- IOD (DMI) — try multiple PSL URLs; lag to latest published month --- iod_dmi: Optional[float] = None dmi_err: Optional[Exception] = None for dmi_url in _PSL_DMI_URLS: try: resp = requests.get(dmi_url, timeout=_TIMEOUT_S) resp.raise_for_status() iod_dmi, dmi_label = _lookup_dmi_with_lag( resp.text, valid_date.year, valid_date.month ) if "lagged" in dmi_label: logger.info( "fetch_basin_context: DMI using %s for %s (from %s)", dmi_label, valid_date.date(), dmi_url, ) dmi_err = None break except Exception as e: dmi_err = e continue if iod_dmi is None: if strict: raise RuntimeError( f"require_real_basin_context=True: DMI fetch failed for " f"{valid_date.date()}: {dmi_err}" ) logger.warning( "fetch_basin_context: DMI fetch/parse failed (%s) — synthetic IOD for %s", dmi_err, valid_date.date(), ) iod_dmi = make_synthetic_basin_context( valid_date=valid_date, seed=_stable_seed(f"iod_{valid_date.date().isoformat()}") ).iod_dmi # --- Helio (SWPC) --- helio = _fetch_swpc_helio(valid_date) required_helio = ("solar_wind_speed_kms", "kp_index", "goes_xray_flux") missing_helio = [k for k in required_helio if k not in helio] if strict and missing_helio: raise RuntimeError( f"require_real_basin_context=True: SWPC helio incomplete for " f"{valid_date.date()} — missing {missing_helio}. " f"Got keys: {sorted(helio.keys())}" ) return BasinContext( valid_date=valid_date, enso_oni=enso_oni, iod_dmi=iod_dmi, solar_wind_speed_kms=float(helio.get("solar_wind_speed_kms", 400.0)), kp_index=float(helio.get("kp_index", 2.0)), goes_xray_flux=float(helio.get("goes_xray_flux", 1e-7)), helio_regime=str(helio.get("helio_regime", "quiet")), source=DataSource.PUBLISHED_INDEX, ) # --------------------------------------------------------------------------- # Source selection # --------------------------------------------------------------------------- def _select_source(cfg: ForecastConfig, rng: random.Random) -> DataSource: real_ratio = getattr(cfg, "real_data_ratio", 0.7) era5_ratio = getattr(cfg, "era5_ratio", 0.5) if rng.random() > real_ratio: return DataSource.SYNTHETIC # Opt-in satellite sources (schema v3+). Default False on ForecastConfig, # so existing callers get exactly the pre-v3 ERA5/Open-Meteo split unless # they explicitly enable one of these. If both are enabled, choose # between them rather than always preferring one silently. satellite_options = [] if getattr(cfg, "use_satellite_precip", False): satellite_options.append(DataSource.SATELLITE_PRECIP) if getattr(cfg, "use_satellite_soil", False): satellite_options.append(DataSource.SATELLITE_SOIL) if satellite_options: return satellite_options[rng.randrange(len(satellite_options))] return ( DataSource.ERA5_REANALYSIS if rng.random() < era5_ratio else DataSource.OPENMETEO_LIVE ) def _fallback_chain(primary: DataSource) -> List[DataSource]: """Ordered fallback chain for a given primary source selection. Satellite sources fall back through ERA5 before synthetic, matching the existing "degrade gracefully" philosophy (ERA5 -> Open-Meteo -> synthetic) rather than jumping straight to synthetic on the first failure. All chains terminate at SYNTHETIC. """ if primary in (DataSource.SATELLITE_PRECIP, DataSource.SATELLITE_SOIL): return [primary, DataSource.ERA5_REANALYSIS, DataSource.SYNTHETIC] return [primary, DataSource.SYNTHETIC] # --------------------------------------------------------------------------- # Public API # --------------------------------------------------------------------------- def fetch_zone_obs( zone_id: str, date_range: Tuple[datetime, datetime], config: Optional[ForecastConfig] = None, ) -> ZoneObs: """Fetch a validated ZoneObs for one zone and date window. Source priority: 1. cfg.force_data_source (debug/test override) 2. Stochastic selection based on real_data_ratio / era5_ratio 3. Synthetic fallback if all real sources fail Always validates the result with ZoneObs.validate(strict=True). Re-validates after noise injection. Bug 2.4 fix: noise injection is now performed outside the fetch/validate try-except block. Previously a noise injection failure (or a re-validation failure after noise) was caught by the same except clause as a fetch failure, silently triggering a redundant retry on DataSource.SYNTHETIC and masking the real error. Noise injection errors now propagate directly to the caller so they are visible and debuggable. """ cfg = config or ForecastConfig() seed = _stable_seed(zone_id + date_range[0].isoformat()) rng = random.Random(seed) if getattr(cfg, "force_data_source", None) is not None: primary = cfg.force_data_source else: primary = _select_source(cfg, rng) fetchers = { DataSource.OPENMETEO_LIVE: _fetch_openmeteo, DataSource.ERA5_REANALYSIS: _fetch_era5, DataSource.SATELLITE_PRECIP: _fetch_imerg, DataSource.SATELLITE_SOIL: _fetch_smap, DataSource.SYNTHETIC: _fetch_synthetic, } obs: Optional[ZoneObs] = None for source in _fallback_chain(primary): try: obs = fetchers[source](zone_id, date_range) ZoneObs.validate(obs, strict=True) break # fetch + validation succeeded; exit retry loop except Exception as e: logger.warning("%s failed for %s: %s", source.value, zone_id, e) obs = None if obs is None: raise RuntimeError( f"All data sources failed for zone '{zone_id}'. " f"Check zone registration and network connectivity." ) # --- Bug 2.4 fix: noise injection is outside the fetch/validate loop --- # A failure here is a pipeline bug (bad noise config or invariant violation # introduced by _inject_noise itself) and should propagate to the caller, # not silently trigger another synthetic fetch. if getattr(cfg, "inject_noise", False): obs = _inject_noise(obs, rng, cfg.noise_scale) # Re-validate: noise can violate invariants even with monotonic guards ZoneObs.validate(obs, strict=True) # --- Climatology anomalies (opt-in via cfg.use_climatology_anomalies) --- # Real fetchers leave precip_anomaly_idx / temp_anomaly_idx / # soil_moisture_anom at 0.0 ("requires climatology"). Without this step # those stay 0.0 forever, which zeroes most of drought_signal() and the # primary term of flood_signal() -- i.e. real observations score as # near-zero risk regardless of actual conditions. climatology.py is the # missing layer. Applied AFTER noise injection so anomalies reflect the # final values the env will see. Synthetic-source obs are skipped inside # apply_climatology_anomalies (their anomalies are event-injected). if getattr(cfg, "use_climatology_anomalies", False): try: _lat, _lon = _resolve_latlon(zone_id) except KeyError: logger.warning( "use_climatology_anomalies=True but zone '%s' is not " "registered -- cannot locate the zone for its climatology. " "Anomaly fields left at fetcher defaults. Call " "register_zone() first.", zone_id, ) else: # Local import: keeps climatology.py an optional dependency and # avoids any module-load-order coupling. from climatology import ( apply_climatology_anomalies, get_zone_climatology, ) _clim = get_zone_climatology( zone_id, _lat, _lon, years=getattr(cfg, "climatology_years", 10), ) obs = apply_climatology_anomalies(obs, _clim) # Anomalies are derived fields of validated inputs, but # re-validate cheaply for defence in depth. ZoneObs.validate(obs, strict=True) return obs def _build_context_forecast( zone_id: str, obs: ZoneObs, cfg: ForecastConfig, seed: int, ) -> "ForecastResult": """Build the ForecastResult for an EpisodeContext. Backend is selected by cfg.forecast_backend (default 'synthetic'): 'synthetic' -- deterministic synthetic forecast. NOTE (behaviour change vs the pre-integration version, deliberate): the synthetic forecast is now anchored to obs.valid_time and cfg.horizon_days instead of a decoupled synthetic EpisodeContext's own clock and a hardcoded 30-day horizon. Still fully deterministic (seeded), but now temporally coherent with the obs it describes. 'baseline' -- timesfm_wrapper.BaselineBackend (deterministic statistical forecast from obs fields; offline). 'openmeteo' -- timesfm_wrapper.OpenMeteoBackend (real API forecast; zone must be registered for lat/lon; falls back to synthetic on ANY failure, logged). 'timesfm' -- LocalTimesFMBackend. Requires the caller to have the checkpoint + sha + timesfm package; this pipeline layer has no way to supply those, so this mode always falls back to synthetic here with a loud warning. Use timesfm_wrapper.create_forecast_backend directly if you need the TimesFM tier. """ from zone_observation import ForecastResult # noqa: F401 (type hint only) mode = getattr(cfg, "forecast_backend", "synthetic") def _synthetic() -> "ForecastResult": from zone_observation import make_synthetic_forecast_result return make_synthetic_forecast_result( zone_id=zone_id, valid_time=obs.valid_time, horizon_days=cfg.horizon_days, seed=seed, ) if mode == "synthetic": return _synthetic() if mode == "timesfm": logger.warning( "forecast_backend='timesfm' cannot be constructed inside " "fetch_episode_context (no checkpoint/sha available at this " "layer) -- falling back to synthetic for zone=%s. Build the " "backend via timesfm_wrapper.create_forecast_backend() and " "call backend.forecast(obs) directly if you need TimesFM.", zone_id, ) return _synthetic() try: from timesfm_wrapper import create_forecast_backend lat = lon = None if mode == "openmeteo": lat, lon = _resolve_latlon(zone_id) backend = create_forecast_backend( mode=mode, lat=lat, lon=lon, horizon_days=cfg.horizon_days, ) return backend.forecast(obs) except Exception as e: logger.warning( "forecast backend %r failed for zone=%s (%s) -- synthetic " "forecast fallback", mode, zone_id, e, ) return _synthetic() def fetch_episode_context( zone_id: str, date_range: Tuple[datetime, datetime], config: Optional[ForecastConfig] = None, ) -> EpisodeContext: """Fetch a complete EpisodeContext for one zone and date window. obs comes from the real data pipeline (ERA5 / Open-Meteo / satellite / synthetic, per cfg source-selection fields). forecast comes from the backend named by cfg.forecast_backend (default 'synthetic'; see _build_context_forecast for the full mode matrix) -- this is the timesfm_wrapper integration the previous version's docstring flagged as not yet done. basin_context is attached only when cfg.include_basin_context is True (default False — existing callers see no change). Soft mode (``require_real_basin_context=False``, default): failures fetching basin context never fail the whole EpisodeContext — ``basin_context`` is left as None (a valid, handled state throughout this pipeline; see weather_forecast_env.py's neutral default). Strict mode (``require_real_basin_context=True``): any basin/helio fetch failure is re-raised so callers cannot silently score with synthetic or missing teleconnection context. """ cfg = config or ForecastConfig() obs = fetch_zone_obs(zone_id, date_range, cfg) seed = _stable_seed(zone_id + date_range[0].isoformat() + "_forecast") fcast = _build_context_forecast(zone_id, obs, cfg, seed) basin_context = None if getattr(cfg, "include_basin_context", False): strict_basin = bool(getattr(cfg, "require_real_basin_context", False)) try: basin_context = fetch_basin_context(date_range[0], cfg) except Exception as e: if strict_basin: raise # real-only: do not swallow logger.warning( "fetch_episode_context: basin context fetch failed (%s) — " "leaving basin_context=None for zone=%s", e, zone_id, ) basin_context = None return EpisodeContext( obs=obs, forecast=fcast, config=cfg, zone_ids=[zone_id], data_source=obs.source, basin_context=basin_context, ) # --------------------------------------------------------------------------- # Dynamics model integration # --------------------------------------------------------------------------- # The following classes extend era5_data_pipeline.py to support pre-training # TemporalDynamicsModel (physics_dynamics.py) on ERA5 reanalysis sequences. # # Key design decisions: # # 1. ZoneStateTensor construction from ForecastResult # ForecastResult.precip_mm is a tuple of horizon_days floats (mm/day). # ForecastResult.precip_p10/p90 are uncertainty bounds. # forecast_uncertainty is derived as normalised inter-quartile spread: # uncertainty_i = clip((p90_i - p10_i) / max(p90_i, 1e-6), 0, 1) # zone_uncertainty = mean(uncertainty_i over horizon) # This matches the exact calculation in WeatherForecastEnv._update_forecast_arrays(). # # 2. Consecutive pair extraction # ERA5 daily sequences are loaded as ZoneStateTensor snapshots, then # paired as (day_t, day_t+1). Each pair is one training sample. # The dynamics model learns to predict tomorrow's forecast from today's. # # 3. Multi-zone batching # All zones in a date window are fetched together. The ZoneStateTensor # batch dimension corresponds to different starting dates (not zones). # Zones are always the second dimension: [batch, n_zones, horizon_days]. # # 4. Synthetic fallback # When ERA5/Open-Meteo data is unavailable (no CDS credentials, offline # testing), get_consecutive_pairs() falls back to synthetic sequences. # The synthetic sequences are still useful for verifying the training loop # before real data is configured. try: import numpy as _np _NUMPY_FOR_DYNAMICS = True except ImportError: _NUMPY_FOR_DYNAMICS = False try: import torch as _torch _TORCH_FOR_DYNAMICS = True except ImportError: _TORCH_FOR_DYNAMICS = False def _forecast_result_to_arrays( forecast_result: Any, horizon_days: int, ) -> Tuple["np.ndarray", float]: """ Extract precipitation forecast array and scalar uncertainty from a ForecastResult, matching the exact logic in WeatherForecastEnv. Returns: precip_arr: np.float32 array of shape [horizon_days] (mm/day, clipped [0, 500]) uncertainty: float in [0, 1] — normalised inter-quartile spread This mirrors _update_forecast_arrays() in weather_forecast_env.py exactly, so the dynamics model sees the same representation as the RL policy. """ import numpy as np # --- Precipitation forecast --- precip_seq = list(forecast_result.precip_mm) if forecast_result.precip_mm else [] if len(precip_seq) < horizon_days: # Pad with zeros if forecast is shorter than horizon precip_seq = precip_seq + [0.0] * (horizon_days - len(precip_seq)) precip_arr = np.clip( np.array(precip_seq[:horizon_days], dtype=np.float32), 0.0, 500.0, ) # --- Uncertainty: normalised p90-p10 spread (matches env exactly) --- uncertainty = 0.5 # default if uncertainty bounds unavailable if forecast_result.precip_p90 and forecast_result.precip_p10: p90 = np.array(list(forecast_result.precip_p90)[:horizon_days], dtype=np.float32) p10 = np.array(list(forecast_result.precip_p10)[:horizon_days], dtype=np.float32) spread = np.clip( (p90 - p10) / np.maximum(np.abs(p90), 1e-6), 0.0, 1.0, ) uncertainty = float(np.mean(spread)) return precip_arr, uncertainty def _episode_context_to_state_tensor( contexts: List[Any], horizon_days: int, prior_belief: float = 0.5, ) -> "ZoneStateTensor": """ Convert a list of EpisodeContexts (one per zone) to a ZoneStateTensor with batch size 1. Args: contexts: List of EpisodeContext, one per zone. Length = n_zones. horizon_days: Forecast horizon to extract. prior_belief: Used when composite_risk() is unavailable. Returns: ZoneStateTensor with shapes [1, n_zones, horizon_days], [1, n_zones], [1, n_zones] """ import numpy as np # Import here to avoid circular dependency at module level from physics_dynamics import ZoneStateTensor n_zones = len(contexts) precip_arr = np.zeros((1, n_zones, horizon_days), dtype=np.float32) uncert_arr = np.zeros((1, n_zones), dtype=np.float32) belief_arr = np.zeros((1, n_zones), dtype=np.float32) for zi, ctx in enumerate(contexts): p_arr, unc = _forecast_result_to_arrays(ctx.forecast, horizon_days) precip_arr[0, zi, :] = p_arr uncert_arr[0, zi] = unc # Belief: blend composite_risk with prior (mirrors _per_zone_beliefs) try: signal = float(ctx.obs.composite_risk()) belief = 0.7 * prior_belief + 0.3 * signal except Exception: belief = prior_belief belief_arr[0, zi] = float(np.clip(belief, 0.0, 1.0)) import torch return ZoneStateTensor( precip=torch.from_numpy(precip_arr), uncertainty=torch.from_numpy(uncert_arr), belief=torch.from_numpy(belief_arr), ) def _build_synthetic_sequence( zone_ids: List[str], start_date: "datetime", n_days: int, horizon_days: int, config: Optional["ForecastConfig"] = None, ) -> List["ZoneStateTensor"]: """ Build a synthetic daily sequence of ZoneStateTensors for n_days. Each element represents the zone state on one day. Used as fallback when ERA5 data is unavailable, and for offline unit testing. The sequence is deterministic given zone_ids and start_date. """ from zone_observation import make_synthetic_episode_context, _stable_seed from physics_dynamics import ZoneStateTensor import numpy as np import torch cfg = config or ForecastConfig() sequence: List[ZoneStateTensor] = [] for day_offset in range(n_days): current_date = start_date + timedelta(days=day_offset) contexts = [] for zi, zid in enumerate(zone_ids): # Deterministic seed: zone + date + day_offset seed = _stable_seed( zid + current_date.isoformat() + str(day_offset) ) ctx = make_synthetic_episode_context(zone_id=zid, seed=seed) contexts.append(ctx) state = _episode_context_to_state_tensor( contexts, horizon_days, prior_belief=cfg.prior_belief ) sequence.append(state) return sequence def _build_era5_sequence( zone_ids: List[str], start_date: "datetime", n_days: int, horizon_days: int, config: Optional["ForecastConfig"] = None, ) -> List["ZoneStateTensor"]: """ Fetch n_days consecutive daily ERA5 snapshots for all zones. Returns a list of n_days ZoneStateTensors. Each tensor has shape [1, n_zones, horizon_days] for precip, [1, n_zones] for uncertainty/belief. Falls back to synthetic for any day/zone where ERA5 fetch fails. The fallback is per-day-per-zone, so partial ERA5 coverage is handled gracefully — days where ERA5 succeeded are real, failed days are synthetic. """ cfg = config or ForecastConfig() sequence = [] for day_offset in range(n_days): current_date = _ensure_utc(start_date + timedelta(days=day_offset)) date_range = (current_date, current_date + timedelta(days=horizon_days)) contexts = [] for zid in zone_ids: try: ctx = fetch_episode_context(zid, date_range, cfg) except Exception as e: logger.warning( "_build_era5_sequence: failed for zone=%s date=%s: %s — using synthetic", zid, current_date.date().isoformat(), e, ) from zone_observation import make_synthetic_episode_context, _stable_seed seed = _stable_seed(zid + current_date.isoformat()) ctx = make_synthetic_episode_context(zone_id=zid, seed=seed) contexts.append(ctx) state = _episode_context_to_state_tensor( contexts, horizon_days, prior_belief=cfg.prior_belief ) sequence.append(state) return sequence def get_consecutive_pairs( zone_ids: List[str], start_date: "datetime", n_days: int = 365, horizon_days: int = 14, config: Optional["ForecastConfig"] = None, use_real_data: bool = True, synthetic_fallback: bool = True, ) -> List[Tuple["ZoneStateTensor", "ZoneStateTensor"]]: """ Build a list of (current, next) ZoneStateTensor pairs for dynamics model training. This is the primary entry point for physics_dynamics.DynamicsTrainer. Each pair represents consecutive daily snapshots: current = zone state at day t next = zone state at day t+1 The dynamics model learns: given state at day t, predict state at day t+1. Args: zone_ids: Zone IDs to include. Must be registered via register_zone() if use_real_data=True. start_date: First day of the sequence window. n_days: Total number of days to fetch. Produces n_days-1 pairs. Recommended: 365 (one year) for meaningful coverage. horizon_days: Forecast horizon — must match WeatherForecastEnv config. config: ForecastConfig for source selection and prior. use_real_data: If True, attempt ERA5/Open-Meteo fetch before synthetic. Set False for offline testing or when CDS is unavailable. synthetic_fallback: If True (default), fall back to synthetic when real data fails. If False, raises on failure. Returns: List of (current, next) ZoneStateTensor tuples. Length = n_days - 1. Raises: RuntimeError: If use_real_data=True, synthetic_fallback=False, and any day fails to fetch from real sources. ImportError: If torch or numpy are not installed. Example: from datetime import datetime, timezone from era5_data_pipeline import register_zone, get_consecutive_pairs from zone_observation import GeoPolygon register_zone(GeoPolygon(zone_id="wheat_belt", coordinates=[...])) register_zone(GeoPolygon(zone_id="rice_delta", coordinates=[...])) pairs = get_consecutive_pairs( zone_ids=["wheat_belt", "rice_delta"], start_date=datetime(2022, 1, 1, tzinfo=timezone.utc), n_days=365, horizon_days=14, ) trainer = DynamicsTrainer(n_zones=2, horizon_days=14) history = trainer.train(pairs, epochs=100) trainer.save("./dynamics/pretrained.pt") """ if not _NUMPY_FOR_DYNAMICS: raise ImportError("numpy required for get_consecutive_pairs(). pip install numpy") if not _TORCH_FOR_DYNAMICS: raise ImportError("torch required for get_consecutive_pairs(). pip install torch") if n_days < 2: raise ValueError(f"n_days must be >= 2 to produce at least one pair, got {n_days}") start_date = _ensure_utc(start_date) logger.info( "get_consecutive_pairs: zones=%s start=%s n_days=%d horizon=%d real=%s", zone_ids, start_date.date().isoformat(), n_days, horizon_days, use_real_data, ) # Build sequence if use_real_data: try: sequence = _build_era5_sequence( zone_ids, start_date, n_days, horizon_days, config ) except Exception as e: if not synthetic_fallback: raise RuntimeError( f"ERA5 sequence build failed and synthetic_fallback=False: {e}" ) from e logger.warning( "ERA5 sequence build failed (%s) — falling back to full synthetic sequence", e, ) sequence = _build_synthetic_sequence( zone_ids, start_date, n_days, horizon_days, config ) else: sequence = _build_synthetic_sequence( zone_ids, start_date, n_days, horizon_days, config ) # Pair consecutive days pairs = [ (sequence[i], sequence[i + 1]) for i in range(len(sequence) - 1) ] logger.info( "get_consecutive_pairs: built %d pairs from %d-day sequence", len(pairs), n_days, ) return pairs def get_consecutive_pairs_multi_year( zone_ids: List[str], years: List[int], horizon_days: int = 14, config: Optional["ForecastConfig"] = None, use_real_data: bool = True, skip_on_failure: bool = True, ) -> List[Tuple["ZoneStateTensor", "ZoneStateTensor"]]: """ Fetch consecutive pairs across multiple years, concatenating them. Useful for building a large training dataset covering different climate regimes (El Niño, La Niña, drought years, anomalously wet years). Year boundaries are excluded (Dec 31 → Jan 1 pairs are dropped) since forecast continuity across year boundaries is not guaranteed in ERA5. Args: zone_ids: Zone IDs to include. years: List of calendar years to include, e.g. [2019, 2020, 2021]. horizon_days: Forecast horizon — must match env config. config: ForecastConfig. use_real_data: Attempt ERA5 before synthetic. skip_on_failure: If True, skip years that fail entirely and continue. If False, raise on any year failure. Returns: Concatenated list of (current, next) pairs from all years. Example: # 3 years × ~364 pairs/year ≈ 1092 training pairs pairs = get_consecutive_pairs_multi_year( zone_ids=["wheat_belt", "rice_delta"], years=[2019, 2020, 2021], horizon_days=14, ) """ all_pairs: List[Tuple["ZoneStateTensor", "ZoneStateTensor"]] = [] for year in years: start = _ensure_utc(datetime(year, 1, 1, tzinfo=timezone.utc)) # 365 days: excludes Dec 31 → Jan 1 boundary pair n_days = 366 if (year % 4 == 0 and (year % 100 != 0 or year % 400 == 0)) else 365 try: year_pairs = get_consecutive_pairs( zone_ids=zone_ids, start_date=start, n_days=n_days, horizon_days=horizon_days, config=config, use_real_data=use_real_data, synthetic_fallback=True, ) all_pairs.extend(year_pairs) logger.info("Year %d: added %d pairs (total=%d)", year, len(year_pairs), len(all_pairs)) except Exception as e: if not skip_on_failure: raise logger.warning("Year %d failed (%s) — skipped", year, e) if not all_pairs: raise RuntimeError( f"No pairs collected across years {years}. " "Check zone registration and data availability." ) logger.info( "get_consecutive_pairs_multi_year: %d total pairs from %d years", len(all_pairs), len(years), ) return all_pairs def compute_dataset_statistics( pairs: List[Tuple["ZoneStateTensor", "ZoneStateTensor"]], ) -> Dict[str, Any]: """ Compute normalisation statistics over a pairs dataset. Returns mean and std for precip, uncertainty, and belief across all pairs. These can be used to normalise inputs to the dynamics model for more stable training (especially for precipitation, which has high variance). Args: pairs: Output of get_consecutive_pairs() or get_consecutive_pairs_multi_year(). Returns: Dict with keys: precip_mean, precip_std, uncert_mean, uncert_std, belief_mean, belief_std, n_pairs, n_zones, horizon_days. """ if not pairs: raise ValueError("pairs is empty") import numpy as np import torch all_precip = [] all_uncert = [] all_belief = [] for curr, nxt in pairs: # Include both current and next in statistics for state in (curr, nxt): all_precip.append(state.precip.numpy().flatten()) all_uncert.append(state.uncertainty.numpy().flatten()) all_belief.append(state.belief.numpy().flatten()) precip_all = np.concatenate(all_precip) uncert_all = np.concatenate(all_uncert) belief_all = np.concatenate(all_belief) stats = { "precip_mean": float(np.mean(precip_all)), "precip_std": float(np.std(precip_all)) + 1e-8, "uncert_mean": float(np.mean(uncert_all)), "uncert_std": float(np.std(uncert_all)) + 1e-8, "belief_mean": float(np.mean(belief_all)), "belief_std": float(np.std(belief_all)) + 1e-8, "n_pairs": len(pairs), "n_zones": pairs[0][0].n_zones, "horizon_days": pairs[0][0].horizon_days, "precip_p95": float(np.percentile(precip_all, 95)), # useful for clipping "precip_max": float(np.max(precip_all)), } logger.info( "Dataset stats: n_pairs=%d precip_mean=%.1f±%.1f mm " "uncert_mean=%.3f belief_mean=%.3f", stats["n_pairs"], stats["precip_mean"], stats["precip_std"], stats["uncert_mean"], stats["belief_mean"], ) return stats