diff --git a/cf_xarray/accessor.py b/cf_xarray/accessor.py index ed0c216d..b3d57be3 100644 --- a/cf_xarray/accessor.py +++ b/cf_xarray/accessor.py @@ -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 ( @@ -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 @@ -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) @@ -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()} diff --git a/cf_xarray/tests/test_accessor.py b/cf_xarray/tests/test_accessor.py index 6657be46..24512172 100644 --- a/cf_xarray/tests/test_accessor.py +++ b/cf_xarray/tests/test_accessor.py @@ -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():