Skip to content
57 changes: 57 additions & 0 deletions tests/integration_test_10_cdsapi_era5_reanalysis.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
95 changes: 94 additions & 1 deletion tests/test_10_client_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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",
Expand Down
47 changes: 45 additions & 2 deletions xarray_ecmwf/client_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = []
Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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[
Expand Down
Loading