Skip to content

Commit f71233e

Browse files
committed
Improve validation of vector_fields
1 parent 3a85df3 commit f71233e

2 files changed

Lines changed: 54 additions & 2 deletions

File tree

src/parcels/_core/model.py

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -161,7 +161,7 @@ def from_sgrid_conventions(
161161
# ds["lat"] = ds[node_dimensions[1]]
162162

163163
vector_fields = resolve_vector_fields(ds, vector_fields)
164-
assert_vector_field_components_in_dataset(ds, vector_fields)
164+
assert_valid_vector_fields(ds, vector_fields)
165165

166166
model = cls(ds, mesh=mesh, vector_field_components=vector_fields)
167167
model._fields = model.construct_fields()
@@ -181,6 +181,29 @@ def resolve_vector_fields(
181181
return vector_fields
182182

183183

184+
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}.")
187+
188+
for vfield_name, components in vector_fields.items():
189+
if not isinstance(vfield_name, str):
190+
raise ValueError(
191+
f"Invalid `vector_fields` argument. Vector field name in `vector_fields` should be a string. Got field name {vfield_name!r}."
192+
)
193+
if not (2 <= len(components) <= 3):
194+
raise ValueError(
195+
f"Invalid `vector_fields` argument. Vector fields must have either 2 or 3 components. Vector field {vfield_name} has {len(components)} components."
196+
)
197+
for c in components:
198+
if not isinstance(c, str):
199+
raise ValueError(
200+
f"Invalid `vector_fields` argument. Component names must be strings. Got component name of value {c!r}."
201+
)
202+
203+
assert_vector_field_components_in_dataset(ds, vector_fields)
204+
return
205+
206+
184207
def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: ptyping.VectorFields) -> None:
185208
for components in vector_fields.values():
186209
for c in components:
@@ -273,7 +296,7 @@ def from_ugrid_conventions(
273296
ds = _discover_ux_U_and_V(ds)
274297

275298
vector_fields = resolve_vector_fields(ds, vector_fields)
276-
assert_vector_field_components_in_dataset(ds, vector_fields)
299+
assert_valid_vector_fields(ds, vector_fields)
277300

278301
model = cls(ds, grid, vector_fields)
279302
model._fields = model.construct_fields()

tests/test_fieldset.py

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from contextlib import nullcontext
12
from datetime import timedelta
23

34
import cf_xarray # noqa: F401
@@ -97,6 +98,34 @@ def test_fieldset_from_structured_generic_datasets(ds):
9798
assert len(fieldset.gridset) == 1
9899

99100

101+
@pytest.mark.parametrize(
102+
"vector_fields,ctx",
103+
[
104+
pytest.param(
105+
{"UV": ("U",)},
106+
pytest.raises(ValueError, match="must have either 2 or 3 components"),
107+
id="single-component",
108+
),
109+
pytest.param(
110+
{"UV": ("U", "missing")},
111+
pytest.raises(ValueError, match="not present in the source dataset"),
112+
id="component-not-in-dataset",
113+
),
114+
pytest.param(
115+
{"UV": ("U", "U", "U", "U")},
116+
pytest.raises(ValueError, match="must have either 2 or 3 components"),
117+
id="too-many-components",
118+
),
119+
pytest.param(None, nullcontext(), id="None"),
120+
],
121+
)
122+
def test_fieldset_invalid_vector_fields(vector_fields, ctx):
123+
ds = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename({"U_A_grid": "U", "V_A_grid": "V"})
124+
125+
with ctx:
126+
FieldSet.from_sgrid_conventions(ds, mesh="flat", vector_fields=vector_fields)
127+
128+
100129
def test_fieldset_structured_vectorfield_default():
101130
ds = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename({"U_A_grid": "U", "V_A_grid": "V"})
102131

0 commit comments

Comments
 (0)