forcingkit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,458 @@
1
+ """
2
+ NYOFS (NOAA New York / New Jersey Operational Forecast System) fetcher.
3
+
4
+ Provides Initial Conditions (IC) and Open Boundary Conditions (OBC) from Princeton Ocean
5
+ Model (POM) structured curvilinear grid via OPeNDAP.
6
+ """
7
+
8
+ import numpy as np
9
+ import xarray as xr
10
+ import pandas as pd
11
+ import logging
12
+ from typing import Optional, Tuple
13
+ from forcingkit import settings
14
+
15
+ logger = logging.getLogger(__name__)
16
+
17
+
18
+ def get_metadata() -> dict:
19
+ """Returns metadata for the NYOFS system."""
20
+ return {
21
+ "id": "nyofs",
22
+ "name": "NOAA NYOFS (NY/NJ Harbor)",
23
+ "resolution_approx_m": 100.0,
24
+ "type_desc": "Structured curvilinear POM grid",
25
+ "domain_bbox": [-74.3, 40.2, -73.3, 41.1],
26
+ }
27
+
28
+
29
+ def supports_bbox(bbox: list[float]) -> bool:
30
+ """Check if the requested bbox is within NYOFS domain."""
31
+ min_lon, min_lat, max_lon, max_lat = bbox
32
+ domain_bbox = get_metadata()["domain_bbox"]
33
+ domain_min_lon, domain_min_lat, domain_max_lon, domain_max_lat = domain_bbox
34
+
35
+ # Request bbox must be fully contained within NYOFS domain
36
+ if (
37
+ min_lon < domain_min_lon
38
+ or max_lon > domain_max_lon
39
+ or min_lat < domain_min_lat
40
+ or max_lat > domain_max_lat
41
+ ):
42
+ return False
43
+ return True
44
+
45
+
46
+ def _to_dap_url(url: str) -> str:
47
+ """Convert http(s) URL to pydap dap2:// scheme."""
48
+ return url.replace("https://", "dap2://").replace("http://", "dap2://")
49
+
50
+
51
+ def _get_nyofs_url(target_dt: pd.Timestamp) -> Tuple[str, str]:
52
+ """
53
+ Resolve NYOFS data access URL and mode.
54
+
55
+ Returns (access_mode, url_or_pattern) where access_mode is one of:
56
+ - "fmrc": single FMRC aggregated OPeNDAP URL
57
+ - "ncei": NCEI THREDDS file pattern (requires enumeration)
58
+
59
+ FMRC is preferred for recent data (< 31 days).
60
+ NCEI is fallback for historical data (> 31 days).
61
+ """
62
+ now = pd.Timestamp.now(tz="UTC").tz_localize(None)
63
+ age_days = (now - target_dt).total_seconds() / 86400
64
+
65
+ # Prefer FMRC aggregation for recent data (< 7 days)
66
+ if 0 <= age_days <= 6:
67
+ fmrc_url = (
68
+ "https://opendap.co-ops.nos.noaa.gov/thredds/dodsC/"
69
+ "NYOFS/fmrc/Aggregated_7_day_NYOFS_Fields_Forecast_best.ncd"
70
+ )
71
+ return ("fmrc", fmrc_url)
72
+
73
+ # Historical: AWS S3 or NCEI file-per-hour. Naming convention changed 2024-09-09
74
+ # and the archive migrated to AWS S3 starting 2024-01-01.
75
+ if target_dt >= pd.Timestamp("2024-01-01"):
76
+ pattern = (
77
+ "s3://noaa-nos-ofs-pds/nyofs/netcdf/"
78
+ "{yyyy}/{mm}/{dd}/nyofs.t{cc}z.{yyyymmdd}.fields.{type}{hhh:03d}.nc"
79
+ )
80
+ return ("aws_s3", pattern)
81
+ elif target_dt >= pd.Timestamp("2024-09-09"):
82
+ # Post-Sept 9 2024 naming: nyofs.tCCz.YYYYMMDD.fields.nHHH.nc
83
+ pattern = (
84
+ "https://www.ncei.noaa.gov/thredds/dodsC/model-nyofs-files/"
85
+ "{yyyy}/{mm}/{dd}/nyofs.t{cc}z.{yyyymmdd}.fields.{type}{hhh:03d}.nc"
86
+ )
87
+ else:
88
+ # Legacy naming: nos.nyofs.fields.nHHH.YYYYMMDD.tCCz.nc
89
+ pattern = (
90
+ "https://www.ncei.noaa.gov/thredds/dodsC/model-nyofs-files/"
91
+ "{yyyy}/{mm}/nos.nyofs.fields.{type}{hhh:03d}.{yyyymmdd}.t{cc}z.nc"
92
+ )
93
+ return ("ncei", pattern)
94
+
95
+
96
+ def _enumerate_ncei_nyofs_files(
97
+ pattern: str, start_dt: pd.Timestamp, end_dt: pd.Timestamp
98
+ ) -> list[str]:
99
+ """
100
+ Enumerate NCEI NYOFS file URLs for a time range.
101
+
102
+ Picks the best 6-hourly cycle (00/06/12/18Z) and builds per-hour file URLs.
103
+ Spans cycle boundaries if duration_hours > remaining hours in first cycle.
104
+ """
105
+ # Find the most recent 6-hourly cycle before start_dt
106
+ cycle_hours = [0, 6, 12, 18]
107
+ cycle_before = None
108
+ for ch in reversed(cycle_hours):
109
+ test_dt = start_dt.replace(hour=ch, minute=0, second=0, microsecond=0)
110
+ if test_dt <= start_dt:
111
+ cycle_before = test_dt
112
+ break
113
+
114
+ if cycle_before is None:
115
+ # Fall back to previous day's last cycle
116
+ cycle_before = (start_dt - pd.Timedelta(days=1)).replace(
117
+ hour=18, minute=0, second=0, microsecond=0
118
+ )
119
+
120
+ # Enumerate files from cycle_before to end_dt
121
+ files = []
122
+ current_dt = cycle_before
123
+ current_cycle_dt = cycle_before
124
+
125
+ while current_dt <= end_dt:
126
+ # Check if we've crossed into a new cycle (6 hours apart)
127
+ if (current_dt - current_cycle_dt).total_seconds() >= 6 * 3600:
128
+ current_cycle_dt = current_dt.replace(minute=0, second=0, microsecond=0)
129
+ # Snap to 6-hourly boundaries
130
+ cycle_hour = (current_cycle_dt.hour // 6) * 6
131
+ current_cycle_dt = current_cycle_dt.replace(hour=cycle_hour)
132
+
133
+ # Calculate file index (hour offset from cycle start)
134
+ hour_offset = int((current_dt - current_cycle_dt).total_seconds() / 3600) + 1
135
+ forecast_or_nowcast = "f" if current_dt > current_cycle_dt else "n"
136
+
137
+ # Format template variables
138
+ fmt_vars = {
139
+ "yyyy": current_cycle_dt.strftime("%Y"),
140
+ "mm": current_cycle_dt.strftime("%m"),
141
+ "dd": current_cycle_dt.strftime("%d"),
142
+ "yyyymmdd": current_cycle_dt.strftime("%Y%m%d"),
143
+ "cc": current_cycle_dt.strftime("%H"),
144
+ "hhh": hour_offset,
145
+ "type": forecast_or_nowcast,
146
+ }
147
+
148
+ # Build URL
149
+ url = pattern.format(**fmt_vars)
150
+ files.append(url)
151
+
152
+ current_dt += pd.Timedelta(hours=1)
153
+
154
+ return files
155
+
156
+
157
+ def _open_nyofs_dataset(
158
+ access_mode: str,
159
+ url_or_pattern: str,
160
+ target_dt: pd.Timestamp,
161
+ end_dt: Optional[pd.Timestamp] = None,
162
+ ) -> Optional[xr.Dataset]:
163
+ """
164
+ Open NYOFS dataset via OPeNDAP (FMRC or NCEI).
165
+
166
+ Args:
167
+ access_mode: "fmrc" or "ncei"
168
+ url_or_pattern: Full URL (fmrc) or pattern (ncei)
169
+ target_dt: Start datetime for slicing
170
+ end_dt: End datetime (required for ncei, optional for fmrc)
171
+
172
+ Returns:
173
+ xr.Dataset or None if fetch fails
174
+ """
175
+ try:
176
+ if access_mode == "fmrc":
177
+ logger.info(f"Opening FMRC aggregation: {url_or_pattern}")
178
+ dap_url = _to_dap_url(url_or_pattern)
179
+ ds = xr.open_dataset(dap_url, engine="pydap")
180
+
181
+ # Determine time variable name
182
+ time_var = "time" if "time" in ds.coords else "ocean_time"
183
+
184
+ # FMRC aggregations can have non-monotonic time indices; sort before slicing
185
+ ds = ds.sortby(time_var)
186
+
187
+ # Slice to requested time range
188
+ if end_dt is None:
189
+ ds_t = ds.sel({time_var: target_dt}, method="nearest")
190
+ else:
191
+ ds_t = ds.sel({time_var: slice(target_dt, end_dt)})
192
+ if ds_t.sizes[time_var] == 0:
193
+ ds_t = ds.sel({time_var: target_dt}, method="nearest").expand_dims(
194
+ time_var
195
+ )
196
+ logger.warning("Exact time range empty, fell back to nearest.")
197
+
198
+ return ds_t
199
+
200
+ elif access_mode in ("ncei", "aws_s3"):
201
+ mode_name = "AWS S3" if access_mode == "aws_s3" else "NCEI"
202
+ logger.info(
203
+ f"Enumerating NYOFS {mode_name} files from {target_dt} to {end_dt}"
204
+ )
205
+ if end_dt is None:
206
+ end_dt = target_dt + pd.Timedelta(hours=1)
207
+
208
+ files = _enumerate_ncei_nyofs_files(url_or_pattern, target_dt, end_dt)
209
+ logger.info(f"Opening {len(files)} NYOFS {mode_name} files...")
210
+
211
+ if access_mode == "aws_s3":
212
+ try:
213
+ logger.info("Using xarray.open_mfdataset for parallel S3 access...")
214
+ ds_t = xr.open_mfdataset(
215
+ files,
216
+ engine="h5netcdf",
217
+ parallel=True,
218
+ storage_options={"anon": True},
219
+ data_vars="minimal",
220
+ coords="minimal",
221
+ compat="override",
222
+ )
223
+ return ds_t
224
+ except Exception as e:
225
+ logger.error(
226
+ f"Failed to open/concat NYOFS S3 files via mfdataset: {e}"
227
+ )
228
+ return None
229
+ else:
230
+ import concurrent.futures
231
+
232
+ datasets: list[Optional[xr.Dataset]] = [None] * len(files)
233
+ fail_counts = [0]
234
+
235
+ def _fetch_file(args):
236
+ i, f = args
237
+ if fail_counts[0] >= 3:
238
+ return i, None
239
+ try:
240
+ if i % 10 == 0 or i == 1 or i == len(files):
241
+ logger.info(
242
+ f"[{i}/{len(files)}] Fetching/Opening NYOFS {mode_name} file: {f.split('/')[-1] if 's3' in f else f}"
243
+ )
244
+ ds_file = xr.open_dataset(_to_dap_url(f), engine="pydap")
245
+ return i, ds_file
246
+ except Exception as e:
247
+ logger.warning(
248
+ f"Failed to open NYOFS {mode_name} file {f}: {e}"
249
+ )
250
+ fail_counts[0] += 1
251
+ return i, None
252
+
253
+ max_workers = settings.max_workers()
254
+ with concurrent.futures.ThreadPoolExecutor(
255
+ max_workers=max_workers
256
+ ) as executor:
257
+ for i, ds_file in executor.map(_fetch_file, enumerate(files, 1)):
258
+ if ds_file is not None:
259
+ datasets[i - 1] = ds_file
260
+
261
+ opened = [ds for ds in datasets if ds is not None]
262
+
263
+ if not opened:
264
+ logger.error(f"No NYOFS {mode_name} files could be opened.")
265
+ return None
266
+
267
+ # Concatenate along time dimension
268
+ return xr.concat(opened, dim="time", join="override")
269
+
270
+ else:
271
+ logger.error(f"Unknown access_mode: {access_mode}")
272
+ return None
273
+
274
+ except Exception as e:
275
+ logger.error(f"Failed to open NYOFS dataset ({access_mode}): {e}")
276
+ return None
277
+
278
+
279
+ _VAR_CANDIDATES: dict[str, list[str]] = {
280
+ "u": ["u", "water_u", "u_eastward"],
281
+ "v": ["v", "water_v", "v_northward"],
282
+ "temp": ["temp", "water_temp", "temperature", "sea_water_temperature"],
283
+ "salt": ["salt", "salinity", "sea_water_salinity"],
284
+ "zeta": ["zeta", "sea_surface_height", "ssh"],
285
+ }
286
+
287
+
288
+ def _resolve_var(ds: xr.Dataset, role: str) -> str:
289
+ """Return the first candidate name for *role* that exists in *ds*."""
290
+ for name in _VAR_CANDIDATES[role]:
291
+ if name in ds:
292
+ return name
293
+ available = list(ds.data_vars)
294
+ raise KeyError(
295
+ f"No variable found for role '{role}'. "
296
+ f"Tried: {_VAR_CANDIDATES[role]}. Available: {available}"
297
+ )
298
+
299
+
300
+ def _c_grid_to_rho(
301
+ u_raw: np.ndarray,
302
+ v_raw: np.ndarray,
303
+ ) -> Tuple[np.ndarray, np.ndarray]:
304
+ """
305
+ Interpolate Arakawa C-grid u,v face values to rho-points by averaging
306
+ adjacent pairs. Works on already-subset arrays of any leading shape:
307
+
308
+ u_rho[..., j, i] = 0.5 * (u[..., j, i] + u[..., j, i+1])
309
+ v_rho[..., j, i] = 0.5 * (v[..., j, i] + v[..., j+1, i])
310
+
311
+ The boundary column/row is filled by extrapolation (copy-edge) so the
312
+ output shape always matches the input shape. This is valid for both
313
+ IC arrays (sigma, eta, xi) and OBC arrays (time, sigma, eta, xi).
314
+ """
315
+ u_rho = np.empty_like(u_raw)
316
+ u_rho[..., :-1] = 0.5 * (u_raw[..., :-1] + u_raw[..., 1:])
317
+ u_rho[..., -1] = u_raw[..., -1] # boundary: copy edge
318
+
319
+ v_rho = np.empty_like(v_raw)
320
+ v_rho[..., :-1, :] = 0.5 * (v_raw[..., :-1, :] + v_raw[..., 1:, :])
321
+ v_rho[..., -1, :] = v_raw[..., -1, :] # boundary: copy edge
322
+
323
+ return u_rho, v_rho
324
+
325
+
326
+ def fetch_nyofs_boundary_conditions(
327
+ start_date: str, duration_hours: int, bbox: list[float]
328
+ ) -> Optional[xr.Dataset]:
329
+ """
330
+ Fetch 4D Ocean State (u, v) from NOAA NYOFS over a time range for OBC.
331
+
332
+ Output dimensions: (time, depth, eta, xi) matching the standard OBC contract.
333
+
334
+ Args:
335
+ start_date: ISO format datetime string
336
+ duration_hours: Duration of the boundary condition period
337
+ bbox: [min_lon, min_lat, max_lon, max_lat]
338
+
339
+ Returns:
340
+ xr.Dataset with dims (time, depth, eta, xi) or None if fetch fails
341
+ """
342
+ min_lon, min_lat, max_lon, max_lat = bbox
343
+
344
+ if not supports_bbox(bbox):
345
+ logger.info(f"Bounding box {bbox} outside NYOFS domain.")
346
+ return None
347
+
348
+ target_dt = pd.to_datetime(start_date)
349
+ if target_dt.tzinfo is not None:
350
+ target_dt = target_dt.tz_convert("UTC").tz_localize(None)
351
+
352
+ end_dt = target_dt + pd.Timedelta(hours=duration_hours)
353
+
354
+ access_mode, url_or_pattern = _get_nyofs_url(target_dt)
355
+ logger.info(
356
+ f"Attempting to fetch NYOFS OBC ({duration_hours}h) from {access_mode.upper()}"
357
+ )
358
+
359
+ ds_t = _open_nyofs_dataset(access_mode, url_or_pattern, target_dt, end_dt)
360
+ if ds_t is None:
361
+ return None
362
+
363
+ try:
364
+ # Spatial subsetting
365
+ lon_var = "lon"
366
+ lat_var = "lat"
367
+
368
+ mask = (
369
+ (ds_t[lon_var] >= min_lon)
370
+ & (ds_t[lon_var] <= max_lon)
371
+ & (ds_t[lat_var] >= min_lat)
372
+ & (ds_t[lat_var] <= max_lat)
373
+ & (ds_t.get("mask", 1) == 1)
374
+ )
375
+
376
+ ds_sub = ds_t.where(mask, drop=True)
377
+
378
+ # Check if the bounding box yielded zero valid water points
379
+ if any(size == 0 for size in ds_sub.sizes.values()):
380
+ logger.warning(
381
+ "No valid NYOFS ocean points found in bounding box (size is 0)."
382
+ )
383
+ return None
384
+
385
+ logger.info("Executing OPeNDAP download for NYOFS OBC subset...")
386
+ ds_sub = ds_sub.compute()
387
+
388
+ # Resolve actual variable names
389
+ u_var = _resolve_var(ds_sub, "u")
390
+ v_var = _resolve_var(ds_sub, "v")
391
+ logger.info(f"NYOFS OBC variable mapping: u={u_var}, v={v_var}")
392
+
393
+ # Detect dimensions
394
+ all_dims = list(ds_sub[u_var].dims)
395
+ time_var = "time" if "time" in ds_sub.coords else "ocean_time"
396
+ sigma_dim = None
397
+ for candidate in ["sigma", "s_rho", "depth", "siglay"]:
398
+ if candidate in all_dims:
399
+ sigma_dim = candidate
400
+ break
401
+
402
+ if sigma_dim is None:
403
+ logger.error(
404
+ f"No recognized sigma dimension in NYOFS OBC. Dims: {all_dims}"
405
+ )
406
+ return None
407
+
408
+ # Spatial dimensions
409
+ spatial_dims = [d for d in all_dims if d != time_var and d != sigma_dim]
410
+ if len(spatial_dims) < 2:
411
+ logger.error(f"Unexpected spatial dims after subset: {spatial_dims}")
412
+ return None
413
+
414
+ eta_dim = spatial_dims[0] # noqa: F841 - kept for readability / debug logging
415
+ xi_dim = spatial_dims[1] # noqa: F841
416
+
417
+ # C-grid interpolation: vectorized over all leading dims (time, sigma, ...)
418
+ u_raw = ds_sub[u_var].values
419
+ v_raw = ds_sub[v_var].values
420
+ u_rho, v_rho = _c_grid_to_rho(
421
+ u_raw.astype(np.float32), v_raw.astype(np.float32)
422
+ )
423
+
424
+ n_sigma = u_rho.shape[1]
425
+ n_eta = u_rho.shape[2]
426
+ n_xi = u_rho.shape[3]
427
+
428
+ # Map sigma to pseudo-depth
429
+ n_depth = n_sigma
430
+ depths = np.linspace(-50, 0, n_depth).astype(np.float32)
431
+
432
+ # Get time coordinates
433
+ out_times = ds_sub[time_var].values
434
+
435
+ # Build output dataset
436
+ ds_out = xr.Dataset(
437
+ data_vars={
438
+ "u": (("time", "depth", "eta", "xi"), u_rho),
439
+ "v": (("time", "depth", "eta", "xi"), v_rho),
440
+ },
441
+ coords={
442
+ "time": out_times,
443
+ "depth": depths,
444
+ "eta": np.arange(n_eta, dtype=np.float32),
445
+ "xi": np.arange(n_xi, dtype=np.float32),
446
+ },
447
+ attrs={"type": "NOAA NYOFS OBC", "source": "NOAA CO-OPS"},
448
+ )
449
+
450
+ logger.info(
451
+ f"Successfully processed NYOFS OBC data. "
452
+ f"Shape: u{ds_out['u'].shape}, {len(out_times)} time steps"
453
+ )
454
+ return ds_out
455
+
456
+ except Exception as e:
457
+ logger.error(f"Failed to process NYOFS OBC data: {e}")
458
+ return None
forcingkit/settings.py ADDED
@@ -0,0 +1,72 @@
1
+ """Environment settings, read under the forcingkit names.
2
+
3
+ forcingkit was ecodata-cache until 2026-10-05. For one release the old environment variable names
4
+ and the old default cache directory still work, each with a FutureWarning (shown by default, as
5
+ it is meant for whoever runs the service) naming the replacement. The next release drops them.
6
+ """
7
+
8
+ import logging
9
+ import os
10
+ import warnings
11
+ from pathlib import Path
12
+
13
+ logger = logging.getLogger("forcingkit")
14
+
15
+ # New name -> old names still read, in order. The cache directory had two names: the service
16
+ # routes read COASTAL_SIM_DATA_CACHE_DIR and the fetchers ECODATA_CACHE_CACHE_DIR.
17
+ LEGACY_ENV: dict[str, tuple[str, ...]] = {
18
+ "FORCINGKIT_CACHE_DIR": ("ECODATA_CACHE_CACHE_DIR", "COASTAL_SIM_DATA_CACHE_DIR"),
19
+ "FORCINGKIT_MAX_WORKERS": ("ECODATA_CACHE_MAX_WORKERS",),
20
+ # The elevation service, renamed topobathysim -> topobathykit on 2026-10-05.
21
+ "TOPOBATHYKIT_URL": ("TOPOBATHYSIM_URL",),
22
+ }
23
+
24
+ DEFAULT_CACHE_DIR = Path("~/.cache/forcingkit")
25
+ LEGACY_CACHE_DIR = Path("~/.cache/ecodata-cache")
26
+
27
+ _warned: set[str] = set()
28
+
29
+
30
+ def _deprecated(old: str, new: str) -> None:
31
+ if old in _warned:
32
+ return
33
+ _warned.add(old)
34
+ message = f"{old} is deprecated; use {new} (forcingkit was ecodata-cache)."
35
+ logger.warning(message)
36
+ warnings.warn(message, FutureWarning, stacklevel=3)
37
+
38
+
39
+ def env(name: str, default: str | None = None) -> str | None:
40
+ """The value of `name`, else of its first set legacy name (with a warning), else `default`."""
41
+ value = os.environ.get(name)
42
+ if value is not None:
43
+ return value
44
+ for old in LEGACY_ENV.get(name, ()):
45
+ value = os.environ.get(old)
46
+ if value is not None:
47
+ _deprecated(old, name)
48
+ return value
49
+ return default
50
+
51
+
52
+ def cache_dir(*parts: str) -> str:
53
+ """The cache root (FORCINGKIT_CACHE_DIR, default ~/.cache/forcingkit), joined with `parts`.
54
+
55
+ While the default directory does not exist and ~/.cache/ecodata-cache does, the old directory
56
+ is used, with a warning to move it.
57
+ """
58
+ configured = env("FORCINGKIT_CACHE_DIR")
59
+ if configured is not None:
60
+ root = Path(configured).expanduser()
61
+ else:
62
+ root = DEFAULT_CACHE_DIR.expanduser()
63
+ legacy = LEGACY_CACHE_DIR.expanduser()
64
+ if not root.exists() and legacy.exists():
65
+ _deprecated(str(legacy), f"{root} (move the directory)")
66
+ root = legacy
67
+ return os.path.join(str(root), *parts)
68
+
69
+
70
+ def max_workers(default: int = 4) -> int:
71
+ """Threads for parallel remote reads (FORCINGKIT_MAX_WORKERS)."""
72
+ return int(env("FORCINGKIT_MAX_WORKERS", str(default)) or default)
@@ -0,0 +1,146 @@
1
+ """Write a time series to Zarr one record at a time, and publish it only when it is complete.
2
+
3
+ Long deliveries (a 168 h parent ocean, an hourly atmosphere) used to be held in memory and written
4
+ in one shot, so a failure part way lost everything, and a store left behind by a crashed write was
5
+ served as a cache hit by the existence check. Here records are appended to `<path>.partial` with
6
+ one time record per chunk (so a reader that needs one hour reads one chunk), the store is checked
7
+ for the expected number of contiguous hourly records, marked `complete`, consolidated, and only
8
+ then renamed to its final path.
9
+ """
10
+
11
+ import logging
12
+ import os
13
+ import shutil
14
+ from typing import Any, Optional
15
+
16
+ import numpy as np
17
+ import pandas as pd
18
+ import xarray as xr
19
+ import zarr
20
+
21
+ logger = logging.getLogger(__name__)
22
+
23
+ PARTIAL_SUFFIX = ".partial"
24
+
25
+
26
+ def store_is_complete(path: str, schemas: tuple[str, ...] = ()) -> bool:
27
+ """True when `path` is a finished store a cache may serve.
28
+
29
+ A store whose `schema` attribute is one of `schemas` (those written by `StreamingZarrWriter`)
30
+ must also carry `complete = True`; any other store counts as complete if it exists, since it was
31
+ written in one shot after its data was assembled.
32
+ """
33
+ if not os.path.isdir(path):
34
+ return False
35
+ try:
36
+ attrs = dict(zarr.open_group(path, mode="r", zarr_format=2).attrs)
37
+ except Exception:
38
+ return False
39
+ if attrs.get("schema") in schemas:
40
+ return bool(attrs.get("complete", False))
41
+ return True
42
+
43
+
44
+ class StreamingZarrWriter:
45
+ """Append hourly records to a Zarr v2 store and publish it atomically.
46
+
47
+ `static` holds the variables without a time dimension (coordinates, masks); `record_dims` maps
48
+ each per-record variable to its dimensions without time. Times are encoded as seconds since
49
+ the first record.
50
+ """
51
+
52
+ def __init__(
53
+ self,
54
+ final_path: str,
55
+ static: xr.Dataset,
56
+ record_dims: dict[str, tuple[str, ...]],
57
+ expected_records: int,
58
+ attrs: Optional[dict] = None,
59
+ record_attrs: Optional[dict[str, dict]] = None,
60
+ ):
61
+ self.final_path = final_path
62
+ self.partial_path = final_path + PARTIAL_SUFFIX
63
+ self.static = static
64
+ self.record_dims = record_dims
65
+ self.expected_records = expected_records
66
+ self.attrs = dict(attrs or {})
67
+ self.record_attrs = record_attrs or {}
68
+ self.times: list[pd.Timestamp] = []
69
+ if os.path.exists(self.partial_path):
70
+ shutil.rmtree(self.partial_path)
71
+
72
+ def append(self, time: pd.Timestamp, record: dict[str, np.ndarray]) -> None:
73
+ missing = set(self.record_dims) - set(record)
74
+ if missing:
75
+ raise ValueError(f"record at {time} lacks {sorted(missing)}")
76
+ t = pd.Timestamp(time)
77
+ if self.times and t <= self.times[-1]:
78
+ raise ValueError(f"record at {t} is not after {self.times[-1]}")
79
+ data_vars = {
80
+ name: (("time",) + dims, np.asarray(record[name], dtype=np.float32)[None])
81
+ for name, dims in self.record_dims.items()
82
+ }
83
+ ds = xr.Dataset(data_vars, coords={"time": [t.to_datetime64()]})
84
+ for name, a in self.record_attrs.items():
85
+ if name in ds:
86
+ ds[name].attrs.update(a)
87
+ if not self.times:
88
+ ds = xr.merge([ds, self.static])
89
+ ds.attrs.update(self.attrs)
90
+ encoding: dict[str, dict[str, Any]] = {
91
+ name: {"chunks": (1,) + ds[name].shape[1:]} for name in self.record_dims
92
+ }
93
+ encoding["time"] = {
94
+ "units": f"seconds since {t.strftime('%Y-%m-%dT%H:%M:%S')}",
95
+ "dtype": "float64",
96
+ }
97
+ ds.to_zarr(
98
+ self.partial_path,
99
+ mode="w",
100
+ zarr_format=2,
101
+ consolidated=False,
102
+ encoding=encoding,
103
+ )
104
+ else:
105
+ ds.to_zarr(
106
+ self.partial_path,
107
+ append_dim="time",
108
+ zarr_format=2,
109
+ consolidated=False,
110
+ )
111
+ self.times.append(t)
112
+
113
+ def close(self) -> str:
114
+ """Verify, mark complete, consolidate and publish. Raises (and removes the partial store)
115
+ if the record count or the hourly spacing is wrong."""
116
+ try:
117
+ if len(self.times) != self.expected_records:
118
+ raise RuntimeError(
119
+ f"{self.final_path}: {len(self.times)} records written, "
120
+ f"{self.expected_records} expected"
121
+ )
122
+ gaps = np.diff(np.array([t.value for t in self.times])) / 3.6e12
123
+ if len(gaps) and not np.allclose(gaps, 1.0):
124
+ raise RuntimeError(
125
+ f"{self.final_path}: records are not hourly and contiguous "
126
+ f"(spacings {sorted(set(np.round(gaps, 3)))} h)"
127
+ )
128
+ # Appends rewrite the group's attributes from the appended record, so the store's
129
+ # own attributes are set here, once, with the completion mark.
130
+ group = zarr.open_group(self.partial_path, mode="r+", zarr_format=2)
131
+ group.attrs.update(self.attrs)
132
+ group.attrs["complete"] = True
133
+ group.attrs["records"] = len(self.times)
134
+ zarr.consolidate_metadata(self.partial_path, zarr_format=2)
135
+ except Exception:
136
+ self.abort()
137
+ raise
138
+ if os.path.exists(self.final_path):
139
+ shutil.rmtree(self.final_path)
140
+ os.rename(self.partial_path, self.final_path)
141
+ logger.info(f"Published {self.final_path} ({len(self.times)} records)")
142
+ return self.final_path
143
+
144
+ def abort(self) -> None:
145
+ if os.path.exists(self.partial_path):
146
+ shutil.rmtree(self.partial_path)