From 3157d398f89ea03b6012ccbd91e82c772abc1c22 Mon Sep 17 00:00:00 2001 From: McDougall Date: Fri, 21 Aug 2026 09:26:21 +0100 Subject: [PATCH 1/4] Adds open_dataarray to BackendEntrypoint. #10562 This adds the capacity to implement open_dataarray in BackendEntrypoint. This is done in a backwards compatable way: if it is not implemented open_dataset is used instead (which is the current behaviour.) The documentation on open_dataset has been updated. The how-to-add-new-backend.md has been updated. --- doc/internals/how-to-add-new-backend.md | 25 +++ xarray/backends/api.py | 200 ++++++++++++++++-- xarray/backends/common.py | 14 +- .../tests/test_backed_entrypoint_calling.py | 101 +++++++++ 4 files changed, 320 insertions(+), 20 deletions(-) create mode 100644 xarray/tests/test_backed_entrypoint_calling.py diff --git a/doc/internals/how-to-add-new-backend.md b/doc/internals/how-to-add-new-backend.md index 3bc6fcbe0b3..646d02fcf33 100644 --- a/doc/internals/how-to-add-new-backend.md +++ b/doc/internals/how-to-add-new-backend.md @@ -35,6 +35,8 @@ it should implement the following attributes and methods: - the `guess_can_open` method (optional) - the `description` attribute (optional) - the `url` attribute (optional). +- the `open_dataarray` method (optional) +- the `open_datatree` method (optional) This is what a `BackendEntrypoint` subclass should look like: @@ -144,6 +146,29 @@ If you don't want to support the lazy loading, then the {py:class}`~xarray.Dataset` shall contain values as a {py:class}`numpy.ndarray` and your work is almost done. +(rst-open-dataarray)= + +### open_dataarray + +The backend `open_dataarray` may shall reading from file, the variables +decoding and it shall instantiate the output Xarray class {py:class}`~xarray.DataArray`. + +If `MyBackendEntrypoint.open_dataarray` is not implemented and `xarray.open_dataarray(engine='my_engine')` is called then `MyBackendEntrypoint.open_dataset` is used instead. +If `open_dataset` is used to open a `DataArray`, if the `Dataset` contains a single variable, that is returned. If the `Dataset` contains multiple variables then a `ValueError` is raised. + +All other processing and requirements are the same as for {ref}`rst-open_dataset`. + +(rst-open-datatree)= + +### open_datatree + +The backend `open_datatree` may shall reading from file, the variables +decoding and it shall instantiate the output Xarray class {py:class}`~xarray.DataTree`. + +If `MyBackendEntrypoint.open_datatree` is not implemented and `xarray.open_datatree(engine='my_engine')` is called a `NotImplementedError` is raised. + +All other processing and requirements are the same as for {ref}`rst-open_dataset`. + (rst-open-dataset-parameters)= ### open_dataset_parameters diff --git a/xarray/backends/api.py b/xarray/backends/api.py index 4330eb21cc6..62945971dc3 100644 --- a/xarray/backends/api.py +++ b/xarray/backends/api.py @@ -220,6 +220,67 @@ def load_datatree(filename_or_obj: T_PathFileOrDataStore, **kwargs) -> DataTree: return dt.load() +def _chunk_da( + backend_da, + filename_or_obj, + engine, + chunks, + overwrite_encoded_chunks, + inline_array, + chunked_array_type, + from_array_kwargs, + name=None, + chunkmanager=None, + token=(None,), + name_prefix=None, + **extra_tokens, +): + + if chunkmanager is None: + chunkmanager = guess_chunkmanager(chunked_array_type) + + # TODO refactor to move this dask-specific logic inside the DaskManager class + is_dask_chunkmanager = isinstance(chunkmanager, DaskManager) or any( + name == "dask" and manager is chunkmanager + for name, manager in list_chunkmanagers().items() + ) + if is_dask_chunkmanager: + from dask.base import tokenize + + mtime = _get_mtime(filename_or_obj) + token = tokenize(filename_or_obj, mtime, engine, chunks, **extra_tokens) + name_prefix = "open_dataset-" + else: + # not used + token = (None,) + name_prefix = None + + if backend_da._in_memory: + return backend_da + var_chunks = _get_chunk( + backend_da._data, + chunks, + chunkmanager, + preferred_chunks=backend_da.encoding.get("preferred_chunks", {}), + dims=backend_da.dims, + ) + if name is None and hasattr(backend_da, "name"): + name = backend_da.name + + return _maybe_chunk( + name, + backend_da, + var_chunks, + overwrite_encoded_chunks=overwrite_encoded_chunks, + name_prefix=name_prefix, + token=token, + inline_array=inline_array, + chunked_array_type=chunkmanager, + from_array_kwargs=from_array_kwargs.copy(), + just_use_token=True, + ) + + def _chunk_ds( backend_ds, filename_or_obj, @@ -251,28 +312,22 @@ def _chunk_ds( variables = {} for name, var in backend_ds.variables.items(): - if var._in_memory: - variables[name] = var - continue - var_chunks = _get_chunk( - var._data, - chunks, - chunkmanager, - preferred_chunks=var.encoding.get("preferred_chunks", {}), - dims=var.dims, - ) - variables[name] = _maybe_chunk( - name, + variables[name] = _chunk_da( var, - var_chunks, - overwrite_encoded_chunks=overwrite_encoded_chunks, - name_prefix=name_prefix, + filename_or_obj, + engine, + chunks, + overwrite_encoded_chunks, + inline_array, + chunked_array_type, + from_array_kwargs, + name=name, + chunkmanager=chunkmanager, token=token, - inline_array=inline_array, - chunked_array_type=chunkmanager, - from_array_kwargs=from_array_kwargs.copy(), - just_use_token=True, + name_prefix=name_prefix, + **extra_tokens, ) + return backend_ds._replace(variables) @@ -285,6 +340,61 @@ def _maybe_create_default_indexes(ds): return ds.assign_coords(Coordinates(to_index)) +def _dataarray_from_backend_dataarray( + backend_da, + filename_or_obj, + engine, + chunks, + cache, + overwrite_encoded_chunks, + inline_array, + chunked_array_type, + from_array_kwargs, + create_default_indexes, + **extra_tokens, +): + if not isinstance(chunks, int | dict) and chunks not in {None, "auto"}: + raise ValueError( + f"chunks must be an int, dict, 'auto', or None. Instead found {chunks}." + ) + + # Protect data inplace + data: indexing.ExplicitlyIndexedNDArrayMixin + data = indexing.CopyOnWriteArray(backend_da._data) + if cache: + data = indexing.MemoryCachedArray(data) + backend_da.data = data + + if create_default_indexes: + da = _maybe_create_default_indexes(backend_da) + else: + da = backend_da + + if chunks is not None: + da = _chunk_da( + da, + filename_or_obj, + engine, + chunks, + overwrite_encoded_chunks, + inline_array, + chunked_array_type, + from_array_kwargs, + **extra_tokens, + ) + + da.set_close(backend_da._close) + + # Ensure source filename always stored in dataset object + if "source" not in da.encoding: + path = getattr(filename_or_obj, "path", filename_or_obj) + + if isinstance(path, str | os.PathLike): + da.encoding["source"] = _normalize_path(path) + + return da + + def _dataset_from_backend_dataset( backend_ds, filename_or_obj, @@ -823,6 +933,58 @@ class (a subclass of ``BackendEntrypoint``) can also be used. open_dataset """ + try: + if cache is None: + cache = chunks is None + + if backend_kwargs is not None: + kwargs.update(backend_kwargs) + + if engine is None: + engine = plugins.guess_engine(filename_or_obj) + + if from_array_kwargs is None: + from_array_kwargs = {} + + backend = plugins.get_backend(engine) + + decoders = _resolve_decoders_kwargs( + decode_cf, + open_backend_dataset_parameters=backend.open_dataset_parameters, + mask_and_scale=mask_and_scale, + decode_times=decode_times, + decode_timedelta=decode_timedelta, + concat_characters=concat_characters, + use_cftime=use_cftime, + decode_coords=decode_coords, + ) + + overwrite_encoded_chunks = kwargs.pop("overwrite_encoded_chunks", None) + backend_da = backend.open_dataarray( + filename_or_obj, + drop_variables=drop_variables, + **decoders, + **kwargs, + ) + da = _dataarray_from_backend_dataarray( + backend_da, + filename_or_obj, + engine, + chunks, + cache, + overwrite_encoded_chunks, + inline_array, + chunked_array_type, + from_array_kwargs, + drop_variables=drop_variables, + create_default_indexes=create_default_indexes, + **decoders, + **kwargs, + ) + return da + except NotImplementedError: + pass + dataset = open_dataset( filename_or_obj, decode_cf=decode_cf, diff --git a/xarray/backends/common.py b/xarray/backends/common.py index fd818a79961..7fb4893274b 100644 --- a/xarray/backends/common.py +++ b/xarray/backends/common.py @@ -775,6 +775,18 @@ def __repr__(self) -> str: txt += f"\n Learn more at {self.url}" return txt + def open_dataarray( + self, + filename_or_obj: T_PathFileOrDataStore, + *, + drop_variables: str | Iterable[str] | None = None, + ) -> Dataset: + """ + Backend open_dataarray method used by Xarray in :py:func:`~xarray.open_dataarray`. + """ + + raise NotImplementedError() + def open_dataset( self, filename_or_obj: T_PathFileOrDataStore, @@ -782,7 +794,7 @@ def open_dataset( drop_variables: str | Iterable[str] | None = None, ) -> Dataset: """ - Backend open_dataset method used by Xarray in :py:func:`~xarray.open_dataset`. + Backend open_dataset method used by Xarray in :py:func:`~xarray.open_dataset` and :py:func:`~xarray.open_dataarray` of open_dataarray si not implemented. """ raise NotImplementedError() diff --git a/xarray/tests/test_backed_entrypoint_calling.py b/xarray/tests/test_backed_entrypoint_calling.py new file mode 100644 index 00000000000..1e97fbc1058 --- /dev/null +++ b/xarray/tests/test_backed_entrypoint_calling.py @@ -0,0 +1,101 @@ +import pytest + +from xarray import ( + open_dataarray, + open_dataset, + open_datatree, +) +from xarray.backends import BackendEntrypoint + + +class ArgsCalled(Exception): + def __init__(self, *args, **kwargs): + self.args = args + self.kwargs = kwargs + + +class DataArrayCalled(ArgsCalled): + pass + + +class DatasetCalled(ArgsCalled): + pass + + +class DataTreeCalled(ArgsCalled): + pass + + +class DummyBackendEntrypointDataset(BackendEntrypoint): + def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DatasetCalled(*args, **kwargs) + + +class DummyBackendEntrypointDatasetDataTree(BackendEntrypoint): + def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DatasetCalled(*args, **kwargs) + + def open_datatree(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DataTreeCalled(*args, **kwargs) + + +class DummyBackendEntrypointAll(BackendEntrypoint): + def open_dataarray(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DataArrayCalled(*args, **kwargs) + + def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DatasetCalled(*args, **kwargs) + + def open_datatree(filename_or_obj, *args, **kwargs): # type: ignore[override] + raise DataTreeCalled(*args, **kwargs) + + +def test_dataset(tmp_path): + + dataset_engine = DummyBackendEntrypointDataset + + existing_file = tmp_path / "test.unknown" + existing_file.write_bytes(b"") + + try: + open_dataarray(existing_file, engine=dataset_engine) + except DatasetCalled as e: + assert e.args[0] == existing_file + assert e.kwargs == {"drop_variables": None} + + try: + open_dataset(existing_file, engine=dataset_engine) + except DatasetCalled as e: + assert e.args[0] == existing_file + assert e.kwargs == {"drop_variables": None} + + with pytest.raises(NotImplementedError): + open_datatree(existing_file, engine=dataset_engine) + + +def test_datatree(tmp_path): + + dataset_engine = DummyBackendEntrypointDatasetDataTree + + existing_file = tmp_path / "test.unknown" + existing_file.write_bytes(b"") + + try: + open_datatree(existing_file, engine=dataset_engine) + except DataTreeCalled as e: + assert e.args[0] == existing_file + assert e.kwargs == {"drop_variables": None} + + +def test_dataarray(tmp_path): + + dataset_engine = DummyBackendEntrypointAll + + existing_file = tmp_path / "test.unknown" + existing_file.write_bytes(b"") + + try: + open_dataarray(existing_file, engine=dataset_engine) + except DataArrayCalled as e: + assert e.args[0] == existing_file + assert e.kwargs == {"drop_variables": None} From 978f83f42ff835c609b66688c4a91e19bcdab9a6 Mon Sep 17 00:00:00 2001 From: McDougall Date: Fri, 21 Aug 2026 09:41:08 +0100 Subject: [PATCH 2/4] Update whats-new.rst --- doc/whats-new.rst | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/doc/whats-new.rst b/doc/whats-new.rst index d1505bfa081..666746151fb 100644 --- a/doc/whats-new.rst +++ b/doc/whats-new.rst @@ -25,6 +25,11 @@ New Features silently being written uncompressed (:issue:`10657`, :pull:`11067`). By `Mark Harfouche `_. +- ``xarray.BackendEntrypoint`` now supports implementing ``open_dataarray``. + Previously ``open_dataset`` was used when ``xarray.open_dataarray(file, engine='my-engine')`` was called. + Now, if ``BackendEntrypoint.open_dataarray`` is implemented, it will be used. (:issue:`10562`, :pull:`11537`). + By `Duncan McDougall `_. + Breaking Changes ~~~~~~~~~~~~~~~~ From 3cda0c293ece65d6c609b34b9b499e66573f76aa Mon Sep 17 00:00:00 2001 From: McDougall Date: Fri, 21 Aug 2026 10:47:51 +0100 Subject: [PATCH 3/4] Fixes for tests. --- xarray/backends/common.py | 3 ++- ...calling.py => test_backend_entrypoint_calling.py} | 12 ++++++------ 2 files changed, 8 insertions(+), 7 deletions(-) rename xarray/tests/{test_backed_entrypoint_calling.py => test_backend_entrypoint_calling.py} (81%) diff --git a/xarray/backends/common.py b/xarray/backends/common.py index 7fb4893274b..746197cf32a 100644 --- a/xarray/backends/common.py +++ b/xarray/backends/common.py @@ -39,6 +39,7 @@ from xarray.namedarray.utils import is_duck_dask_array if TYPE_CHECKING: + from xarray.core.dataarray import DataArray from xarray.core.dataset import Dataset from xarray.core.types import NestedSequence @@ -780,7 +781,7 @@ def open_dataarray( filename_or_obj: T_PathFileOrDataStore, *, drop_variables: str | Iterable[str] | None = None, - ) -> Dataset: + ) -> DataArray: """ Backend open_dataarray method used by Xarray in :py:func:`~xarray.open_dataarray`. """ diff --git a/xarray/tests/test_backed_entrypoint_calling.py b/xarray/tests/test_backend_entrypoint_calling.py similarity index 81% rename from xarray/tests/test_backed_entrypoint_calling.py rename to xarray/tests/test_backend_entrypoint_calling.py index 1e97fbc1058..1fb0e96e4a6 100644 --- a/xarray/tests/test_backed_entrypoint_calling.py +++ b/xarray/tests/test_backend_entrypoint_calling.py @@ -27,26 +27,26 @@ class DataTreeCalled(ArgsCalled): class DummyBackendEntrypointDataset(BackendEntrypoint): - def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_dataset(filename_or_obj, *args, **kwargs): raise DatasetCalled(*args, **kwargs) class DummyBackendEntrypointDatasetDataTree(BackendEntrypoint): - def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_dataset(filename_or_obj, *args, **kwargs): raise DatasetCalled(*args, **kwargs) - def open_datatree(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_datatree(filename_or_obj, *args, **kwargs): raise DataTreeCalled(*args, **kwargs) class DummyBackendEntrypointAll(BackendEntrypoint): - def open_dataarray(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_dataarray(filename_or_obj, *args, **kwargs): raise DataArrayCalled(*args, **kwargs) - def open_dataset(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_dataset(filename_or_obj, *args, **kwargs): raise DatasetCalled(*args, **kwargs) - def open_datatree(filename_or_obj, *args, **kwargs): # type: ignore[override] + def open_datatree(filename_or_obj, *args, **kwargs): raise DataTreeCalled(*args, **kwargs) From 4ca124bc173e58796d3fbfa88d396879b74853ea Mon Sep 17 00:00:00 2001 From: McDougall Date: Fri, 21 Aug 2026 11:33:30 +0100 Subject: [PATCH 4/4] Moves inplace protection to typed function, to correct tests. --- xarray/backends/api.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/xarray/backends/api.py b/xarray/backends/api.py index 62945971dc3..19a5ec30e25 100644 --- a/xarray/backends/api.py +++ b/xarray/backends/api.py @@ -109,6 +109,14 @@ def _get_mtime(filename_or_obj): return mtime +def _protect_dataarray_variables_inplace(dataarray: DataArray, cache: bool) -> None: + data: indexing.ExplicitlyIndexedNDArrayMixin + data = indexing.CopyOnWriteArray(dataarray._data) + if cache: + data = indexing.MemoryCachedArray(data) + dataarray.data = data + + def _protect_dataset_variables_inplace(dataset: Dataset, cache: bool) -> None: for name, variable in dataset.variables.items(): if name not in dataset._indexes: @@ -358,12 +366,7 @@ def _dataarray_from_backend_dataarray( f"chunks must be an int, dict, 'auto', or None. Instead found {chunks}." ) - # Protect data inplace - data: indexing.ExplicitlyIndexedNDArrayMixin - data = indexing.CopyOnWriteArray(backend_da._data) - if cache: - data = indexing.MemoryCachedArray(data) - backend_da.data = data + _protect_dataarray_variables_inplace(backend_da, cache) if create_default_indexes: da = _maybe_create_default_indexes(backend_da)