|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
3 | 3 | from abc import ABC, abstractmethod |
4 | | -from collections.abc import Mapping, Sequence |
| 4 | +from collections.abc import Hashable, Sequence |
5 | 5 | from typing import Any, Self |
6 | 6 |
|
7 | 7 | import cf_xarray # noqa: F401 |
|
34 | 34 | ) |
35 | 35 | from parcels.interpolators._base import ScalarInterpolator, VectorInterpolator |
36 | 36 |
|
37 | | -TVectorFieldMapping = Mapping[str, tuple[str, str] | tuple[str, str, str]] |
| 37 | +TVectorFieldMapping = dict[str, tuple[str, str] | tuple[str, str, str]] |
38 | 38 |
|
39 | 39 |
|
40 | 40 | class ModelData(ABC): |
@@ -126,7 +126,7 @@ def construct_fields(self) -> list[Field | VectorField]: |
126 | 126 | ) |
127 | 127 |
|
128 | 128 | 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] |
130 | 130 |
|
131 | 131 | fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields} |
132 | 132 | return list(fields.values()) |
@@ -178,7 +178,7 @@ def resolve_vector_fields( |
178 | 178 | if vector_fields is None: |
179 | 179 | return {} |
180 | 180 | 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)) |
182 | 182 | return vector_fields |
183 | 183 |
|
184 | 184 |
|
@@ -247,7 +247,7 @@ def construct_fields(self) -> list[Field | VectorField]: |
247 | 247 | interp_method = Ux_Velocity() |
248 | 248 |
|
249 | 249 | 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] |
251 | 251 |
|
252 | 252 | fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields} |
253 | 253 | return list(fields.values()) |
@@ -304,9 +304,9 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh: |
304 | 304 | return "spherical" if _is_coordinate_in_degrees(ds_sgrid[fpoint_x]) else "flat" |
305 | 305 |
|
306 | 306 |
|
307 | | -def _default_vector_field_components(data_vars: Sequence[str]) -> TVectorFieldMapping: |
| 307 | +def _default_vector_field_components(data_vars: Sequence[Hashable]) -> TVectorFieldMapping: |
308 | 308 | vars = set(data_vars) |
309 | | - ret = {} |
| 309 | + ret: TVectorFieldMapping = {} |
310 | 310 |
|
311 | 311 | if {"U", "V"}.issubset(vars): |
312 | 312 | ret["UV"] = ("U", "V") |
|
0 commit comments