Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions cf_xarray/accessor.py
Original file line number Diff line number Diff line change
Expand Up @@ -326,7 +326,8 @@ def _get_custom_criteria(
if key in criteria_map:
for criterion, patterns in criteria_map[key].items():
for var in variables:
if regex_match(patterns, variables[var].attrs.get(criterion, "")):
metadata = ChainMap(variables[var].attrs, variables[var].encoding)
if regex_match(patterns, metadata.get(criterion, "")):
results.update((var,))
# also check name specifically since not in attributes
elif (
Expand Down Expand Up @@ -393,9 +394,10 @@ def _get_axis_coord(obj: DataArray | Dataset, key: str) -> list[str]:
results: set = set()
for coord in search_in:
var = crds[coord]
metadata = ChainMap(var.attrs, var.encoding)
if key in coordinate_criteria:
for criterion, expected in coordinate_criteria[key].items():
if var.attrs.get(criterion, None) in expected:
if metadata.get(criterion, None) in expected:
results.update((coord,))
if criterion == "units":
# deal with pint-backed objects
Expand Down Expand Up @@ -815,7 +817,7 @@ def _get_with_standard_name(
if isinstance(obj, DataArray):
obj = obj.coords.to_dataset()
for vname, var in obj._variables.items():
stdname = var.attrs.get("standard_name", None)
stdname = ChainMap(var.attrs, var.encoding).get("standard_name", None)
if stdname == name:
varnames.append(vname)

Expand Down Expand Up @@ -2152,8 +2154,9 @@ def standard_names(self) -> dict[str, list[Hashable]]:

vardict: dict[str, list[Hashable]] = {}
for k, v in variables.items():
if "standard_name" in v.attrs:
std_name = v.attrs["standard_name"]
metadata = ChainMap(v.attrs, v.encoding)
if "standard_name" in metadata:
std_name = metadata["standard_name"]
vardict[std_name] = vardict.setdefault(std_name, []) + [k]

return {std: sort_maybe_hashable(v) for std, v in vardict.items()}
Expand Down
37 changes: 37 additions & 0 deletions cf_xarray/tests/test_accessor.py
Original file line number Diff line number Diff line change
Expand Up @@ -569,6 +569,43 @@ def test_keys(obj, expected):
assert actual == expected


@pytest.mark.parametrize(
("key", "metadata"),
[
("longitude", {"standard_name": "longitude"}),
("X", {"axis": "X"}),
("longitude", {"units": "degrees_east"}),
],
)
def test_coordinate_criteria_in_encoding(key, metadata):
array = xr.DataArray([0, 1], dims="x", coords={"x": [0, 1]})
array.x.encoding.update(metadata)

assert_identical(array.cf[key], array.x)


def test_coordinate_attrs_take_precedence_over_encoding():
array = xr.DataArray([0, 1], dims="x", coords={"x": [0, 1]})
array.x.attrs["axis"] = "X"
array.x.encoding["axis"] = "Y"

assert_identical(array.cf["X"], array.x)
with pytest.raises(KeyError):
array.cf["Y"]


def test_standard_names_and_custom_criteria_in_encoding():
dataset = xr.Dataset({"temperature": ("time", [280.0, 281.0])})
dataset.temperature.encoding["standard_name"] = "air_temperature"

assert dataset.cf.standard_names == {"air_temperature": ["temperature"]}
assert_identical(dataset.cf["air_temperature"], dataset.temperature)

criteria = {"temperature": {"standard_name": "air_temp.*"}}
with cf_xarray.set_options(custom_criteria=criteria):
assert_identical(dataset.cf["temperature"], dataset.temperature)


@pytest.mark.parametrize("obj", objects)
def test_args_methods(obj):
with raise_if_dask_computes():
Expand Down