Skip to content
Merged
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
100 changes: 90 additions & 10 deletions xarray/tests/test_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,17 +171,25 @@ def create_store(self):
raise NotImplementedError

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
with create_tmp_file(
allow_cleanup_failure=allow_cleanup_failure) as path:
self.save(data, path, **save_kwargs)
with self.open(path, **open_kwargs) as ds:
yield ds

@contextlib.contextmanager
def roundtrip_append(self, data, save_kwargs={}, open_kwargs={},
def roundtrip_append(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
with create_tmp_file(
allow_cleanup_failure=allow_cleanup_failure) as path:
for i, key in enumerate(data.variables):
Expand Down Expand Up @@ -1194,6 +1202,17 @@ def test_read_variable_len_strings(self):
with open_dataset(tmp_file, **kwargs) as actual:
assert_identical(expected, actual)

def test_encoding_unlimited_dims(self):
ds = Dataset({'x': ('y', np.arange(10.0))})
with self.roundtrip(ds,
save_kwargs=dict(unlimited_dims=['y'])) as actual:
assert actual.encoding['unlimited_dims'] == set('y')
assert_equal(ds, actual)
ds.encoding = {'unlimited_dims': ['y']}
with self.roundtrip(ds) as actual:
assert actual.encoding['unlimited_dims'] == set('y')
assert_equal(ds, actual)


@requires_netCDF4
class TestNetCDF4Data(NetCDF4Base):
Expand Down Expand Up @@ -1276,12 +1295,17 @@ def test_autoclose_future_warning(self):
@pytest.mark.filterwarnings('ignore:deallocating CachingFileManager')
class TestNetCDF4ViaDaskData(TestNetCDF4Data):
@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if open_kwargs is None:
open_kwargs = {}
if save_kwargs is None:
save_kwargs = {}
open_kwargs.setdefault('chunks', -1)
with TestNetCDF4Data.roundtrip(
self, data, save_kwargs, open_kwargs,
allow_cleanup_failure) as ds:
yield ds.chunk()
yield ds

def test_unsorted_index_raises(self):
# Skip when using dask because dask rewrites indexers to getitem,
Expand Down Expand Up @@ -1329,15 +1353,19 @@ def open(self, store_target, **kwargs):
yield ds

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
with self.create_zarr_target() as store_target:
self.save(data, store_target, **save_kwargs)
with self.open(store_target, **open_kwargs) as ds:
yield ds

@contextlib.contextmanager
def roundtrip_append(self, data, save_kwargs={}, open_kwargs={},
def roundtrip_append(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
pytest.skip("zarr backend does not support appending")

Expand Down Expand Up @@ -1618,8 +1646,12 @@ def create_store(self):
yield backends.ScipyDataStore(fobj, 'w')

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
with create_tmp_file() as tmp_file:
with open(tmp_file, 'wb') as f:
self.save(data, f, **save_kwargs)
Expand Down Expand Up @@ -1914,6 +1946,46 @@ def test_dump_encodings_h5py(self):
assert actual.x.encoding['compression_opts'] is None


@requires_h5netcdf
@requires_dask
@pytest.mark.filterwarnings('ignore:deallocating CachingFileManager')
class TestH5NetCDFViaDaskData(TestH5NetCDFData):

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
open_kwargs.setdefault('chunks', -1)
with TestH5NetCDFData.roundtrip(
self, data, save_kwargs, open_kwargs,
allow_cleanup_failure) as ds:
yield ds

def test_dataset_caching(self):
# caching behavior differs for dask
pass

def test_write_inconsistent_chunks(self):
# Construct two variables with the same dimensions, but different
# chunk sizes.
x = da.zeros((100, 100), dtype='f4', chunks=(50, 100))
x = DataArray(data=x, dims=('lat', 'lon'), name='x')
x.encoding['chunksizes'] = (50, 100)
x.encoding['original_shape'] = (100, 100)
y = da.ones((100, 100), dtype='f4', chunks=(100, 50))
y = DataArray(data=y, dims=('lat', 'lon'), name='y')
y.encoding['chunksizes'] = (100, 50)
y.encoding['original_shape'] = (100, 100)
# Put them both into the same dataset
ds = Dataset({'x': x, 'y': y})
with self.roundtrip(ds) as actual:
assert actual['x'].encoding['chunksizes'] == (50, 100)
assert actual['y'].encoding['chunksizes'] == (100, 50)


@pytest.fixture(params=['scipy', 'netcdf4', 'h5netcdf', 'pynio'])
def readengine(request):
return request.param
Expand Down Expand Up @@ -2098,7 +2170,7 @@ def create_store(self):
yield Dataset()

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
yield data.chunk()

Expand Down Expand Up @@ -2591,8 +2663,12 @@ def open(self, path, **kwargs):
return open_dataset(path, engine='pseudonetcdf', **kwargs)

@contextlib.contextmanager
def roundtrip(self, data, save_kwargs={}, open_kwargs={},
def roundtrip(self, data, save_kwargs=None, open_kwargs=None,
allow_cleanup_failure=False):
if save_kwargs is None:
save_kwargs = {}
if open_kwargs is None:
open_kwargs = {}
with create_tmp_file(
allow_cleanup_failure=allow_cleanup_failure) as path:
self.save(data, path, **save_kwargs)
Expand Down Expand Up @@ -2799,10 +2875,14 @@ def create_tmp_geotiff(nx=4, ny=3, nz=3,
transform_args=[5000, 80000, 1000, 2000.],
crs={'units': 'm', 'no_defs': True, 'ellps': 'WGS84',
'proj': 'utm', 'zone': 18},
open_kwargs={}):
open_kwargs=None):
# yields a temporary geotiff file and a corresponding expected DataArray
import rasterio
from rasterio.transform import from_origin

if open_kwargs is None:
open_kwargs = {}

with create_tmp_file(suffix='.tif',
allow_cleanup_failure=ON_WINDOWS) as tmp_file:
# allow 2d or 3d shapes
Expand Down