diff --git a/src/xarray_grass/__init__.py b/src/xarray_grass/__init__.py index 7a94871..538cc85 100644 --- a/src/xarray_grass/__init__.py +++ b/src/xarray_grass/__init__.py @@ -1,10 +1,10 @@ -from xarray_grass.grass_interface import GrassConfig as GrassConfig -from xarray_grass.grass_interface import GrassInterface as GrassInterface -from xarray_grass.xarray_grass import GrassBackendEntrypoint as GrassBackendEntrypoint +from xarray_grass.coord_utils import RegionData as RegionData from xarray_grass.grass_backend_array import ( GrassSTDSBackendArray as GrassSTDSBackendArray, ) +from xarray_grass.grass_interface import GrassConfig as GrassConfig +from xarray_grass.grass_interface import GrassInterface as GrassInterface from xarray_grass.to_grass import to_grass as to_grass -from xarray_grass.coord_utils import RegionData as RegionData +from xarray_grass.xarray_grass import GrassBackendEntrypoint as GrassBackendEntrypoint __version__ = "0.4.0" diff --git a/src/xarray_grass/coord_utils.py b/src/xarray_grass/coord_utils.py index 1a47912..83f64d7 100644 --- a/src/xarray_grass/coord_utils.py +++ b/src/xarray_grass/coord_utils.py @@ -13,12 +13,12 @@ GNU General Public License for more details. """ -from collections import namedtuple -from typing import Mapping +from collections.abc import Mapping +from typing import NamedTuple + import numpy as np import xarray as xr # For type hinting xr.DataArray - region_type_dict = { "projection": str, "zone": str, @@ -41,14 +41,34 @@ "cells": int, "cells3": int, } -RegionData = namedtuple( - "RegionData", - region_type_dict.keys(), - defaults=[None for _ in region_type_dict.keys()], -) -def get_region_from_xarray(data_array: xr.DataArray, dims: Mapping[str, str]) -> dict: +class RegionData(NamedTuple): + projection: str | None = None + zone: str | None = None + n: float | None = None + s: float | None = None + w: float | None = None + e: float | None = None + t: float | None = None + b: float | None = None + nsres: float | None = None + nsres3: float | None = None + ewres: float | None = None + ewres3: float | None = None + tbres: float | None = None + rows: int | None = None + rows3: int | None = None + cols: int | None = None + cols3: int | None = None + depths: int | None = None + cells: int | None = None + cells3: int | None = None + + +def get_region_from_xarray( + data_array: xr.DataArray, dims: Mapping[str, str] +) -> RegionData: """ Calculates GRASS GIS region parameters from an xarray DataArray. @@ -73,8 +93,8 @@ def get_region_from_xarray(data_array: xr.DataArray, dims: Mapping[str, str]) -> Returns ------- - dict - A dictionary containing GRASS region parameters: + RegionData + GRASS region parameters: 'n', 's', 'e', 'w': float or None (geographical limits) 't', 'b': float or None (top, bottom limits for 3D) 'nsres', 'ewres': float or None (2D resolutions) @@ -84,7 +104,7 @@ def get_region_from_xarray(data_array: xr.DataArray, dims: Mapping[str, str]) -> """ region = {} - def _calculate_res(coords_arr_np: np.ndarray) -> float | None: + def _calculate_res(coords_arr_np: np.ndarray | None) -> float | None: if coords_arr_np is not None and len(coords_arr_np) >= 2: # Ensure consistent dtype for subtraction, then convert to float res = np.abs( @@ -95,7 +115,7 @@ def _calculate_res(coords_arr_np: np.ndarray) -> float | None: # Determine if it's 3D based on presence of z-coordinate name in dims and data_array z_name = dims.get("z") - is_3d = z_name and z_name in data_array.coords + is_3d = z_name is not None and z_name in data_array.coords x_coords_np, y_coords_np, z_coords_np = None, None, None @@ -177,4 +197,16 @@ def _calculate_res(coords_arr_np: np.ndarray) -> float | None: region["b"] = float(z_coords_np[0] - region["tbres"] / 2) region["t"] = float(z_coords_np[-1] + region["tbres"] / 2) - return RegionData(**region) + return RegionData( + n=region.get("n"), + s=region.get("s"), + w=region.get("w"), + e=region.get("e"), + t=region.get("t"), + b=region.get("b"), + nsres=region.get("nsres"), + nsres3=region.get("nsres3"), + ewres=region.get("ewres"), + ewres3=region.get("ewres3"), + tbres=region.get("tbres"), + ) diff --git a/src/xarray_grass/grass_backend_array.py b/src/xarray_grass/grass_backend_array.py index 27b25e2..ca5172e 100644 --- a/src/xarray_grass/grass_backend_array.py +++ b/src/xarray_grass/grass_backend_array.py @@ -13,16 +13,16 @@ """ from __future__ import annotations -from typing import TYPE_CHECKING + import threading +from typing import TYPE_CHECKING, Any, Literal import numpy as np -import xarray as xr - from xarray.backends import BackendArray +from xarray.core import indexing if TYPE_CHECKING: - from xarray_grass.grass_interface import GrassInterface + from xarray_grass.grass_interface import GrassInterface, MapData class GrassSTDSBackendArray(BackendArray): @@ -30,30 +30,30 @@ class GrassSTDSBackendArray(BackendArray): def __init__( self, - shape, - dtype, - map_list: list, # List of map metadata objects - map_type: str, + shape: tuple[int, ...], + dtype: np.dtype[Any], + map_list: list[MapData], + map_type: Literal["raster", "raster3d"], grass_interface: GrassInterface, - ): + ) -> None: self.shape = shape self.dtype = dtype self._lock = threading.Lock() - self.map_list = map_list # List with .id attribute - self.map_type = map_type # "raster" or "raster3d" + self.map_list = map_list + self.map_type = map_type self.grass_interface = grass_interface - self._cached_maps = {} # Cache loaded maps by index + self._cached_maps: dict[int, np.ndarray] = {} - def __getitem__(self, key: xr.core.indexing.ExplicitIndexer) -> np.typing.ArrayLike: + def __getitem__(self, key: indexing.ExplicitIndexer) -> np.ndarray: """takes in input an index and returns a NumPy array""" - return xr.core.indexing.explicit_indexing_adapter( + return indexing.explicit_indexing_adapter( key, self.shape, - xr.core.indexing.IndexingSupport.BASIC, + indexing.IndexingSupport.BASIC, self._raw_indexing_method, ) - def _raw_indexing_method(self, key: tuple): + def _raw_indexing_method(self, key: tuple[Any, ...]) -> np.ndarray: """Load only the maps needed for the requested slice""" with self._lock: # key is a tuple of slices/indices for each dimension @@ -70,7 +70,7 @@ def _raw_indexing_method(self, key: tuple): time_indices = list(time_key) # Load only the needed maps - result_list = [] + result_list: list[np.ndarray] = [] for t_idx in time_indices: if t_idx not in self._cached_maps: map_data = self.map_list[t_idx] diff --git a/src/xarray_grass/grass_interface.py b/src/xarray_grass/grass_interface.py index 5d5c7b0..7174c8b 100644 --- a/src/xarray_grass/grass_interface.py +++ b/src/xarray_grass/grass_interface.py @@ -12,29 +12,29 @@ GNU General Public License for more details. """ -import os import math -from collections import namedtuple +import os from dataclasses import dataclass from datetime import datetime, timedelta from pathlib import Path -from typing import Self, Optional - -import numpy as np -import pandas as pd +from typing import Any, Literal, NamedTuple, Self +import grass.pygrass.utils as gutils # ty: ignore[unresolved-import] # Needed to import grass modules -import grass.script as gs -from grass.script import array as garray -import grass.pygrass.utils as gutils -from grass.pygrass import raster as graster -from grass.pygrass.raster.abstract import Info, RasterAbstractBase -import grass.temporal as tgis +import grass.script as gs # ty: ignore[unresolved-import] +import grass.temporal as tgis # ty: ignore[unresolved-import] +import numpy as np +from grass.pygrass import raster as graster # ty: ignore[unresolved-import] +from grass.pygrass.raster.abstract import ( # ty: ignore[unresolved-import] + Info, + RasterAbstractBase, +) +from grass.script import array as garray # ty: ignore[unresolved-import] from xarray_grass.coord_utils import ( - region_type_dict, RegionData, + region_type_dict, ) gs.core.set_raise_on_error(True) @@ -45,11 +45,18 @@ class GrassConfig: gisdb: str | Path project: str | Path mapset: str | Path - grassbin: str | Path + grassbin: str | Path | None strds_cols = ["id", "start_time", "end_time"] -MapData = namedtuple("MapData", strds_cols + ["dtype"]) + + +class MapData(NamedTuple): + id: str + start_time: datetime | int + end_time: datetime | int | None + dtype: np.dtype[Any] + strds_infos = [ "id", @@ -66,7 +73,22 @@ class GrassConfig: "top", "bottom", ] -STRDSInfos = namedtuple("STRDSInfos", strds_infos) + + +class STRDSInfos(NamedTuple): + id: str + title: str + temporal_type: str + time_unit: str | None + start_time: datetime | int | None + end_time: datetime | int | None + time_granularity: str | int + north: float + south: float + east: float + west: float + top: float + bottom: float class GrassInterface(object): @@ -104,7 +126,7 @@ def __init__(self, overwrite: bool = False): tgis.init() @staticmethod - def get_gisenv() -> dict[str]: + def get_gisenv() -> dict[str, str]: """Return the current GRASS environment.""" return gs.gisenv() @@ -114,29 +136,53 @@ def get_accessible_mapsets() -> list[str]: return gs.parse_command("g.mapsets", flags="p", format="json")["mapsets"] @staticmethod - def get_region() -> namedtuple: + def get_region() -> RegionData: """Return the current GRASS region.""" region_raw = gs.parse_command("g.region", flags="g3") - region = { - k: region_type_dict[k](v) for k, v in region_raw.items() if v is not None - } - region = RegionData(**region) - return region + def value(key: str) -> Any: + raw_value = region_raw.get(key) + if raw_value is None: + return None + return region_type_dict[key](raw_value) + + return RegionData( + projection=value("projection"), + zone=value("zone"), + n=value("n"), + s=value("s"), + w=value("w"), + e=value("e"), + t=value("t"), + b=value("b"), + nsres=value("nsres"), + nsres3=value("nsres3"), + ewres=value("ewres"), + ewres3=value("ewres3"), + tbres=value("tbres"), + rows=value("rows"), + rows3=value("rows3"), + cols=value("cols"), + cols3=value("cols3"), + depths=value("depths"), + cells=value("cells"), + cells3=value("cells3"), + ) @staticmethod def set_region(region_data: RegionData) -> None: # 2D region if region_data.tbres is None: - if not all( - [ + if any( + value is None + for value in ( region_data.n, region_data.s, region_data.e, region_data.w, region_data.nsres, region_data.ewres, - ] + ) ): raise ValueError( "n, s, e, w, nsres and ewres must be set for 2D regions." @@ -153,6 +199,24 @@ def set_region(region_data: RegionData) -> None: ) # 3D region else: + if any( + value is None + for value in ( + region_data.n, + region_data.s, + region_data.e, + region_data.w, + region_data.t, + region_data.b, + region_data.nsres3, + region_data.ewres3, + ) + ): + raise ValueError( + "n, s, e, w, t, b, nsres3 and ewres3 must be set for 3D regions." + ) + assert region_data.nsres3 is not None + assert region_data.ewres3 is not None # TODO: remove when grass 8.5 is released tolerance = 1e-9 if not math.isclose( @@ -177,10 +241,10 @@ def set_region(region_data: RegionData) -> None: ) @staticmethod - def is_latlon(): - return gs.locn_is_latlong() + def is_latlon() -> bool: + return bool(gs.locn_is_latlong()) - def is_xy(self): + def is_xy(self) -> bool: """return True if the location is neither projected or latlon""" proj_code = gs.parse_command("g.region", flags="pug")["projection"] if int(proj_code) == 0: @@ -188,8 +252,8 @@ def is_xy(self): else: return False - def get_spatial_units(self): - if self.is_xy: + def get_spatial_units(self) -> str | None: + if self.is_xy(): return None else: return gs.parse_command("g.proj", flags="g")["units"] @@ -264,7 +328,7 @@ def grass_dtype(self, dtype: str) -> str: return mtype @staticmethod - def numpy_dtype(mtype: str) -> np.dtype: + def numpy_dtype(mtype: str) -> np.dtype[Any]: if mtype == "CELL": dtype = np.dtype("int64") elif mtype == "FCELL": @@ -281,29 +345,29 @@ def has_mask() -> bool: return bool(gs.read_command("g.list", type="raster", pattern="MASK")) @staticmethod - def list_strds(mapset: str = None) -> list[str]: + def list_strds(mapset: str | None = None) -> list[str]: if mapset: return tgis.tlist_grouped("strds")[mapset] else: return tgis.tlist("strds") @staticmethod - def list_str3ds(mapset: str = None) -> list[str]: + def list_str3ds(mapset: str | None = None) -> list[str]: if mapset: return tgis.tlist_grouped("str3ds")[mapset] else: return tgis.tlist("str3ds") @staticmethod - def list_raster(mapset: str = None) -> list[str]: + def list_raster(mapset: str | None = None) -> list[str]: """List raster maps in the given mapset""" return gs.list_strings("raster", mapset=mapset) @staticmethod - def list_raster3d(mapset: str = None) -> list[str]: + def list_raster3d(mapset: str | None = None) -> list[str]: return gs.list_strings("raster_3d", mapset=mapset) - def list_grass_objects(self, mapset: str = None) -> dict[list[str]]: + def list_grass_objects(self, mapset: str | None = None) -> dict[str, list[str]]: """Return all GRASS objects in a given mapset.""" objects_dict = {} objects_dict["raster"] = self.list_raster(mapset) @@ -313,18 +377,20 @@ def list_grass_objects(self, mapset: str = None) -> dict[list[str]]: return objects_dict @staticmethod - def get_raster_info(raster_id: str) -> Info: + def get_raster_info(raster_id: str) -> dict[str, Any]: result = gs.parse_command("r.info", map=raster_id, flags="ge") # Strip quotes from string values (r.info returns quoted strings) return {k: v.strip('"') if isinstance(v, str) else v for k, v in result.items()} @staticmethod - def get_raster3d_info(raster3d_id): + def get_raster3d_info(raster3d_id: str) -> dict[str, Any]: result = gs.parse_command("r3.info", map=raster3d_id, flags="gh") # Strip quotes from string values (r3.info -gh returns quoted strings) return {k: v.strip('"') if isinstance(v, str) else v for k, v in result.items()} - def get_stds_infos(self, strds_name, stds_type) -> STRDSInfos: + def get_stds_infos( + self, strds_name: str, stds_type: Literal["strds", "str3ds"] + ) -> STRDSInfos: strds_id = self.get_id_from_name(strds_name) if stds_type not in ["strds", "str3ds"]: raise ValueError( @@ -443,7 +509,7 @@ def register_maps_in_stds( semantic: str, t_type: str, stds_type: str, - time_unit: Optional[str] = None, + time_unit: str | None = None, ) -> Self: """Create a STDS, create one mapdataset for each map and register them in the temporal database. @@ -479,8 +545,17 @@ def register_maps_in_stds( raise TypeError("relative time requires a timedelta object.") if not time_unit: raise TypeError("relative time requires a time_unit.") - # Convert timedelta to numeric value in the specified unit - rel_time = map_time / pd.Timedelta(1, unit=time_unit) + unit_deltas = { + "days": timedelta(days=1), + "hours": timedelta(hours=1), + "minutes": timedelta(minutes=1), + "seconds": timedelta(seconds=1), + } + try: + unit_delta = unit_deltas[time_unit] + except KeyError: + raise ValueError(f"Unsupported relative time unit: {time_unit}") + rel_time = map_time / unit_delta map_dts.set_relative_time(rel_time, None, time_unit) elif t_type == "absolute": if not isinstance(map_time, datetime): @@ -511,7 +586,7 @@ def register_maps_in_stds( ) return self - def get_coordinates(self, raster_3d: bool) -> dict[str : np.ndarray]: + def get_coordinates(self, raster_3d: bool) -> dict[str, np.ndarray]: """return np.ndarray of coordinates from the GRASS region.""" current_region = self.get_region() lim_e = current_region.e @@ -527,6 +602,18 @@ def get_coordinates(self, raster_3d: bool) -> dict[str : np.ndarray]: else: dx = current_region.ewres dy = current_region.nsres + if ( + lim_e is None + or lim_w is None + or lim_n is None + or lim_s is None + or lim_t is None + or lim_b is None + or dx is None + or dy is None + or dz is None + ): + raise ValueError("The current GRASS region is missing coordinate metadata.") # GRASS limits are at the edge of the region. # In the exported arrays, coordinates are at the center of the cell # Stop not changed to include it in the range diff --git a/src/xarray_grass/to_grass.py b/src/xarray_grass/to_grass.py index 60004f5..6bcd963 100644 --- a/src/xarray_grass/to_grass.py +++ b/src/xarray_grass/to_grass.py @@ -14,13 +14,16 @@ """ from __future__ import annotations + import os -from typing import TYPE_CHECKING, Mapping, Optional +from collections.abc import Mapping +from datetime import datetime, timedelta +from typing import TYPE_CHECKING -from pyproj import CRS -import xarray as xr import numpy as np import pandas as pd +import xarray as xr +from pyproj import CRS from xarray_grass.coord_utils import get_region_from_xarray @@ -28,9 +31,15 @@ from xarray_grass.grass_interface import GrassInterface +def _validate_name(name: object) -> str: + if not isinstance(name, str) or not name: + raise ValueError("GRASS object names must be non-empty strings.") + return name + + def to_grass( dataset: xr.Dataset | xr.DataArray, - dims: Optional[Mapping[str, Mapping[str, str]]] = None, + dims: Mapping[str, Mapping[str, str]] | None = None, overwrite: bool = False, ) -> None: """Convert an xarray.Dataset or xarray.DataArray to GRASS GIS maps. @@ -69,9 +78,11 @@ def to_grass( from xarray_grass.grass_interface import GrassInterface if isinstance(dataset, xr.Dataset): - input_var_names = [var_name for var_name, _ in dataset.data_vars.items()] + input_var_names: list[str] = [ + _validate_name(var_name) for var_name in dataset.data_vars + ] elif isinstance(dataset, xr.DataArray): - input_var_names = [dataset.name] + input_var_names = [_validate_name(dataset.name)] else: raise TypeError( f"'dataset must be either an Xarray DataArray or Dataset, not {type(dataset)}" @@ -91,7 +102,7 @@ class DimensionsFormatter: """Populate the dimension mapping based on default values and user-provided ones""" # Default dimension names - default_dims = { + default_dims: dict[str, str] = { "start_time": "start_time", "end_time": "end_time", "x": "x", @@ -101,10 +112,14 @@ class DimensionsFormatter: "z": "z", } - def __init__(self, input_var_names, input_dims): + def __init__( + self, + input_var_names: list[str], + input_dims: Mapping[str, Mapping[str, str]] | None, + ) -> None: self.input_var_names = input_var_names self.input_dims = input_dims - self._dataset_dims = {} + self._dataset_dims: dict[str, dict[str, str]] = {} # Instantiate the dimensions with default values for var_name in input_var_names: @@ -115,7 +130,7 @@ def __init__(self, input_var_names, input_dims): if self.input_dims is not None: self.check_input_dims() - def check_input_dims(self): + def check_input_dims(self) -> None: """Check conformity of provided dims Mapping""" if not isinstance(self.input_dims, Mapping): raise TypeError( @@ -132,13 +147,15 @@ def check_input_dims(self): f"Variables found: {self.input_var_names}" ) - def fill_dims(self): + def fill_dims(self) -> None: """Replace the default values with those given by the user.""" + if self.input_dims is None: + return for var_name, dims in self.input_dims.items(): for dim_key, dim_value in dims.items(): self._dataset_dims[var_name][dim_key] = dim_value - def get_formatted_dims(self): + def get_formatted_dims(self) -> dict[str, dict[str, str]]: if self.input_dims is not None: self.fill_dims() return self._dataset_dims @@ -149,7 +166,7 @@ def __init__( self, dataset: xr.Dataset | xr.DataArray, grass_interface: GrassInterface, - dims: Mapping[str, str] = None, + dims: Mapping[str, Mapping[str, str]], ): self.dataset = dataset self.grass_interface = grass_interface @@ -168,11 +185,13 @@ def to_grass(self) -> None: f"CRS mismatch: GRASS project CRS is {grass_crs}, " f"but dataset CRS is {dataset_crs}." ) - try: + if isinstance(self.dataset, xr.Dataset): for var_name, data in self.dataset.data_vars.items(): - self._datarray_to_grass(data, self.dataset_dims[var_name]) - except AttributeError: # DataArray - self._datarray_to_grass(self.dataset, self.dataset_dims[self.dataset.name]) + name = _validate_name(var_name) + self._datarray_to_grass(data, self.dataset_dims[name]) + else: + name = _validate_name(self.dataset.name) + self._datarray_to_grass(self.dataset, self.dataset_dims[name]) def _datarray_to_grass( self, @@ -180,6 +199,7 @@ def _datarray_to_grass( dims: Mapping[str, str], ) -> None: """Convert an xarray DataArray to GRASS maps.""" + data_name = _validate_name(data.name) if len(data.dims) > 4 or len(data.dims) < 2: raise ValueError( f"Only DataArray with 2 to 4 dimensions are supported. " @@ -216,12 +236,12 @@ def _datarray_to_grass( try: if is_raster: data = self.transpose(data, dims, arr_type="raster") - self.grass_interface.write_raster_map(data, data.name) + self.grass_interface.write_raster_map(data.values, data_name) elif is_strds: self._write_stds(data, dims) elif is_raster_3d: data = self.transpose(data, dims, arr_type="raster3d") - self.grass_interface.write_raster3d_map(data, data.name) + self.grass_interface.write_raster3d_map(data.values, data_name) elif is_str3ds: self._write_stds(data, dims) else: @@ -234,7 +254,10 @@ def _datarray_to_grass( self.grass_interface.set_region(current_region) def transpose( - self, da: xr.DataArray, dims, arr_type: str = "raster" + self, + da: xr.DataArray, + dims: Mapping[str, str], + arr_type: str = "raster", ) -> xr.DataArray: """Force dimension order to conform with grass expectation.""" if "raster" == arr_type: @@ -246,7 +269,7 @@ def transpose( f"Unknown array type: {arr_type}. Must be 'raster' or 'raster3d'." ) - def _write_stds(self, data: xr.DataArray, dims: Mapping): + def _write_stds(self, data: xr.DataArray, dims: Mapping[str, str]) -> None: # 1. Determine the temporal coordinate and type time_coord = data[dims["start_time"]] time_dtype = time_coord.dtype @@ -255,15 +278,16 @@ def _write_stds(self, data: xr.DataArray, dims: Mapping): temporal_type = "absolute" elif np.issubdtype(time_dtype, np.integer): temporal_type = "relative" - time_unit = time_coord.attrs.get("units", None) - if not time_unit: + raw_time_unit = time_coord.attrs.get("units") + if not isinstance(raw_time_unit, str) or not raw_time_unit: raise ValueError( f"Relative time coordinate '{dims['start_time']}' in DataArray '{data.name}' " "requires a 'units' attribute. " "Accepted values: 'days', 'hours', 'minutes', 'seconds'." ) + time_unit = raw_time_unit # Validate that the unit is supported by both pandas and GRASS - supported_units = ["days", "hours", "minutes", "seconds"] + supported_units: list[str] = ["days", "hours", "minutes", "seconds"] if time_unit not in supported_units: raise ValueError( f"Unsupported time unit '{time_unit}' for relative time in DataArray '{data.name}'. " @@ -284,27 +308,28 @@ def _write_stds(self, data: xr.DataArray, dims: Mapping): arr_type = "raster3d" # Check if exists + data_name = _validate_name(data.name) if "strds" == stds_type: if ( not self.grass_interface.overwrite - and self.grass_interface.name_is_strds(data.name) + and self.grass_interface.name_is_strds(data_name) ): raise RuntimeError( - f"STRDS {data.name} already exists and will not be overwritten." + f"STRDS {data_name} already exists and will not be overwritten." ) elif "str3ds" == stds_type: if ( not self.grass_interface.overwrite - and self.grass_interface.name_is_str3ds(data.name) + and self.grass_interface.name_is_str3ds(data_name) ): raise RuntimeError( - f"STR3DS {data.name} already exists and will not be overwritten." + f"STR3DS {data_name} already exists and will not be overwritten." ) else: raise ValueError(f"Unknown STDS type '{stds_type}'.") # 3. Loop through the time dim: - map_list = [] + map_list: list[tuple[str, datetime | timedelta]] = [] for index, time in enumerate(time_coord): darray = data.sel({dims["start_time"]: time}) darray = self.transpose(darray, dims, arr_type=arr_type) @@ -318,7 +343,7 @@ def _write_stds(self, data: xr.DataArray, dims: Mapping): nd_array[np.isinf(nd_array)] = np.nan # 3.1 Write each map individually - raster_name = f"{data.name}_{temporal_type}_{index}" + raster_name = f"{data_name}_{temporal_type}_{index}" if not is_3d: self.grass_interface.write_raster_map( arr=nd_array, rast_name=raster_name @@ -329,16 +354,25 @@ def _write_stds(self, data: xr.DataArray, dims: Mapping): ) # 3.2 populate an iterable[tuple[str, datetime | timedelta]] time_value = time.values.item() + if pd.isna(time_value): + raise ValueError( + f"Temporal coordinate '{dims['start_time']}' contains a missing value." + ) if temporal_type == "absolute": absolute_time = pd.Timestamp(time_value) - map_list.append((raster_name, absolute_time.to_pydatetime())) + python_time = absolute_time.to_pydatetime() + if not isinstance(python_time, datetime): + raise ValueError("Absolute timestamps cannot contain NaT values.") + map_list.append((raster_name, python_time)) else: relative_time = pd.Timedelta(time_value, unit=time_unit) + if not isinstance(relative_time, pd.Timedelta): + raise ValueError("Relative timestamps cannot contain NaT values.") map_list.append((raster_name, relative_time.to_pytimedelta())) # 4. Create STDS and register the maps in it self.grass_interface.register_maps_in_stds( stds_title="", - stds_name=data.name, + stds_name=data_name, stds_desc="", map_list=map_list, semantic=semantic_type, diff --git a/src/xarray_grass/xarray_grass.py b/src/xarray_grass/xarray_grass.py index 409db35..02569d0 100644 --- a/src/xarray_grass/xarray_grass.py +++ b/src/xarray_grass/xarray_grass.py @@ -12,17 +12,42 @@ GNU General Public License for more details. """ +from __future__ import annotations + import os +from collections.abc import Iterable from datetime import datetime, timezone from pathlib import Path -from typing import Iterable, Optional +from typing import TYPE_CHECKING, Any, Callable, TypedDict -from xarray.backends import BackendEntrypoint import xarray as xr +from xarray.backends import BackendEntrypoint +from xarray.core.indexing import LazilyIndexedArray import xarray_grass -from xarray_grass.grass_interface import GrassInterface from xarray_grass.grass_backend_array import GrassSTDSBackendArray +from xarray_grass.grass_interface import GrassInterface + +if TYPE_CHECKING: + from xarray.backends.common import AbstractDataStore + from xarray.core.types import ReadBuffer + + +class OpenMapParams(TypedDict): + raster_list: list[str] + raster_3d_list: list[str] + strds_list: list[str] + str3ds_list: list[str] + + +def _path_from_input(filename_or_obj: object) -> Path: + if isinstance(filename_or_obj, str): + return Path(filename_or_obj) + if isinstance(filename_or_obj, os.PathLike): + path_value = filename_or_obj.__fspath__() + if isinstance(path_value, str): + return Path(path_value) + raise TypeError("The GRASS backend requires a text filesystem path.") class GrassBackendEntrypoint(BackendEntrypoint): @@ -42,13 +67,13 @@ class GrassBackendEntrypoint(BackendEntrypoint): def open_dataset( self, - filename_or_obj, + filename_or_obj: str | os.PathLike[Any] | ReadBuffer | AbstractDataStore, *, - raster: Optional[str | Iterable[str]] = None, - raster_3d: Optional[str | Iterable[str]] = None, - strds: Optional[str | Iterable[str]] = None, - str3ds: Optional[str | Iterable[str]] = None, - drop_variables: Iterable[str], + raster: str | Iterable[str] | None = None, + raster_3d: str | Iterable[str] | None = None, + strds: str | Iterable[str] | None = None, + str3ds: str | Iterable[str] | None = None, + drop_variables: str | Iterable[str] | None = None, ) -> xr.Dataset: """Open GRASS project or mapset as an xarray.Dataset. Requires an active GRASS session. @@ -60,46 +85,70 @@ def open_dataset( "Please setup a GRASS session before trying to access GRASS data." ) - if filename_or_obj: - dirpath = Path(filename_or_obj) - if not dir_is_grass_mapset(dirpath): - raise ValueError(f"{filename_or_obj} is not a GRASS mapset") - self.check_accessible_mapset(filename_or_obj) + dirpath = _path_from_input(filename_or_obj) + if not dir_is_grass_mapset(dirpath): + raise ValueError(f"{filename_or_obj} is not a GRASS mapset") self.grass_interface = GrassInterface() - - open_func_params = dict( - raster_list=raster, - raster_3d_list=raster_3d, - strds_list=strds, - str3ds_list=str3ds, + self.check_accessible_mapset(dirpath) + + def as_list(value: str | Iterable[str] | None) -> list[str]: + if isinstance(value, str): + return [value] + if value is None: + return [] + return list(value) + + open_func_params = OpenMapParams( + raster_list=as_list(raster), + raster_3d_list=as_list(raster_3d), + strds_list=as_list(strds), + str3ds_list=as_list(str3ds), ) if not any([raster, raster_3d, strds, str3ds]): self._list_all_mapset(open_func_params) - else: - # Format str inputs into list - for object_type, elem in open_func_params.items(): - if isinstance(elem, str): - open_func_params[object_type] = [elem] - elif elem is None: - open_func_params[object_type] = [] - else: - open_func_params[object_type] = list(elem) # drop requested variables if drop_variables is not None: - for object_type, grass_obj_name_list in open_func_params.items(): - open_func_params[object_type] = [ - name for name in grass_obj_name_list if name not in drop_variables - ] + dropped_names = ( + {drop_variables} + if isinstance(drop_variables, str) + else set(drop_variables) + ) + open_func_params = OpenMapParams( + raster_list=[ + name + for name in open_func_params["raster_list"] + if name not in dropped_names + ], + raster_3d_list=[ + name + for name in open_func_params["raster_3d_list"] + if name not in dropped_names + ], + strds_list=[ + name + for name in open_func_params["strds_list"] + if name not in dropped_names + ], + str3ds_list=[ + name + for name in open_func_params["str3ds_list"] + if name not in dropped_names + ], + ) - return self._open_grass_maps(filename_or_obj, **open_func_params) + return self._open_grass_maps(dirpath, **open_func_params) - def guess_can_open(self, filename_or_obj) -> bool: + def guess_can_open(self, filename_or_obj: object) -> bool: """infer if the path is a GRASS mapset. TODO: add support for whole project.""" - return dir_is_grass_mapset(filename_or_obj) + try: + dirpath = _path_from_input(filename_or_obj) + except TypeError: + return False + return dir_is_grass_mapset(dirpath) - def _list_all_mapset(self, open_func_params): + def _list_all_mapset(self, open_func_params: OpenMapParams) -> None: """List map objects in the whole mapset. If a map is part of a STDS, do not list it as a single map. """ @@ -109,19 +158,13 @@ def _list_all_mapset(self, open_func_params): for strds_name in grass_objects["strds"]: maps_in_strds = self.grass_interface.list_maps_in_strds(strds_name) rasters_in_strds.extend([map_data.id for map_data in maps_in_strds]) - if open_func_params["strds_list"] is None: - open_func_params["strds_list"] = [strds_name] - else: - open_func_params["strds_list"].append(strds_name) + open_func_params["strds_list"].append(strds_name) raster3ds_in_str3ds = [] # str3ds for str3ds_name in grass_objects["str3ds"]: maps_in_str3ds = self.grass_interface.list_maps_in_str3ds(str3ds_name) raster3ds_in_str3ds.extend([map_data.id for map_data in maps_in_str3ds]) - if open_func_params["str3ds_list"] is None: - open_func_params["str3ds_list"] = [str3ds_name] - else: - open_func_params["str3ds_list"].append(str3ds_name) + open_func_params["str3ds_list"].append(str3ds_name) # rasters not in strds open_func_params["raster_list"] = [ name for name in grass_objects["raster"] if name not in rasters_in_strds @@ -133,8 +176,7 @@ def _list_all_mapset(self, open_func_params): if name not in raster3ds_in_str3ds ] - def check_accessible_mapset(self, filename_or_obj): - dirpath = Path(filename_or_obj) + def check_accessible_mapset(self, dirpath: Path) -> None: mapset = dirpath.stem project_path = dirpath.parent gisdb_path = project_path.parent @@ -158,52 +200,66 @@ def check_accessible_mapset(self, filename_or_obj): def _open_grass_maps( self, - filename_or_obj: str | Path, - raster_list: Iterable[str] = None, - raster_3d_list: Iterable[str] = None, - strds_list: Iterable[str] = None, - str3ds_list: Iterable[str] = None, + filename_or_obj: Path, + raster_list: list[str], + raster_3d_list: list[str], + strds_list: list[str], + str3ds_list: list[str], ) -> xr.Dataset: """ Open a GRASS mapset and return an xarray dataset. """ # Configuration for processing different GRASS map types - map_processing_configs = [ - { - "input_list": raster_list, - "existence_check_method": self.grass_interface.name_is_raster, - "open_function": self._open_grass_raster, - "not_found_key": "raster", - }, - { - "input_list": raster_3d_list, - "existence_check_method": self.grass_interface.name_is_raster_3d, - "open_function": self._open_grass_raster_3d, - "not_found_key": "raster_3d", - }, - { - "input_list": strds_list, - "existence_check_method": self.grass_interface.name_is_strds, - "open_function": self._open_grass_strds, - "not_found_key": "strds", - }, - { - "input_list": str3ds_list, - "existence_check_method": self.grass_interface.name_is_str3ds, - "open_function": self._open_grass_str3ds, - "not_found_key": "str3ds", - }, + map_processing_configs: list[ + tuple[ + list[str], + Callable[[str], bool], + Callable[[str], xr.DataArray], + str, + ] + ] = [ + ( + raster_list, + self.grass_interface.name_is_raster, + self._open_grass_raster, + "raster", + ), + ( + raster_3d_list, + self.grass_interface.name_is_raster_3d, + self._open_grass_raster_3d, + "raster_3d", + ), + ( + strds_list, + self.grass_interface.name_is_strds, + self._open_grass_strds, + "strds", + ), + ( + str3ds_list, + self.grass_interface.name_is_str3ds, + self._open_grass_str3ds, + "str3ds", + ), ] # Open all given maps and identify non-existent data - not_found = {config["not_found_key"]: [] for config in map_processing_configs} - data_array_list = [] + not_found: dict[str, list[str]] = { + not_found_key: [] for _, _, _, not_found_key in map_processing_configs + } + data_array_list: list[xr.DataArray] = [] raw_coords_list = [] - for config in map_processing_configs: - for map_name in config["input_list"]: - if not config["existence_check_method"](map_name): - not_found[config["not_found_key"]].append(map_name) + for ( + input_list, + existence_check, + open_function, + not_found_key, + ) in map_processing_configs: + for map_name in input_list: + if not existence_check(map_name): + not_found[not_found_key].append(map_name) continue - data_array = config["open_function"](map_name) + data_array = open_function(map_name) raw_coords_list.append(data_array.coords) data_array_list.append(data_array) if any(not_found.values()): @@ -240,7 +296,7 @@ def _set_cf_coordinates_attributes( da: xr.DataArray, is_3d: bool, z_unit: str = "", - time_dims: Optional[list[str, str]] = None, + time_dims: tuple[str, str] | None = None, time_unit: str = "", ): """Set coordinate attributes according to CF conventions""" @@ -355,12 +411,14 @@ def _open_grass_strds(self, strds_name: str) -> xr.DataArray: if strds_infos.temporal_type == "absolute": time_unit = "" else: - time_unit = strds_infos.time_unit + time_unit = strds_infos.time_unit or "" start_time_dim = f"start_time_{strds_name}" end_time_dim = f"end_time_{strds_name}" map_list = self.grass_interface.list_maps_in_strds(strds_id) region = self.grass_interface.get_region() + if region.rows is None or region.cols is None: + raise ValueError("The current GRASS region is missing its 2D shape.") # Create a single backend array for the entire STRDS backend_array = GrassSTDSBackendArray( @@ -370,7 +428,7 @@ def _open_grass_strds(self, strds_name: str) -> xr.DataArray: map_type="raster", grass_interface=self.grass_interface, ) - lazy_array = xr.core.indexing.LazilyIndexedArray(backend_array) + lazy_array = LazilyIndexedArray(backend_array) # Create Variable with lazy array var = xr.Variable(dims=[start_time_dim, "y", "x"], data=lazy_array) @@ -399,7 +457,7 @@ def _open_grass_strds(self, strds_name: str) -> xr.DataArray: da_with_attrs = self._set_cf_coordinates_attributes( data_array, is_3d=False, - time_dims=[start_time_dim, end_time_dim], + time_dims=(start_time_dim, end_time_dim), time_unit=time_unit, ) da_with_attrs.attrs["long_name"] = strds_infos.title @@ -424,12 +482,14 @@ def _open_grass_str3ds(self, str3ds_name: str) -> xr.DataArray: if strds_infos.temporal_type == "absolute": time_unit = "" else: - time_unit = strds_infos.time_unit + time_unit = strds_infos.time_unit or "" start_time_dim = f"start_time_{str3ds_name}" end_time_dim = f"end_time_{str3ds_name}" map_list = self.grass_interface.list_maps_in_str3ds(str3ds_id) region = self.grass_interface.get_region() + if region.depths is None or region.rows3 is None or region.cols3 is None: + raise ValueError("The current GRASS region is missing its 3D shape.") # Create a single backend array for the entire STR3DS backend_array = GrassSTDSBackendArray( @@ -439,7 +499,7 @@ def _open_grass_str3ds(self, str3ds_name: str) -> xr.DataArray: map_type="raster3d", grass_interface=self.grass_interface, ) - lazy_array = xr.core.indexing.LazilyIndexedArray(backend_array) + lazy_array = LazilyIndexedArray(backend_array) # Create Variable with lazy array var = xr.Variable(dims=[start_time_dim, "z", "y_3d", "x_3d"], data=lazy_array) @@ -470,7 +530,7 @@ def _open_grass_str3ds(self, str3ds_name: str) -> xr.DataArray: data_array, is_3d=True, z_unit=r3_infos["vertical_units"], - time_dims=[start_time_dim, end_time_dim], + time_dims=(start_time_dim, end_time_dim), time_unit=time_unit, ) da_with_attrs.attrs["long_name"] = strds_infos.title diff --git a/tests/test_tograss_error_handling.py b/tests/test_tograss_error_handling.py index 61f4feb..039f184 100644 --- a/tests/test_tograss_error_handling.py +++ b/tests/test_tograss_error_handling.py @@ -105,6 +105,18 @@ def test_mapset_not_accessible_simplified(self, grass_i: GrassInterface): @pytest.mark.usefixtures("grass_session_fixture") class TestToGrassInputValidation: + def test_unnamed_dataarray(self): + unnamed = xr.DataArray( + np.ones((2, 2)), + coords={"y": [0, 1], "x": [0, 1]}, + dims=("y", "x"), + ) + + with pytest.raises( + ValueError, match="GRASS object names must be non-empty strings" + ): + to_grass(unnamed) + def test_invalid_dataset_type(self, temp_gisdb, grass_i: GrassInterface): """Test error handling for invalid 'dataset' parameter type. That a first try. Let's see how it goes considering that the tested code uses duck typing.""" diff --git a/tests/test_xarray_grass.py b/tests/test_xarray_grass.py index adbade5..ddfb548 100644 --- a/tests/test_xarray_grass.py +++ b/tests/test_xarray_grass.py @@ -259,6 +259,21 @@ def test_drop_variables(self, grass_i, temp_gisdb) -> None: assert len(test_dataset.x) == region.cols assert len(test_dataset.y) == region.rows + def test_drop_variable_string(self, grass_i, temp_gisdb) -> None: + mapset_path = os.path.join( + str(temp_gisdb.gisdb), str(temp_gisdb.project), str(temp_gisdb.mapset) + ) + test_dataset = xr.open_dataset( + mapset_path, + raster=[ACTUAL_RASTER_MAP, ACTUAL_RASTER_MAP2], + drop_variables=ACTUAL_RASTER_MAP, + ) + + dropped_name = grass_i.get_name_from_id(ACTUAL_RASTER_MAP) + retained_name = grass_i.get_name_from_id(ACTUAL_RASTER_MAP2) + assert dropped_name not in test_dataset + assert retained_name in test_dataset + def test_attributes_separation(self, grass_i, temp_gisdb) -> None: """Test that DataArray attributes don't leak to Dataset level.""" mapset_path = os.path.join(