Skip to content

Commit 848e9d2

Browse files
committed
Remove None as option for vector_fields
1 parent f71233e commit 848e9d2

3 files changed

Lines changed: 19 additions & 24 deletions

File tree

src/parcels/_core/fieldset.py

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -211,7 +211,7 @@ def from_ugrid_conventions(
211211
cls,
212212
ds: ux.UxDataset,
213213
mesh: str = "spherical",
214-
vector_fields: ptyping.VectorFields | None | NotSetType = NOTSET,
214+
vector_fields: ptyping.VectorFields | NotSetType = NOTSET,
215215
):
216216
"""Create a FieldSet from a Parcels compliant uxarray.UxDataset.
217217
@@ -229,8 +229,7 @@ def from_ugrid_conventions(
229229
vector_fields : Mapping[str, tuple[str, ...]] or None, optional
230230
Mapping of vector field names to tuples of component variable names in the dataset.
231231
For example, ``{"UV": ("U", "V"), "UVW": ("U", "V", "W")}``.
232-
If ``None``, no vector fields are constructed. If omitted (default), vector fields
233-
are auto-discovered from standard variable names (``U``/``V``/``W``).
232+
If omitted (default), vector fields are auto-discovered from standard variable names (``U``/``V``/``W``).
234233
235234
Returns
236235
-------
@@ -249,7 +248,7 @@ def from_sgrid_conventions(
249248
cls,
250249
ds: xr.Dataset,
251250
mesh: ptyping.Mesh | None = None,
252-
vector_fields: ptyping.VectorFields | None | NotSetType = NOTSET,
251+
vector_fields: ptyping.VectorFields | NotSetType = NOTSET,
253252
): # TODO: Update mesh to be discovered from the dataset metadata
254253
"""Create a FieldSet from a dataset using SGRID convention metadata.
255254
@@ -267,8 +266,7 @@ def from_sgrid_conventions(
267266
vector_fields : Mapping[str, tuple[str, ...]] or None, optional
268267
Mapping of vector field names to tuples of component variable names in the dataset.
269268
For example, ``{"UV": ("U", "V"), "UVW": ("U", "V", "W")}``.
270-
If ``None``, no vector fields are constructed. If omitted (default), vector fields
271-
are auto-discovered from standard variable names (``U``/``V``/``W``).
269+
If omitted (default), vector fields are auto-discovered from standard variable names (``U``/``V``/``W``).
272270
273271
Returns
274272
-------

src/parcels/_core/model.py

Lines changed: 6 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -132,7 +132,7 @@ def construct_fields(self) -> list[Field | VectorField]:
132132

133133
@classmethod
134134
def from_sgrid_conventions(
135-
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: ptyping.VectorFields | None | NotSetType
135+
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: ptyping.VectorFields | NotSetType
136136
) -> Self:
137137
ds = ds.copy()
138138
if mesh is None:
@@ -171,19 +171,15 @@ def from_sgrid_conventions(
171171
return model
172172

173173

174-
def resolve_vector_fields(
175-
ds: xr.Dataset, vector_fields: ptyping.VectorFields | None | NotSetType
176-
) -> ptyping.VectorFields:
177-
if vector_fields is None:
178-
return {}
174+
def resolve_vector_fields(ds: xr.Dataset, vector_fields: ptyping.VectorFields | NotSetType) -> ptyping.VectorFields:
179175
if vector_fields is NOTSET: # i.e., the default vectorfield discovery behaviour
180176
return _default_vector_field_components(list(ds.data_vars))
181177
return vector_fields
182178

183179

184180
def assert_valid_vector_fields(ds: xr.Dataset, vector_fields: ptyping.VectorFields) -> None:
185-
# if not isinstance(vector_fields, dict):
186-
# raise ValueError(f"vector_fields must be a dictionary. Got {type(vector_fields)=!r}.")
181+
if not isinstance(vector_fields, dict):
182+
raise ValueError(f"vector_fields must be a dictionary. Got {type(vector_fields)=!r}.")
187183

188184
for vfield_name, components in vector_fields.items():
189185
if not isinstance(vfield_name, str):
@@ -237,7 +233,7 @@ def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: pty
237233
),
238234
),
239235
mesh=mesh, # type:ignore
240-
vector_fields=None,
236+
vector_fields={},
241237
)
242238
for mesh in ["flat", "spherical"]
243239
}
@@ -283,9 +279,7 @@ def scalar_field_names(self) -> list[str]:
283279
return list(self.data.data_vars)
284280

285281
@classmethod
286-
def from_ugrid_conventions(
287-
cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: ptyping.VectorFields | None | NotSetType
288-
):
282+
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: ptyping.VectorFields | NotSetType):
289283
ds_dims = list(ds.dims)
290284
if not all(dim in ds_dims for dim in ["time", "zf", "zc"]):
291285
raise ValueError(

tests/test_fieldset.py

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,3 @@
1-
from contextlib import nullcontext
21
from datetime import timedelta
32

43
import cf_xarray # noqa: F401
@@ -116,7 +115,11 @@ def test_fieldset_from_structured_generic_datasets(ds):
116115
pytest.raises(ValueError, match="must have either 2 or 3 components"),
117116
id="too-many-components",
118117
),
119-
pytest.param(None, nullcontext(), id="None"),
118+
pytest.param(
119+
None,
120+
pytest.raises(ValueError, match="vector_fields must be a dictionary"),
121+
id="None",
122+
),
120123
],
121124
)
122125
def test_fieldset_invalid_vector_fields(vector_fields, ctx):
@@ -147,10 +150,10 @@ def test_fieldset_structured_vectorfield_custom():
147150
assert "UV_wind" in fset.fields
148151

149152

150-
def test_fieldset_structured_vectorfield_none():
153+
def test_fieldset_structured_vectorfield_empty():
151154
ds = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename({"U_A_grid": "U", "V_A_grid": "V"})
152155

153-
fset = FieldSet.from_sgrid_conventions(ds, mesh="flat", vector_fields=None)
156+
fset = FieldSet.from_sgrid_conventions(ds, mesh="flat", vector_fields={})
154157

155158
assert "U" in fset.fields
156159
assert "V" in fset.fields
@@ -177,10 +180,10 @@ def test_fieldset_unstructured_vectorfield_custom():
177180
assert "UV_wind" in fset.fields
178181

179182

180-
def test_fieldset_unstructured_vectorfield_none():
183+
def test_fieldset_unstructured_vectorfield_empty():
181184
ds = datasets_unstructured["stommel_gyre_delaunay"]
182185

183-
fset = FieldSet.from_ugrid_conventions(ds, mesh="spherical", vector_fields=None)
186+
fset = FieldSet.from_ugrid_conventions(ds, mesh="spherical", vector_fields={})
184187

185188
assert "U" in fset.fields
186189
assert "V" in fset.fields

0 commit comments

Comments
 (0)