diff --git a/tests/integration_test_10_cdsapi_era5_reanalysis.py b/tests/integration_test_10_cdsapi_era5_reanalysis.py index d2704d9..cfeca68 100644 --- a/tests/integration_test_10_cdsapi_era5_reanalysis.py +++ b/tests/integration_test_10_cdsapi_era5_reanalysis.py @@ -14,6 +14,16 @@ "time": ["00:00", "12:00"], } +REQUEST_YEARS = { + "dataset": "reanalysis-era5-single-levels", + "product_type": ["reanalysis"], + "variable": ["2m_temperature"], + "year": ["2021", "2022", "2023"], + "month": ["01"], + "day": ["01"], + "time": ["00:00", "12:00"], +} + def test_open_dataset() -> None: res = xr.open_dataset(REQUEST, engine="ecmwf") # type: ignore @@ -195,6 +205,53 @@ def test_cds_era5_big_slice_time_month() -> None: assert res.size == 1 +def test_compare_chunked_no_chunked_year() -> None: + ds = xr.open_dataset( + REQUEST, # type: ignore + engine="ecmwf", + request_chunks={"year": 1}, + chunks={}, + ) + da = ds.data_vars["t2m"] + + res = da.sel(time=slice("2022-07-01", "2022-07-16")).mean().compute() + + assert isinstance(res, xr.DataArray) + assert res.size == 1 + + +def test_compare_chunked_no_chunked_year_n() -> None: + # no year chunking + ds0 = xr.open_dataset(REQUEST_YEARS, engine="ecmwf", chunks={}) # type: ignore + assert ds0.chunks["time"] == (6,) + res0 = ds0.data_vars["t2m"].load() + + # 1 year chunking + ds1 = xr.open_dataset( + REQUEST_YEARS, # type: ignore + engine="ecmwf", + request_chunks={"year": 1}, + chunks={}, + ) + assert ds1.chunks["time"] == (2, 2, 2) + res1 = ds1.data_vars["t2m"].load() + + # 2 years chunking + ds2 = xr.open_dataset( + REQUEST_YEARS, # type: ignore + engine="ecmwf", + request_chunks={"year": 2}, + chunks={}, + ) + assert ds2.chunks["time"] == (4, 2) + res2 = ds2.data_vars["t2m"].load() + + assert (res0 - res1).shape == res0.shape + assert (res0 == res1).all() + assert (res0 - res2).shape == res0.shape + assert (res0 == res2).all() + + def test_cds_era5_small_slice_time_longitute() -> None: ds = xr.open_dataset( REQUEST, # type: ignore diff --git a/tests/test_10_client_common.py b/tests/test_10_client_common.py index a60c55d..07c80b7 100644 --- a/tests/test_10_client_common.py +++ b/tests/test_10_client_common.py @@ -79,7 +79,7 @@ def test_build_chunk_ymd_requests( k: v for k, v in filter( lambda x: x[1], [("year", years), ("month", months), ("day", days)] - ) + ) # why this filter if they are never empty? } request["time"] = ALL_TIMES @@ -106,6 +106,99 @@ def test_build_chunk_ymd_requests( client_common.build_chunk_ymd_requests(request, request_chunks) +@pytest.mark.parametrize( + "years, months, days", + [ + (["2022"], ALL_MONTHS, ALL_DAYS), + (["2020"], ALL_MONTHS, ALL_DAYS), # leap year + ( + [str(y) for y in range(2015, 2020)], + ALL_MONTHS, + ALL_DAYS, + ), # consecutive years + (["2010", "2015", "2020"], ALL_MONTHS, ALL_DAYS), + (["2010", "2015", "2020"], ["01", "02", "09"], ALL_DAYS), + (["2010", "2015", "2020"], ["01", "02", "09"], ["01", "29", "30", "31"]), + ], +) +def test_build_chunk_ymd_year_requests( + years: list[str], months: list[str], days: list[str] +) -> None: + request_chunks = {"year": 1} + request = { + "year": years, + "month": months, + "day": days, + "time": ALL_TIMES, + } + + ( + time, + time_chunk, + time_chunk_requests, + ) = client_common.build_chunk_ymd_requests(request, request_chunks) + total_days = 0 + for year, month, day in itertools.product( + map(int, request["year"]), + map(int, request["month"]), + map(int, request["day"]), + ): + if day <= calendar.monthrange(year, month)[1]: + total_days += 1 + + assert len(time) == 24 * total_days # 24h default value + assert len(time_chunk_requests) == len(years) + assert isinstance(time_chunk, tuple) + assert sum(time_chunk) == len(time) + + offset = 0 + for (start, chunk_request), year_str, size in zip( + time_chunk_requests, years, time_chunk + ): + assert start == offset + assert chunk_request == {"year": [year_str]} + offset += size + + +@pytest.mark.parametrize( + "year_chunk, expected_year_groups", + [ + (1, [["2015"], ["2016"], ["2017"]]), + (2, [["2015", "2016"], ["2017"]]), # leftover chunk + (5, [["2015", "2016", "2017"]]), # N larger than n years + ], +) +def test_build_chunk_ymd_year_requests_groups_years( + year_chunk: int, expected_year_groups: list[list[str]] +) -> None: + request = { + "year": ["2015", "2016", "2017"], + "month": ["01"], + "day": ["01"], + "time": ["00:00"], + } # one time stamp per year + request_chunks = {"year": year_chunk} + time, time_chunk, time_chunk_requests = client_common.build_chunk_ymd_requests( + request, request_chunks + ) + assert len(time) == 3 + assert [chunk_request["year"] for _, chunk_request in time_chunk_requests] == ( + expected_year_groups + ) + assert time_chunk == tuple(len(group) for group in expected_year_groups) + + +def test_build_chunk_ymd_year_requests_invalid() -> None: + request = { + "year": ["2022"], + "month": ["01"], + "day": ["01"], + "time": ["00:00"], + } + with pytest.raises(ValueError): + client_common.build_chunk_ymd_requests(request, {"year": 0}) + + def test_build_chunk_request() -> None: coord, chunk, chunk_request = client_common.build_chunks_header_requests( dim="x", diff --git a/xarray_ecmwf/client_common.py b/xarray_ecmwf/client_common.py index 99ca461..37519c0 100644 --- a/xarray_ecmwf/client_common.py +++ b/xarray_ecmwf/client_common.py @@ -147,7 +147,7 @@ def build_chunk_ymd_month_requests( list[tuple[int, dict[str, Any]]], ]: if request_chunks["month"] != 1: - raise ValueError("split on day values != 1 not supported") + raise ValueError("split on month values != 1 not supported") datetimes: list[np.datetime64] = [] chunk_requests = [] @@ -189,12 +189,14 @@ def build_chunk_ymd_requests( assert len(time_chunk_keys) == 1 time_chunk_key = time_chunk_keys[0] - assert time_chunk_key in set(["day", "month"]) + assert time_chunk_key in set(["day", "month", "year"]) # this check is redundant if time_chunk_key == "month": out = build_chunk_ymd_month_requests(request, request_chunks) elif time_chunk_key == "day": out = build_chunk_ymd_day_requests(request, request_chunks) + elif time_chunk_key == "year": + out = build_chunk_ymd_year_requests(request, request_chunks) return out @@ -230,6 +232,47 @@ def build_chunk_ymd_day_requests( return np.array(times), len(request["time"]), chunk_requests +def build_chunk_ymd_year_requests( + request: dict[str, Any], request_chunks: dict[str, int] +) -> tuple[ + np.typing.NDArray[np.datetime64], + int | tuple[int, ...], + list[tuple[int, dict[str, Any]]], +]: + year_chunk_size = request_chunks["year"] + if year_chunk_size < 1: + raise ValueError("split on year values < 1 not supported") + + datetimes: list[np.datetime64] = [] + chunk_requests: list[tuple[int, dict[str, Any]]] = [] + chunks: list[int] = [] + years = request["year"] + istart = 0 + while istart < len(years): + istop = min(istart + year_chunk_size, len(years)) + year_values = years[istart:istop] + start = len(datetimes) + chunk = 0 + for year in year_values: + assert len(year) == 4 + for month in request["month"]: + assert len(month) == 2 + ndays = calendar.monthrange(int(year), int(month))[1] + for day in request["day"]: + assert len(day) == 2 + if int(day) > ndays: + break + for time in request["time"]: + assert len(time) == 5 + chunk += 1 + datetime = np.datetime64(f"{year}-{month}-{day}T{time}", "ns") + datetimes.append(datetime) + chunks.append(chunk) + chunk_requests.append((start, {"year": year_values})) + istart += year_chunk_size + return np.array(datetimes), tuple(chunks), chunk_requests + + def build_time_chunk_requests( request: dict[str, Any], request_chunks: dict[str, int], sep: str = "/" ) -> tuple[