Skip to content

Commit 3dbd950

Browse files
committed
Fix mypy issues
1 parent 8043eec commit 3dbd950

2 files changed

Lines changed: 18 additions & 18 deletions

File tree

src/parcels/_core/fieldset.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414
CONSTANT_FIELD_MODELS,
1515
ModelData,
1616
StructuredModelData,
17-
TVectorFieldMapping,
17+
TVectorField,
1818
UnstructuredModelData,
1919
)
2020
from parcels._core.utils.string import _assert_str_and_python_varname
@@ -212,7 +212,7 @@ def from_ugrid_conventions(
212212
cls,
213213
ds: ux.UxDataset,
214214
mesh: str = "spherical",
215-
vector_fields: TVectorFieldMapping | None | NotSetType = NOTSET,
215+
vector_fields: TVectorField | None | NotSetType = NOTSET,
216216
):
217217
"""Create a FieldSet from a Parcels compliant uxarray.UxDataset.
218218
@@ -250,7 +250,7 @@ def from_sgrid_conventions(
250250
cls,
251251
ds: xr.Dataset,
252252
mesh: Mesh | None = None,
253-
vector_fields: TVectorFieldMapping | None | NotSetType = NOTSET,
253+
vector_fields: TVectorField | None | NotSetType = NOTSET,
254254
): # TODO: Update mesh to be discovered from the dataset metadata
255255
"""Create a FieldSet from a dataset using SGRID convention metadata.
256256

src/parcels/_core/model.py

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
from __future__ import annotations
22

33
from abc import ABC, abstractmethod
4-
from collections.abc import Mapping, Sequence
4+
from collections.abc import Hashable, Sequence
55
from typing import Any, Self
66

77
import cf_xarray # noqa: F401
@@ -34,14 +34,14 @@
3434
)
3535
from 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

4040
class 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

8686
class 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

175175
def 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

224224
class 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

Comments
 (0)