11from __future__ import annotations
22
33from abc import ABC , abstractmethod
4- from collections .abc import Mapping , Sequence
4+ from collections .abc import Hashable , Sequence
55from typing import Any , Self
66
77import cf_xarray # noqa: F401
3434)
3535from parcels .interpolators ._base import ScalarInterpolator , VectorInterpolator
3636
37- TVectorFieldMapping = Mapping [str , tuple [str , str ] | tuple [str , str , str ]]
37+ TVectorField = dict [str , tuple [str , str ] | tuple [str , str , str ]]
3838
3939
4040class ModelData (ABC ):
4141 data : Any
4242 grid : BaseGrid
4343 field_to_interpolator : dict [str , ScalarInterpolator | VectorInterpolator ]
44- vector_field_components : TVectorFieldMapping
44+ vector_field_components : TVectorField
4545
4646 @abstractmethod
4747 def construct_fields (self ) -> list [Field | VectorField ]: ...
@@ -84,7 +84,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset:
8484
8585
8686class StructuredModelData (ModelData ):
87- def __init__ (self , data : xr .Dataset , mesh : Mesh , vector_field_components : TVectorFieldMapping ):
87+ def __init__ (self , data : xr .Dataset , mesh : Mesh , vector_field_components : TVectorField ):
8888 if not isinstance (data , xr .Dataset ):
8989 raise ValueError (f"Expected `data` to be an xarray.Dataset . Got { type (data )} " )
9090
@@ -126,14 +126,14 @@ def construct_fields(self) -> list[Field | VectorField]:
126126 )
127127
128128 component_fields = [single_fields [name ] for name in components ]
129- vector_fields [vfield_name ] = VectorField (vfield_name , * component_fields , interp_method = interp_method )
129+ vector_fields [vfield_name ] = VectorField (vfield_name , * component_fields , interp_method = interp_method ) # type:ignore[misc,arg-type]
130130
131131 fields : dict [str , Field | VectorField ] = {** single_fields , ** vector_fields }
132132 return list (fields .values ())
133133
134134 @classmethod
135135 def from_sgrid_conventions (
136- cls , ds : xr .Dataset , mesh : Mesh | None , vector_fields : TVectorFieldMapping | None | NotSetType
136+ cls , ds : xr .Dataset , mesh : Mesh | None , vector_fields : TVectorField | None | NotSetType
137137 ) -> Self :
138138 ds = ds .copy ()
139139 if mesh is None :
@@ -173,16 +173,16 @@ def from_sgrid_conventions(
173173
174174
175175def resolve_vector_fields (
176- ds : xr .Dataset , vector_fields : TVectorFieldMapping | None | NotSetType
177- ) -> TVectorFieldMapping :
176+ ds : xr .Dataset , vector_fields : TVectorField | None | NotSetType
177+ ) -> TVectorField :
178178 if vector_fields is None :
179179 return {}
180180 if vector_fields is NOTSET : # i.e., the default vectorfield discovery behaviour
181- return _default_vector_field_components (ds .data_vars )
181+ return _default_vector_field_components (list ( ds .data_vars ) )
182182 return vector_fields
183183
184184
185- def assert_vector_field_components_in_dataset (ds : xr .Dataset , vector_fields : TVectorFieldMapping ) -> None :
185+ def assert_vector_field_components_in_dataset (ds : xr .Dataset , vector_fields : TVectorField ) -> None :
186186 for components in vector_fields .values ():
187187 for c in components :
188188 if c not in ds .data_vars :
@@ -222,7 +222,7 @@ def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: TVe
222222
223223
224224class UnstructuredModelData (ModelData ):
225- def __init__ (self , data : ux .UxDataset , grid : UxGrid , vector_field_components : TVectorFieldMapping ):
225+ def __init__ (self , data : ux .UxDataset , grid : UxGrid , vector_field_components : TVectorField ):
226226 if not isinstance (data , ux .UxDataset ):
227227 raise ValueError (f"Expected `data` to be an uxarray.UxDataset . Got { type (data )} " )
228228
@@ -247,7 +247,7 @@ def construct_fields(self) -> list[Field | VectorField]:
247247 interp_method = Ux_Velocity ()
248248
249249 component_fields = [single_fields [name ] for name in components ]
250- vector_fields [vfield_name ] = VectorField (vfield_name , * component_fields , interp_method = interp_method )
250+ vector_fields [vfield_name ] = VectorField (vfield_name , * component_fields , interp_method = interp_method ) # type:ignore[misc, arg-type]
251251
252252 fields : dict [str , Field | VectorField ] = {** single_fields , ** vector_fields }
253253 return list (fields .values ())
@@ -262,7 +262,7 @@ def scalar_field_names(self) -> list[str]:
262262
263263 @classmethod
264264 def from_ugrid_conventions (
265- cls , ds : ux .UxDataset , mesh : Mesh , vector_fields : TVectorFieldMapping | None | NotSetType
265+ cls , ds : ux .UxDataset , mesh : Mesh , vector_fields : TVectorField | None | NotSetType
266266 ):
267267 ds_dims = list (ds .dims )
268268 if not all (dim in ds_dims for dim in ["time" , "zf" , "zc" ]):
@@ -304,9 +304,9 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh:
304304 return "spherical" if _is_coordinate_in_degrees (ds_sgrid [fpoint_x ]) else "flat"
305305
306306
307- def _default_vector_field_components (data_vars : Sequence [str ]) -> TVectorFieldMapping :
307+ def _default_vector_field_components (data_vars : Sequence [Hashable ]) -> TVectorField :
308308 vars = set (data_vars )
309- ret = {}
309+ ret : TVectorField = {}
310310
311311 if {"U" , "V" }.issubset (vars ):
312312 ret ["UV" ] = ("U" , "V" )
0 commit comments