From 78fc049dde281bf1a6aa6bfa2e1cb6d571a88c48 Mon Sep 17 00:00:00 2001 From: Kiran Karnam Date: Wed, 15 Jul 2026 16:42:51 -0500 Subject: [PATCH 1/5] docs: design GFS cache integrity fixes Signed-off-by: Kiran Karnam --- .../2026-07-15-gfs-cache-integrity-design.md | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) create mode 100644 docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md diff --git a/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md b/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md new file mode 100644 index 000000000..381c06fb6 --- /dev/null +++ b/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md @@ -0,0 +1,74 @@ +# GFS Cache Integrity and Loop-Safe Coalescing + +## Goal + +Fix two cache regressions in `earth2studio.data.GFS`: + +1. A reused `GFS` instance must not retain synchronization objects bound to an old + event loop. +2. A cached byte range must be reused only when it matches the complete requested + range. + +The implementation must preserve atomic cache writes and coalesce concurrent +same-range downloads within one event loop. + +## Cache Identity and Validation + +Cache filenames will use a SHA-256 digest of an unambiguous serialization of the +remote path, byte offset, and byte length. Including `byte_length` prevents two +requests that begin at the same offset but request different ranges from sharing a +cache entry. This intentionally invalidates filenames produced by the old +path-and-offset-only scheme. + +For a finite byte-range request, an existing cache file is valid only when its size +equals the requested byte length. A wrong-size file is removed and treated as a +cache miss. Full-file requests have no expected size available locally, so they rely +on the atomic-write guarantee and the new cache namespace. + +Downloads continue to write to a unique temporary file and use `os.replace` only +after the complete payload has been written. A downloaded finite range whose size +does not match the request raises `IOError` and is never promoted into the cache. + +## Loop-Safe Miss Coalescing + +The persistent `dict[str, asyncio.Lock]` will be replaced with a registry of +in-flight download tasks keyed by `(current_event_loop, cache_path)`. + +When a cache miss occurs: + +1. The current event loop checks the registry for an existing task for that cache + path. +2. If present, the caller awaits the shared task through `asyncio.shield` so one + cancelled waiter cannot cancel the download for every waiter. +3. If absent, the loop creates and registers a download task. +4. A completion callback removes that exact task from the registry on success, + failure, or cancellation. + +Tasks from different event loops are never shared or awaited across loops. Separate +loops may perform duplicate downloads, but unique temporary paths and atomic replace +keep the final cache entry consistent. Because completed tasks are removed, the +registry does not retain event loops or grow with every historical cache path. + +## Error Handling + +- Wrong-size existing ranges are deleted and refetched. +- Wrong-size downloads raise `IOError` without leaving a final or temporary file. +- Download exceptions propagate to all same-loop waiters and remove the in-flight + registry entry so a later request can retry. +- Temporary files are removed in a `finally` block. + +## Tests + +Focused unit tests will verify: + +- an exact-size existing range is reused without a remote call; +- a wrong-size existing range is removed and refetched; +- different byte lengths produce different cache paths; +- concurrent same-loop misses make one remote call; +- the same `GFS` instance can coalesce misses in two successive event loops after + cache eviction; +- atomic writes, wrong-size download rejection, and failed-download cleanup remain + intact. + +The affected tests, Black, Ruff, MyPy for the changed surface, and `git diff --check` +must pass before handoff. From 3c8d9083775b48b12c70ed5a4002c5732ea18021 Mon Sep 17 00:00:00 2001 From: Kiran Karnam Date: Fri, 17 Jul 2026 15:05:03 -0500 Subject: [PATCH 2/5] feat: add deterministic batch runtime Signed-off-by: Kiran Karnam --- CHANGELOG.md | 2 + docs/modules/workflows.rst | 29 +++ earth2studio/batched_workflows.py | 311 ++++++++++++++++++++++++++++++ earth2studio/data/gfs.py | 117 +++++++++-- test/data/test_gfs.py | 249 +++++++++++++++++++++++- test/test_batched_workflows.py | 283 +++++++++++++++++++++++++++ 6 files changed, 974 insertions(+), 17 deletions(-) create mode 100644 earth2studio/batched_workflows.py create mode 100644 test/test_batched_workflows.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 95b4d1611..590fc005b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -78,6 +78,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 reflectivity, rain rate, and 1-hour accumulation (`OPERA`) - Added support for cumulative variables in ARCO data source - Added DLESyM-v0-ISCCP-ERA5 climate model +- Added a deterministic shared-resource batching API ### Changed @@ -94,6 +95,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Fixed ARCO data source `ARCO_TIME_STOP` fallback to 2025-12-31, reflecting the most recent available data in the bucket - Fixed `ZarrBackend` chunk metadata reload to skip coordinate arrays when reopening Zarr stores. +- Fixed GFS cache identity, validation, atomic writes, and loop-safe download coalescing ### Dependencies diff --git a/docs/modules/workflows.rst b/docs/modules/workflows.rst index 00b55f84b..95e884581 100644 --- a/docs/modules/workflows.rst +++ b/docs/modules/workflows.rst @@ -23,3 +23,32 @@ use cases. run.deterministic run.diagnostic run.ensemble + + +:mod:`earth2studio.batched_workflows`: Batched Workflows +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Utilities for executing deterministic forecast requests with reusable model and +data-source resources. + +.. automodule:: earth2studio.batched_workflows + :no-members: + :no-inherited-members: + +.. currentmodule:: earth2studio.batched_workflows + +.. autosummary:: + :nosignatures: + :toctree: generated/workflows/ + :template: class.rst + + DeterministicBatchRequest + DeterministicBatchResponse + DeterministicBatchRuntime + +.. autosummary:: + :nosignatures: + :toctree: generated/workflows/ + :template: function.rst + + run_deterministic_batch diff --git a/earth2studio/batched_workflows.py b/earth2studio/batched_workflows.py new file mode 100644 index 000000000..998a6e838 --- /dev/null +++ b/earth2studio/batched_workflows.py @@ -0,0 +1,311 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import os +import shutil +import uuid +from collections.abc import Callable, Sequence +from dataclasses import dataclass, replace +from pathlib import Path +from typing import Any, Literal + +from loguru import logger + +__all__ = [ + "DeterministicBatchRequest", + "DeterministicBatchResponse", + "DeterministicBatchRuntime", + "run_deterministic_batch", +] + + +@dataclass(frozen=True) +class DeterministicBatchRequest: + """Single deterministic forecast request in a shared-resource batch. + + Parameters + ---------- + model : str + Name of the prognostic model to run. + start_time : str + Forecast initialization time in ISO 8601 format. + nsteps : int + Number of forecast steps to execute. + output_path : str | Path + Path where the request's forecast dataset will be written. The path must not + already exist. + run_id : str | None, optional + Identifier used to correlate request logs and results, by default None + """ + + model: str + start_time: str + nsteps: int + output_path: str | Path + run_id: str | None = None + + +@dataclass(frozen=True) +class DeterministicBatchResponse: + """Per-request outcome from a deterministic forecast batch. + + Parameters + ---------- + model : str + Name of the prognostic model used for the request. + start_time : str + Forecast initialization time in ISO 8601 format. + nsteps : int + Number of forecast steps requested. + dataset_path : str + Path assigned to the request's forecast dataset. + status : Literal["succeeded", "failed"], optional + Request outcome, by default "succeeded" + error : str | None, optional + Failure details when ``status`` is ``"failed"``, by default None + """ + + model: str + start_time: str + nsteps: int + dataset_path: str + status: Literal["succeeded", "failed"] = "succeeded" + error: str | None = None + + +ModelLoader = Callable[[str], Any] +DataFactory = Callable[[], Any] +ForecastRunner = Callable[[DeterministicBatchRequest, Any, Any, Any], None] + + +class DeterministicBatchRuntime: + """Reusable deterministic inference resources for grouped requests. + + This runtime deliberately has no dependency on the Earth2Studio serve stack. + Callers such as PhysicsNeMo-Serve own queueing, grouping, scheduling, and + output registration; this object owns model/data setup and forecast calls. + + Parameters + ---------- + device : Any | None, optional + Device used for model inference. When None, CUDA is selected when available + and CPU otherwise, by default None + model_loader : ModelLoader | None, optional + Callable that loads a model from a normalized model name. Uses the built-in + DLWP loader when None, by default None + data_factory : DataFactory | None, optional + Callable that creates the shared data source. Uses GFS when None, by default + None. + runner : ForecastRunner | None, optional + Callable that executes one request with the shared model, data source, and + device. Uses the built-in deterministic workflow when None, by default None + """ + + def __init__( + self, + *, + device: Any | None = None, + model_loader: ModelLoader | None = None, + data_factory: DataFactory | None = None, + runner: ForecastRunner | None = None, + ) -> None: + self.device = _resolve_device(device) + self._model_loader = model_loader or _load_default_deterministic_model + self._data_factory = data_factory or _load_default_data_source + self._runner = runner + self._loaded_model_name: str | None = None + self._model: Any | None = None + self._data: Any | None = None + + def _ensure_loaded(self, model_name: str) -> tuple[Any, Any]: + normalized_model = _normalize_model_name(model_name) + if self._loaded_model_name is not None: + if normalized_model != self._loaded_model_name: + raise ValueError( + "DeterministicBatchRuntime can only hold one model at a time; " + f"loaded {self._loaded_model_name!r}, requested {normalized_model!r}" + ) + if self._model is not None and self._data is not None: + return self._model, self._data + + model = self._model_loader(normalized_model) + model = model.to(self.device) + data = self._data_factory() + self._model = model + self._data = data + self._loaded_model_name = normalized_model + return self._model, self._data + + def run(self, request: DeterministicBatchRequest) -> DeterministicBatchResponse: + """Run one deterministic forecast using the cached resources. + + Parameters + ---------- + request : DeterministicBatchRequest + Deterministic forecast request to execute. + + Returns + ------- + DeterministicBatchResponse + Request outcome. Execution errors produce failed responses. + """ + logger.info( + "Earth2 deterministic request start run_id={} model={} start_time={} nsteps={} output_path={}", + request.run_id, + request.model, + request.start_time, + request.nsteps, + request.output_path, + ) + final_path = Path(request.output_path) + staging_path = final_path.with_name( + f".{final_path.name}.tmp-{uuid.uuid4().hex}" + ) + staged_request = replace(request, output_path=staging_path) + try: + if final_path.exists(): + raise FileExistsError(f"output path already exists: {final_path}") + + model, data = self._ensure_loaded(request.model) + if self._runner is not None: + self._runner(staged_request, model, data, self.device) + else: + _run_default_deterministic_forecast( + request=staged_request, + model=model, + data=data, + device=self.device, + ) + os.replace(staging_path, final_path) + except Exception as exc: + shutil.rmtree(staging_path, ignore_errors=True) + logger.exception( + "Earth2 deterministic request failed run_id={} model={} start_time={} nsteps={}", + request.run_id, + request.model, + request.start_time, + request.nsteps, + ) + return DeterministicBatchResponse( + model=request.model, + start_time=request.start_time, + nsteps=request.nsteps, + dataset_path=str(request.output_path), + status="failed", + error=str(exc), + ) + + logger.info( + "Earth2 deterministic request succeeded run_id={} model={} start_time={} nsteps={} output_path={}", + request.run_id, + request.model, + request.start_time, + request.nsteps, + request.output_path, + ) + return DeterministicBatchResponse( + model=request.model, + start_time=request.start_time, + nsteps=request.nsteps, + dataset_path=str(request.output_path), + ) + + +def run_deterministic_batch( + requests: Sequence[DeterministicBatchRequest], + *, + runtime: DeterministicBatchRuntime | None = None, + device: Any | None = None, +) -> list[DeterministicBatchResponse]: + """Run deterministic forecast requests with shared model and data resources. + + Parameters + ---------- + requests : Sequence[DeterministicBatchRequest] + Compatible deterministic forecast requests to execute sequentially. + runtime : DeterministicBatchRuntime | None, optional + Runtime that owns shared inference resources. A runtime is created when None, + by default None + device : Any | None, optional + Device for a newly created runtime. Ignored when ``runtime`` is supplied, by + default None + + Returns + ------- + list[DeterministicBatchResponse] + Per-request outcomes in input order. An empty request sequence returns an + empty list. + """ + if not requests: + return [] + + batch_runtime = runtime or DeterministicBatchRuntime(device=device) + return [batch_runtime.run(request) for request in requests] + + +def _normalize_model_name(model_name: str) -> str: + normalized = model_name.strip().lower() + if not normalized: + raise ValueError("model name cannot be empty") + return normalized + + +def _resolve_device(device: Any | None) -> Any: + if device is not None: + return device + + import torch + + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +def _load_default_deterministic_model(model_name: str) -> Any: + if model_name != "dlwp": + raise ValueError("default deterministic batching supports only model='dlwp'") + + from earth2studio.models.px import DLWP + + package = DLWP.load_default_package() + return DLWP.load_model(package) + + +def _load_default_data_source() -> Any: + from earth2studio.data import GFS + + return GFS() + + +def _run_default_deterministic_forecast( + *, + request: DeterministicBatchRequest, + model: Any, + data: Any, + device: Any, +) -> None: + from earth2studio.io import ZarrBackend + from earth2studio.run import deterministic + from earth2studio.utils.time import to_time_array + + deterministic( + time=to_time_array([request.start_time]), + nsteps=request.nsteps, + prognostic=model, + data=data, + io=ZarrBackend(str(request.output_path)), + device=device, + ) diff --git a/earth2studio/data/gfs.py b/earth2studio/data/gfs.py index 36353701d..71e4b5c96 100644 --- a/earth2studio/data/gfs.py +++ b/earth2studio/data/gfs.py @@ -15,7 +15,9 @@ # limitations under the License. import asyncio +import functools import hashlib +import json import os import pathlib import shutil @@ -39,7 +41,7 @@ cancellable_to_thread, datasource_cache_root, gather_with_concurrency, - obstore_fetch_to_cache, + obstore_read_range, obstore_store_from_url, prep_data_inputs, prep_forecast_inputs, @@ -127,6 +129,9 @@ def __init__( self._async_workers = async_workers self._retries = retries self._tmp_cache_hash: str | None = None + self._cache_downloads: dict[ + tuple[asyncio.AbstractEventLoop, str], asyncio.Task[str] + ] = {} self.store: ObjectStore | None = None if source == "aws": @@ -390,7 +395,13 @@ async def fetch_array( byte_length=byte_length, ) # pygrib decode is blocking and GIL-bound; run in a thread with timeout - values = await cancellable_to_thread(_decode_gfs_grib, grib_file, timeout=30.0) + try: + values = await cancellable_to_thread( + _decode_gfs_grib, grib_file, timeout=30.0 + ) + except Exception: + pathlib.Path(grib_file).unlink(missing_ok=True) + raise return modifier(values) def _validate_time(self, times: list[datetime]) -> None: @@ -457,36 +468,110 @@ async def _fetch_remote_file( self, path: str, byte_offset: int = 0, byte_length: int | None = None ) -> str: """Fetches remote file into cache""" - # Hash the bucket-prefixed path (not the store-relative key) so warm - # caches populated before the obstore migration remain valid - sha = hashlib.sha256((path + str(byte_offset)).encode()) + cache_key = json.dumps( + (path, byte_offset, byte_length), separators=(",", ":") + ).encode() + sha = hashlib.sha256(cache_key) filename = sha.hexdigest() + cache_path = os.path.join(self.cache, filename) + pathlib.Path(cache_path).parent.mkdir(parents=True, exist_ok=True) + + if self._cache_file_is_valid(cache_path, byte_length): + return cache_path + + loop = asyncio.get_running_loop() + task_key = (loop, cache_path) + download_task = self._cache_downloads.get(task_key) + if download_task is None: + download_task = loop.create_task( + self._download_remote_file( + path, + cache_path, + byte_offset, + byte_length, + ) + ) + self._cache_downloads[task_key] = download_task + download_task.add_done_callback( + functools.partial(self._remove_cache_download, task_key) + ) + + return await asyncio.shield(download_task) + + async def _download_remote_file( + self, + path: str, + cache_path: str, + byte_offset: int, + byte_length: int | None, + ) -> str: + if self._cache_file_is_valid(cache_path, byte_length): + return cache_path if self.store is not None: key = path.removeprefix(self.GFS_BUCKET_NAME + "/") - return await obstore_fetch_to_cache( + data = await obstore_read_range( self.store, key, - self.cache, byte_offset=byte_offset, byte_length=byte_length, - cache_key=filename, ) - - if self.fs is None: + elif self.fs is not None: + data = await asyncio.to_thread( + self.fs.read_block, + path, + offset=byte_offset, + length=byte_length, + ) + else: raise ValueError("File system is not initialized") - # ncep FTP source (sync filesystem) - cache_path = os.path.join(self.cache, filename) - if not pathlib.Path(cache_path).is_file(): - data = await asyncio.to_thread( - self.fs.read_block, path, offset=byte_offset, length=byte_length + if byte_length is not None and len(data) != byte_length: + raise OSError( + "GFS cache download size mismatch " + f"remote={path} byte_offset={byte_offset} " + f"expected={byte_length} actual={len(data)}" ) - with open(cache_path, "wb") as file: + + tmp_path = f"{cache_path}.tmp.{os.getpid()}.{uuid.uuid4().hex}" + try: + with open(tmp_path, "wb") as file: await asyncio.to_thread(file.write, data) + await asyncio.to_thread(os.replace, tmp_path, cache_path) + finally: + pathlib.Path(tmp_path).unlink(missing_ok=True) return cache_path + def _cache_file_is_valid(self, cache_path: str, byte_length: int | None) -> bool: + cache_file = pathlib.Path(cache_path) + try: + if not cache_file.is_file(): + return False + actual_size = cache_file.stat().st_size + except FileNotFoundError: + return False + + if byte_length is None or actual_size == byte_length: + return True + + logger.warning( + "GFS cache size mismatch cache_path={} expected={} actual={}; refetching", + cache_path, + byte_length, + actual_size, + ) + cache_file.unlink(missing_ok=True) + return False + + def _remove_cache_download( + self, + task_key: tuple[asyncio.AbstractEventLoop, str], + download_task: asyncio.Future[str], + ) -> None: + if self._cache_downloads.get(task_key) is download_task: + self._cache_downloads.pop(task_key, None) + def _grib_uri(self, time: datetime, lead_time: timedelta) -> str: """Generates the URI for GFS grib files""" lead_hour = int(lead_time.total_seconds() // 3600) diff --git a/test/data/test_gfs.py b/test/data/test_gfs.py index 4ba8eb780..53823d8d2 100644 --- a/test/data/test_gfs.py +++ b/test/data/test_gfs.py @@ -15,17 +15,61 @@ # limitations under the License. import asyncio +import hashlib +import json import pathlib import shutil +import time from datetime import datetime, timedelta -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import numpy as np import pytest +import earth2studio.data.gfs as gfs_module from earth2studio.data import GFS, GFS_FX +def _run(coro): + return asyncio.run(coro) + + +class _FakeFS: + def __init__( + self, + data: bytes = b"fake-grib-bytes", + delay: float = 0.0, + error: Exception | None = None, + ) -> None: + self.data = data + self.delay = delay + self.error = error + self.calls: list[tuple[str, int, int | None]] = [] + + def read_block( + self, path: str, *, offset: int = 0, length: int | None = None + ) -> bytes: + self.calls.append((path, offset, length)) + if self.delay: + time.sleep(self.delay) + if self.error is not None: + raise self.error + return self.data + + +def _gfs_cache_file( + cache: str, + uri: str, + byte_offset: int = 0, + byte_length: int | None = None, +) -> pathlib.Path: + cache_key = json.dumps( + (uri, byte_offset, byte_length), separators=(",", ":") + ).encode() + filename = hashlib.sha256(cache_key).hexdigest() + return pathlib.Path(cache) / filename + + @pytest.mark.slow @pytest.mark.xfail @pytest.mark.timeout(30) @@ -148,6 +192,209 @@ def test_gfs_cache(time, variable, cache): pass +def test_gfs_fetch_remote_file_reuses_existing_cache(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS() + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + cached_file = _gfs_cache_file(ds.cache, uri, byte_offset=12, byte_length=4) + cached_file.parent.mkdir(parents=True, exist_ok=True) + cached_file.write_bytes(b"data") + + result = _run(ds._fetch_remote_file(uri, byte_offset=12, byte_length=4)) + + assert result == str(cached_file) + assert fake_fs.calls == [] + + +def test_gfs_fetch_remote_file_refetches_wrong_size_cache(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"data") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + cached_file = _gfs_cache_file(ds.cache, uri, byte_offset=12, byte_length=4) + cached_file.parent.mkdir(parents=True, exist_ok=True) + cached_file.write_bytes(b"stale-data") + + result = _run(ds._fetch_remote_file(uri, byte_offset=12, byte_length=4)) + + assert result == str(cached_file) + assert cached_file.read_bytes() == b"data" + assert fake_fs.calls == [(uri, 12, 4)] + + +def test_gfs_fetch_remote_file_keys_cache_by_complete_range(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"data") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + short_range = _run(ds._fetch_remote_file(uri, byte_offset=3, byte_length=4)) + fake_fs.data = b"payload" + long_range = _run(ds._fetch_remote_file(uri, byte_offset=3, byte_length=7)) + + assert short_range != long_range + assert pathlib.Path(short_range).read_bytes() == b"data" + assert pathlib.Path(long_range).read_bytes() == b"payload" + assert fake_fs.calls == [(uri, 3, 4), (uri, 3, 7)] + + +def test_gfs_fetch_remote_file_writes_cache_atomically(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"payload") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + replace = MagicMock(wraps=gfs_module.os.replace) + monkeypatch.setattr(gfs_module.os, "replace", replace) + + result = _run(ds._fetch_remote_file(uri, byte_offset=3, byte_length=7)) + + cache_path = pathlib.Path(result) + replace.assert_called_once() + temporary_path, published_path = map(pathlib.Path, replace.call_args.args) + assert published_path == cache_path + assert temporary_path.parent == cache_path.parent + assert temporary_path.name.startswith(f"{cache_path.name}.tmp.") + assert cache_path.read_bytes() == b"payload" + assert fake_fs.calls == [(uri, 3, 7)] + assert not list(cache_path.parent.glob("*.tmp.*")) + + +def test_gfs_fetch_remote_file_rejects_size_mismatch(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"short") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + with pytest.raises(IOError, match="GFS cache download size mismatch"): + _run(ds._fetch_remote_file(uri, byte_offset=3, byte_length=7)) + + cache_dir = pathlib.Path(ds.cache) + assert fake_fs.calls == [(uri, 3, 7)] + assert not [path for path in cache_dir.iterdir() if path.is_file()] + assert not list(cache_dir.glob("*.tmp.*")) + + +def test_gfs_fetch_remote_file_cleans_up_failed_download(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(error=RuntimeError("boom")) + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + with pytest.raises(RuntimeError, match="boom"): + _run(ds._fetch_remote_file(uri, byte_offset=3, byte_length=7)) + + cache_dir = pathlib.Path(ds.cache) + assert not [path for path in cache_dir.iterdir() if path.is_file()] + assert not list(cache_dir.glob("*.tmp.*")) + + +def test_gfs_fetch_remote_file_coalesces_concurrent_cache_misses(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"data", delay=0.05) + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + async def fetch_twice(): + return await asyncio.gather( + ds._fetch_remote_file(uri, byte_offset=5, byte_length=4), + ds._fetch_remote_file(uri, byte_offset=5, byte_length=4), + ) + + first, second = _run(fetch_twice()) + + assert first == second + assert pathlib.Path(first).read_bytes() == b"data" + assert fake_fs.calls == [(uri, 5, 4)] + assert ds._cache_downloads == {} + + +def test_gfs_fetch_remote_file_coalesces_after_event_loop_change(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"data", delay=0.05) + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + async def fetch_twice(): + return await asyncio.gather( + ds._fetch_remote_file(uri, byte_offset=5, byte_length=4), + ds._fetch_remote_file(uri, byte_offset=5, byte_length=4), + ) + + first_paths = _run(fetch_twice()) + pathlib.Path(first_paths[0]).unlink() + second_paths = _run(fetch_twice()) + + assert first_paths[0] == first_paths[1] + assert second_paths[0] == second_paths[1] + assert first_paths[0] == second_paths[0] + assert pathlib.Path(second_paths[0]).read_bytes() == b"data" + assert fake_fs.calls == [(uri, 5, 4), (uri, 5, 4)] + assert ds._cache_downloads == {} + + +def test_gfs_fetch_array_removes_cached_file_on_grib_open_failure( + tmp_path, monkeypatch +): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"badgrb") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + + def _raise_bad_grib(_path): + raise RuntimeError("bad grib") + + monkeypatch.setattr(gfs_module.pygrib, "open", _raise_bad_grib) + + with pytest.raises(RuntimeError, match="bad grib"): + _run(ds.fetch_array(uri, byte_offset=2, byte_length=6, modifier=lambda x: x)) + + cache_file = _gfs_cache_file(ds.cache, uri, byte_offset=2, byte_length=6) + assert fake_fs.calls == [(uri, 2, 6)] + assert not cache_file.exists() + assert not list(cache_file.parent.glob("*.tmp.*")) + + +def test_gfs_fetch_array_keeps_cached_file_on_modifier_failure(tmp_path, monkeypatch): + monkeypatch.setenv("EARTH2STUDIO_DATA_CACHE", str(tmp_path)) + ds = GFS(cache=True) + fake_fs = _FakeFS(data=b"valid!") + ds.fs = fake_fs + + uri = "noaa-gfs-bdp-pds/gfs.20260101/00/atmos/gfs.t00z.pgrb2.0p25.f000" + grbs = MagicMock() + grbs.__getitem__.return_value.values = np.array([1.0]) + monkeypatch.setattr(gfs_module.pygrib, "open", lambda _path: grbs) + + def fail_modifier(_values: np.ndarray) -> np.ndarray: + raise RuntimeError("modifier boom") + + with pytest.raises(RuntimeError, match="modifier boom"): + _run(ds.fetch_array(uri, 2, 6, fail_modifier)) + + cache_file = _gfs_cache_file(ds.cache, uri, byte_offset=2, byte_length=6) + assert cache_file.read_bytes() == b"valid!" + grbs.close.assert_called_once_with() + + @pytest.mark.slow @pytest.mark.xfail @pytest.mark.timeout(15) diff --git a/test/test_batched_workflows.py b/test/test_batched_workflows.py new file mode 100644 index 000000000..a0e8d3ccc --- /dev/null +++ b/test/test_batched_workflows.py @@ -0,0 +1,283 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import Mock + +import pytest + +from earth2studio.batched_workflows import ( + DeterministicBatchRequest, + DeterministicBatchRuntime, + run_deterministic_batch, +) + + +class FakeModel: + def __init__(self, calls: list[tuple[Any, ...]]) -> None: + self.calls = calls + + def to(self, device: Any) -> FakeModel: + self.calls.append(("model.to", device)) + return self + + +def test_run_deterministic_batch_reuses_runtime_resources(tmp_path: Path) -> None: + calls: list[tuple[Any, ...]] = [] + + def model_loader(model_name: str) -> FakeModel: + calls.append(("load_model", model_name)) + return FakeModel(calls) + + def data_factory() -> object: + calls.append(("load_data", None)) + return object() + + def runner( + request: DeterministicBatchRequest, + _model: object, + _data: object, + device: Any, + ) -> None: + Path(request.output_path).mkdir(parents=True) + calls.append(("run", request.run_id, device)) + + runtime = DeterministicBatchRuntime( + device="cpu", + model_loader=model_loader, + data_factory=data_factory, + runner=runner, + ) + requests = [ + DeterministicBatchRequest( + model="custom", + start_time="2026-01-01T00:00:00Z", + nsteps=1, + output_path=tmp_path / f"forecast-{index}.zarr", + run_id=f"run-{index}", + ) + for index in range(2) + ] + + responses = run_deterministic_batch(requests, runtime=runtime) + + assert [response.status for response in responses] == ["succeeded", "succeeded"] + assert calls == [ + ("load_model", "custom"), + ("model.to", "cpu"), + ("load_data", None), + ("run", "run-0", "cpu"), + ("run", "run-1", "cpu"), + ] + assert [response.dataset_path for response in responses] == [ + str(tmp_path / "forecast-0.zarr"), + str(tmp_path / "forecast-1.zarr"), + ] + assert all(Path(response.dataset_path).is_dir() for response in responses) + assert not list(tmp_path.glob(".*.tmp-*")) + + +def test_run_deterministic_batch_uses_default_components( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + import earth2studio.data as data_module + import earth2studio.io as io_module + import earth2studio.models.px as model_module + import earth2studio.run as run_module + import earth2studio.utils.time as time_module + + output_path = tmp_path / "default.zarr" + request = DeterministicBatchRequest( + model="dlwp", + start_time="2026-01-01T00:00:00Z", + nsteps=2, + output_path=output_path, + ) + package = object() + data = object() + converted_time = object() + model = FakeModel([]) + load_package = Mock(return_value=package) + load_model = Mock(return_value=model) + load_data = Mock(return_value=data) + to_time_array = Mock(return_value=converted_time) + + def deterministic(**kwargs: Any) -> None: + Path(kwargs["io"]).mkdir(parents=True) + + forecast = Mock(side_effect=deterministic) + + monkeypatch.setattr(model_module.DLWP, "load_default_package", load_package) + monkeypatch.setattr(model_module.DLWP, "load_model", load_model) + monkeypatch.setattr(data_module, "GFS", load_data) + monkeypatch.setattr(time_module, "to_time_array", to_time_array) + monkeypatch.setattr(io_module, "ZarrBackend", Path) + monkeypatch.setattr(run_module, "deterministic", forecast) + + response = run_deterministic_batch([request], device="cpu")[0] + + assert response.status == "succeeded" + assert response.dataset_path == str(output_path) + assert model.calls == [("model.to", "cpu")] + load_package.assert_called_once_with() + load_model.assert_called_once_with(package) + load_data.assert_called_once_with() + to_time_array.assert_called_once_with([request.start_time]) + staging_path = forecast.call_args.kwargs["io"] + assert staging_path.parent == output_path.parent + assert staging_path.name.startswith(f".{output_path.name}.tmp-") + forecast.assert_called_once_with( + time=converted_time, + nsteps=2, + prognostic=model, + data=data, + io=staging_path, + device="cpu", + ) + assert output_path.is_dir() + assert not list(tmp_path.glob(".*.tmp-*")) + + +def test_run_deterministic_batch_returns_item_failure(tmp_path: Path) -> None: + def model_loader(_model_name: str) -> FakeModel: + return FakeModel([]) + + def runner( + request: DeterministicBatchRequest, + _model: object, + _data: object, + _device: Any, + ) -> None: + output_path = Path(request.output_path) + output_path.mkdir(parents=True) + (output_path / "data").write_text(str(request.run_id)) + if request.run_id == "bad": + raise RuntimeError("boom") + + runtime = DeterministicBatchRuntime( + device="cpu", + model_loader=model_loader, + data_factory=object, + runner=runner, + ) + + responses = run_deterministic_batch( + [ + DeterministicBatchRequest( + model="dlwp", + start_time="2026-01-01T00:00:00Z", + nsteps=1, + output_path=tmp_path / "good.zarr", + run_id="good", + ), + DeterministicBatchRequest( + model="dlwp", + start_time="2026-01-01T00:00:00Z", + nsteps=1, + output_path=tmp_path / "bad.zarr", + run_id="bad", + ), + ], + runtime=runtime, + ) + + assert responses[0].status == "succeeded" + assert responses[1].status == "failed" + assert responses[1].error is not None + assert "boom" in responses[1].error + assert (tmp_path / "good.zarr" / "data").read_text() == "good" + assert not (tmp_path / "bad.zarr").exists() + assert not list(tmp_path.glob(".*.tmp-*")) + + +def test_run_deterministic_batch_preserves_existing_output(tmp_path: Path) -> None: + output_path = tmp_path / "existing.zarr" + output_path.mkdir() + marker = output_path / "marker" + marker.write_text("original") + runner_called = False + + def runner( + _request: DeterministicBatchRequest, + _model: object, + _data: object, + _device: Any, + ) -> None: + nonlocal runner_called + runner_called = True + + runtime = DeterministicBatchRuntime(device="cpu", runner=runner) + response = runtime.run( + DeterministicBatchRequest( + model="dlwp", + start_time="2026-01-01T00:00:00Z", + nsteps=1, + output_path=output_path, + ) + ) + + assert response.status == "failed" + assert response.error == f"output path already exists: {output_path}" + assert marker.read_text() == "original" + assert not runner_called + assert not list(tmp_path.glob(".*.tmp-*")) + + +def test_runtime_does_not_commit_partial_state_on_data_failure() -> None: + calls: list[tuple[Any, ...]] = [] + data_load_attempt = 0 + + def model_loader(model_name: str) -> FakeModel: + calls.append(("load_model", model_name)) + return FakeModel(calls) + + def data_factory() -> object: + nonlocal data_load_attempt + data_load_attempt += 1 + calls.append(("load_data", data_load_attempt)) + if data_load_attempt == 1: + raise RuntimeError("data boom") + return object() + + runtime = DeterministicBatchRuntime( + device="cpu", + model_loader=model_loader, + data_factory=data_factory, + ) + + with pytest.raises(RuntimeError, match="data boom"): + runtime._ensure_loaded("dlwp") + + assert runtime._model is None + assert runtime._data is None + assert runtime._loaded_model_name is None + + model, data = runtime._ensure_loaded("dlwp") + + assert runtime._model is model + assert runtime._data is data + assert runtime._loaded_model_name == "dlwp" + assert calls == [ + ("load_model", "dlwp"), + ("model.to", "cpu"), + ("load_data", 1), + ("load_model", "dlwp"), + ("model.to", "cpu"), + ("load_data", 2), + ] From c12ebf9629fd48364c53d4fec24f60805a9ab0b8 Mon Sep 17 00:00:00 2001 From: Kiran Karnam Date: Fri, 17 Jul 2026 15:25:11 -0500 Subject: [PATCH 3/5] remove design doc --- .../2026-07-15-gfs-cache-integrity-design.md | 74 ------------------- 1 file changed, 74 deletions(-) delete mode 100644 docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md diff --git a/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md b/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md deleted file mode 100644 index 381c06fb6..000000000 --- a/docs/superpowers/specs/2026-07-15-gfs-cache-integrity-design.md +++ /dev/null @@ -1,74 +0,0 @@ -# GFS Cache Integrity and Loop-Safe Coalescing - -## Goal - -Fix two cache regressions in `earth2studio.data.GFS`: - -1. A reused `GFS` instance must not retain synchronization objects bound to an old - event loop. -2. A cached byte range must be reused only when it matches the complete requested - range. - -The implementation must preserve atomic cache writes and coalesce concurrent -same-range downloads within one event loop. - -## Cache Identity and Validation - -Cache filenames will use a SHA-256 digest of an unambiguous serialization of the -remote path, byte offset, and byte length. Including `byte_length` prevents two -requests that begin at the same offset but request different ranges from sharing a -cache entry. This intentionally invalidates filenames produced by the old -path-and-offset-only scheme. - -For a finite byte-range request, an existing cache file is valid only when its size -equals the requested byte length. A wrong-size file is removed and treated as a -cache miss. Full-file requests have no expected size available locally, so they rely -on the atomic-write guarantee and the new cache namespace. - -Downloads continue to write to a unique temporary file and use `os.replace` only -after the complete payload has been written. A downloaded finite range whose size -does not match the request raises `IOError` and is never promoted into the cache. - -## Loop-Safe Miss Coalescing - -The persistent `dict[str, asyncio.Lock]` will be replaced with a registry of -in-flight download tasks keyed by `(current_event_loop, cache_path)`. - -When a cache miss occurs: - -1. The current event loop checks the registry for an existing task for that cache - path. -2. If present, the caller awaits the shared task through `asyncio.shield` so one - cancelled waiter cannot cancel the download for every waiter. -3. If absent, the loop creates and registers a download task. -4. A completion callback removes that exact task from the registry on success, - failure, or cancellation. - -Tasks from different event loops are never shared or awaited across loops. Separate -loops may perform duplicate downloads, but unique temporary paths and atomic replace -keep the final cache entry consistent. Because completed tasks are removed, the -registry does not retain event loops or grow with every historical cache path. - -## Error Handling - -- Wrong-size existing ranges are deleted and refetched. -- Wrong-size downloads raise `IOError` without leaving a final or temporary file. -- Download exceptions propagate to all same-loop waiters and remove the in-flight - registry entry so a later request can retry. -- Temporary files are removed in a `finally` block. - -## Tests - -Focused unit tests will verify: - -- an exact-size existing range is reused without a remote call; -- a wrong-size existing range is removed and refetched; -- different byte lengths produce different cache paths; -- concurrent same-loop misses make one remote call; -- the same `GFS` instance can coalesce misses in two successive event loops after - cache eviction; -- atomic writes, wrong-size download rejection, and failed-download cleanup remain - intact. - -The affected tests, Black, Ruff, MyPy for the changed surface, and `git diff --check` -must pass before handoff. From 45173c286322b228f1aa532b331aa655b20cfb84 Mon Sep 17 00:00:00 2001 From: kirankarnam Date: Fri, 17 Jul 2026 18:51:41 -0500 Subject: [PATCH 4/5] Update earth2studio/data/gfs.py Consolidating both the write and the rename into a single helper passed to one asyncio.to_thread Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- earth2studio/data/gfs.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/earth2studio/data/gfs.py b/earth2studio/data/gfs.py index 71e4b5c96..3c5ec9ab7 100644 --- a/earth2studio/data/gfs.py +++ b/earth2studio/data/gfs.py @@ -535,9 +535,12 @@ async def _download_remote_file( tmp_path = f"{cache_path}.tmp.{os.getpid()}.{uuid.uuid4().hex}" try: - with open(tmp_path, "wb") as file: - await asyncio.to_thread(file.write, data) - await asyncio.to_thread(os.replace, tmp_path, cache_path) + def _write_and_replace() -> None: + with open(tmp_path, "wb") as file: + file.write(data) + os.replace(tmp_path, cache_path) + + await asyncio.to_thread(_write_and_replace) finally: pathlib.Path(tmp_path).unlink(missing_ok=True) From c7e3b772e3d28fb0a5b5c1d5b9669533ee392daf Mon Sep 17 00:00:00 2001 From: Kiran Karnam Date: Thu, 23 Jul 2026 16:22:21 -0500 Subject: [PATCH 5/5] add close call --- earth2studio/batched_workflows.py | 11 +++++++++ test/test_batched_workflows.py | 40 +++++++++++++++++++++++++++++++ 2 files changed, 51 insertions(+) diff --git a/earth2studio/batched_workflows.py b/earth2studio/batched_workflows.py index 998a6e838..53e054050 100644 --- a/earth2studio/batched_workflows.py +++ b/earth2studio/batched_workflows.py @@ -151,6 +151,17 @@ def _ensure_loaded(self, model_name: str) -> tuple[Any, Any]: self._loaded_model_name = normalized_model return self._model, self._data + def close(self) -> None: + """Release the model and data resources cached by this runtime. + + The runtime remains reusable after closing. Its model and data resources are + loaded again when the next request runs. + """ + + self._model = None + self._data = None + self._loaded_model_name = None + def run(self, request: DeterministicBatchRequest) -> DeterministicBatchResponse: """Run one deterministic forecast using the cached resources. diff --git a/test/test_batched_workflows.py b/test/test_batched_workflows.py index a0e8d3ccc..fe4e47107 100644 --- a/test/test_batched_workflows.py +++ b/test/test_batched_workflows.py @@ -93,6 +93,46 @@ def runner( assert not list(tmp_path.glob(".*.tmp-*")) +def test_runtime_close_releases_resources_and_allows_reuse() -> None: + calls: list[tuple[Any, ...]] = [] + + def model_loader(model_name: str) -> FakeModel: + calls.append(("load_model", model_name)) + return FakeModel(calls) + + def data_factory() -> object: + data = object() + calls.append(("load_data", data)) + return data + + runtime = DeterministicBatchRuntime( + device="cpu", + model_loader=model_loader, + data_factory=data_factory, + ) + first_model, first_data = runtime._ensure_loaded("custom") + + runtime.close() + + assert runtime._model is None + assert runtime._data is None + assert runtime._loaded_model_name is None + + runtime.close() + second_model, second_data = runtime._ensure_loaded("custom") + + assert second_model is not first_model + assert second_data is not first_data + assert calls == [ + ("load_model", "custom"), + ("model.to", "cpu"), + ("load_data", first_data), + ("load_model", "custom"), + ("model.to", "cpu"), + ("load_data", second_data), + ] + + def test_run_deterministic_batch_uses_default_components( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: