diff --git a/docs/dev/sdd.md b/docs/dev/sdd.md index cba712e0..a2ef5d11 100644 --- a/docs/dev/sdd.md +++ b/docs/dev/sdd.md @@ -143,7 +143,7 @@ IO is at the product's boundary. Details of any input or output format should no Input file IO is implemented in three layers: 1. **Unified IO layer**: Registry and descriptors implementing `load` and `write` methods on the base `Component` class -2. **Conversion layer**: Uses `cattrs` to map the object model to/from Python primitives and containers (i.e. un/structuring) +2. **Conversion layer**: Maps the object model to/from Python primitives and containers (i.e. un/structuring) 3. **Serialization layer**: Format-specific encoders/decoders translating primitives and containers to/from strings or binary data #### Unified IO @@ -156,9 +156,9 @@ Loaders and writers can be registered for any component class and format. The re The sparse input format used by MODFLOW 6 is in tension with an object model where tables are disaggregated into a separate array variable for each column — this requires a nontrivial conversion at load and write time. -The conversion layer uses `cattrs` to transform between the product's `xarray`/`attrs`-based object model and plain Python data structures suitable for serialization. This layer is format-agnostic and handles structural transformations common across formats. +The conversion layer (`flopy4.mf6.converter`) transforms between the product's object model (pydantic dataclasses holding `xarray`/`numpy` data) and plain Python data structures suitable for serialization. This layer is format-agnostic and handles structural transformations common across formats. -**Unstructuring (write time)**: A `cattrs` converter with appropriate unstructuring hooks converts components to a form suitable for serialization, handling transformations like: +**Unstructuring (write time)**: `converter.egress.unstructure` converts components to a form suitable for serialization, handling transformations like: - Grouping fields into blocks according to their `block` metadata from DFNs - Converting child components to binding tables for parent component name files @@ -167,7 +167,7 @@ The conversion layer uses `cattrs` to transform between the product's `xarray`/` The unstructuring phase aims to avoid a) unnecessary copies and b) materializing data in memory. -**Structuring (load time)**: A `cattrs` converter with appropriate structuring hooks converts dictionaries of primitives into component instances, including: +**Structuring (load time)**: `converter.ingress.structure` converts dictionaries of primitives into component instances, including: - Instantiating child components from bindings - Converting sparse list input data representations to arrays @@ -204,7 +204,7 @@ The reader in `flopy4.mf6.codec.reader` uses [Lark](https://lark-parser.readthed **Parsing**: A minimal *basic* grammar recognizes only block structure — blocks delimited by `BEGIN ` / `END `, each containing lines of whitespace-separated tokens (words and numbers). `BasicTransformer` yields a `{BLOCK_NAME: [token_row, ...]}` mapping. This grammar is component-agnostic, so one parser handles every input file. -**Structuring**: `converter.ingress.structure` reconstructs a component from the parsed block mapping, using the component class's `attrs` field metadata as the specification: +**Structuring**: `converter.ingress.structure` reconstructs a component from the parsed block mapping, using the component class's field metadata as the specification: - Blocks map to fields by name; field metadata identifies each field's kind (scalar, keyword, array, record, list). - List and record blocks are parsed into typed `Item` / `Record` objects (one class per row shape), resolving cellid width, `AUXILIARY` columns, and boundnames. diff --git a/flopy4/attrs_xarray.py b/flopy4/dataclass_xarray.py similarity index 53% rename from flopy4/attrs_xarray.py rename to flopy4/dataclass_xarray.py index 3e8fd754..7f374d26 100644 --- a/flopy4/attrs_xarray.py +++ b/flopy4/dataclass_xarray.py @@ -1,21 +1,23 @@ """ -Generic conversion between `attrs`-decorated classes and `xarray` -`Dataset`/`DataTree` objects. +Generic conversion between `pydantic.dataclasses.dataclass`-decorated +classes and `xarray` `Dataset`/`DataTree` objects. Reads only two things per field: its own instance value (via `getattr`) and, for array values, the field's own `shape` metadata (a tuple of dimension-name strings -- the same convention `flopy4.mf6.spec.field()`/ `array()` already tag leaf DFN packages with). -Child detection (which fields hold nested attrs instances, as opposed to -plain scalar/array leaf values) is done by inspecting instance *values* -with `attrs.has()` when building a tree (mirrors the `attrs.fields()` + -`isinstance` pattern `flopy4.dimensions` already established), and by +Child detection (which fields hold nested dataclass instances, as opposed +to plain scalar/array leaf values) is done by inspecting instance *values* +with `is_dataclass_instance()` when building a tree (mirrors the field-walk ++ `isinstance` pattern `flopy4.dimensions` already established), and by inspecting field *type annotations* when reconstructing one from a tree (a bare `DataTree` node carries no back-pointer to the field it came from, so reconstruction has no values to inspect yet). -Known limitation: a field literally named `dims`, `parent`, or `_parent` +InitVars (e.g. `Component.dims`) aren't stored, so they're skipped. + +Known limitation: a field literally named `parent` or `_parent` is always excluded (see `_RESERVED_FIELD_NAMES` for why each one is there). No real DFN field uses any of these names today; if one ever did, it would need a different name or a second, explicit exclusion @@ -27,99 +29,106 @@ field name, unlike "list"-kind's positional `f"{field_name}{index}"` convention. Namefile binding rows carry this information separately; reconstruction should use that rather than guessing here. Likewise, a -"list"-kind field whose element type is itself a `Union` of attrs -classes (e.g. `list[Union[Chd, Chdg]]`) isn't resolved to a concrete -arm. +"list"-kind field whose element type is itself a `Union` of dataclasses +(e.g. `list[Union[Chd, Chdg]]`) isn't resolved to a concrete arm. """ import types from typing import Any, Union, get_args, get_origin -import attrs import numpy as np import xarray as xr +from pydantic.dataclasses import is_pydantic_dataclass +from flopy4.spec import field_meta, is_dataclass_instance, pydantic_fields -def _is_attrs_instance(value: Any) -> bool: - return attrs.has(type(value)) - - -# Field names to always skip, regardless of what they hold. `dims` -# (Component.dims) is a plain dict of already-resolved dimension sizes, -# not xarray-representable leaf/child data. `parent` and `_parent` +# Field names to always skip, regardless of what they hold. `parent` and `_parent` # (Output.parent, Component._parent -- see their own docstrings) are # back-reference fields whose runtime values would otherwise look # structurally like real leaf/child data and cause infinite recursion # (parent -> child -> same parent) if walked, so they're excluded by # name rather than by (unreliable) type-checking. -_RESERVED_FIELD_NAMES = frozenset({"dims", "parent", "_parent"}) +_RESERVED_FIELD_NAMES = frozenset({"parent", "_parent"}) def _leaf_fields_and_children( obj, ) -> "tuple[dict[str, tuple], dict[str, Any], dict[str, dict | list]]": - """Split `obj`'s attrs fields by instance value into: - - leaf fields: ``{name: (attrs.Attribute, value)}`` + """Split `obj`'s pydantic fields by instance value into: + - leaf fields: ``{name: (FieldInfo, value)}`` - single-child fields: ``{name: child_obj}`` - collection-child fields: ``{name: {key: child_obj}}`` or ``{name: [child_obj, ...]}`` """ leaves: "dict[str, tuple]" = {} single_children: "dict[str, Any]" = {} collection_children: "dict[str, dict | list]" = {} - for field in attrs.fields(type(obj)): - if field.name in _RESERVED_FIELD_NAMES: + for name, finfo in pydantic_fields(type(obj)).items(): + if name in _RESERVED_FIELD_NAMES or finfo.init_var: continue - value = getattr(obj, field.name, None) + # A private field (leading underscore) exposed under an alias -- + # e.g. Context._workspace/alias="workspace" -- is keyed by that + # alias here, not its real (private) name, matching the public API + # the field is actually meant to be read/written through (same + # convention structure.py/unstructure.py already use). + exposed = finfo.alias if (finfo.alias and name.startswith("_")) else name + value = getattr(obj, name, None) if value is None: continue - if _is_attrs_instance(value): - single_children[field.name] = value + if is_dataclass_instance(value): + single_children[exposed] = value elif ( - isinstance(value, dict) and value and all(_is_attrs_instance(v) for v in value.values()) + isinstance(value, dict) + and value + and all(is_dataclass_instance(v) for v in value.values()) ): - collection_children[field.name] = value + collection_children[exposed] = value elif ( - isinstance(value, (list, tuple)) and value and all(_is_attrs_instance(v) for v in value) + isinstance(value, (list, tuple)) + and value + and all(is_dataclass_instance(v) for v in value) ): - collection_children[field.name] = list(value) + collection_children[exposed] = list(value) else: - leaves[field.name] = (field, value) + leaves[exposed] = (finfo, value) return leaves, single_children, collection_children -def _array_dims(field: attrs.Attribute, name: str, ndim: int) -> tuple: - shape_meta = field.metadata.get("shape") +def _array_dims(finfo: Any, name: str, ndim: int) -> tuple: + meta = field_meta(finfo) + shape_meta = meta.get("shape") if isinstance(meta, dict) else None if shape_meta and len(shape_meta) == ndim: return tuple(shape_meta) return tuple(f"{name}_dim{i}" for i in range(ndim)) -def attrs_to_dataset(obj) -> xr.Dataset: +def dataclass_to_dataset(obj) -> xr.Dataset: """Flatten `obj`'s own scalar/array fields into a flat `xr.Dataset`. - Attrs-typed child fields are skipped here -- see `attrs_to_datatree()` - for those. A `numpy.ndarray`-valued field becomes a data variable, - with dims named from its `shape` metadata when present (falling back - to generic per-axis names otherwise); everything else becomes a - dataset-level attr. + Dataclass-typed child fields are skipped here -- see + `dataclass_to_datatree()` for those. A `numpy.ndarray`-valued field becomes + a data variable, with dims named from its `shape` metadata when + present (falling back to generic per-axis names otherwise); everything + else becomes a dataset-level attr. """ leaves, _, _ = _leaf_fields_and_children(obj) data_vars = {} ds_attrs = {} - for name, (field, value) in leaves.items(): + for name, (finfo, value) in leaves.items(): + meta = field_meta(finfo) + has_shape = isinstance(meta, dict) and meta.get("shape") if isinstance(value, xr.DataArray): data_vars[name] = value elif isinstance(value, np.ndarray): - data_vars[name] = xr.DataArray(value, dims=_array_dims(field, name, value.ndim)) - elif field.metadata.get("shape") and isinstance(value, (list, tuple)): + data_vars[name] = xr.DataArray(value, dims=_array_dims(finfo, name, value.ndim)) + elif has_shape and isinstance(value, (list, tuple)): # A shape-tagged field whose value hasn't (yet) been coerced to # a real ndarray -- e.g. assigned directly post-construction, - # bypassing whatever coercion __attrs_post_init__ normally does. - # The field's own metadata says it's array-shaped regardless of - # the value's current runtime type, so honor that rather than + # bypassing whatever coercion __post_init__ normally does. The + # field's own metadata says it's array-shaped regardless of the + # value's current runtime type, so honor that rather than # silently dropping it to a dataset-level attr. arr = np.asarray(value) - data_vars[name] = xr.DataArray(arr, dims=_array_dims(field, name, arr.ndim)) + data_vars[name] = xr.DataArray(arr, dims=_array_dims(finfo, name, arr.ndim)) else: ds_attrs[name] = value return xr.Dataset(data_vars, attrs=ds_attrs) @@ -127,9 +136,18 @@ def attrs_to_dataset(obj) -> xr.Dataset: def _init_field_names(cls: type) -> set: # init=False fields (e.g. Dis's derived nodes/ncpl/nvert) have no - # __init__ parameter -- they're recomputed by __attrs_post_init__, not - # round-tripped through the constructor. - return {f.name for f in attrs.fields(cls) if f.init is not False} + # __init__ parameter -- they're recomputed by __post_init__, not + # round-tripped through the constructor. Exposed under each field's + # alias when it has one and is itself private (leading underscore), + # matching _leaf_fields_and_children's own exposed-name convention -- + # a dataset produced from Context.workspace (backed by the private + # _workspace/alias="workspace" field) is keyed "workspace", so lookups + # here must match that key, not the private real name. + return { + (f.alias if (f.alias and name.startswith("_")) else name) + for name, f in pydantic_fields(cls).items() + if f.init is not False + } def _leaf_kwargs_from_dataset(cls: type, dataset: xr.Dataset) -> dict: @@ -144,9 +162,9 @@ def _leaf_kwargs_from_dataset(cls: type, dataset: xr.Dataset) -> dict: return kwargs -def dataset_to_attrs(cls: type, dataset: xr.Dataset): +def dataset_to_dataclass(cls: type, dataset: xr.Dataset): """Construct a `cls` instance from an `xr.Dataset` shaped like - `attrs_to_dataset()`'s output: data variables and dataset-level attrs + `dataclass_to_dataset()`'s output: data variables and dataset-level attrs are matched to `cls`'s fields by name and passed as constructor kwargs. Unmatched dataset keys are ignored; fields with no matching key fall back to `cls`'s own default. @@ -154,13 +172,13 @@ def dataset_to_attrs(cls: type, dataset: xr.Dataset): return cls(**_leaf_kwargs_from_dataset(cls, dataset)) -def attrs_to_datatree(obj, _ancestors: frozenset = frozenset()) -> xr.DataTree: +def dataclass_to_datatree(obj, _ancestors: frozenset = frozenset()) -> xr.DataTree: """Recursively convert `obj` into an `xr.DataTree`. - `obj`'s own leaf fields become the root dataset (`attrs_to_dataset()`). - Each attrs-typed child field becomes a named child node (recursively - converted the same way); a dict- or list-of-children field expands to - one child node per entry, named by its dict key or + `obj`'s own leaf fields become the root dataset (`dataclass_to_dataset()`). + Each dataclass-typed child field becomes a named child node + (recursively converted the same way); a dict- or list-of-children + field expands to one child node per entry, named by its dict key or ``f"{field_name}{index}"`` respectively. A back-reference field pointing back up the tree (e.g. `Component`'s @@ -176,15 +194,15 @@ def attrs_to_datatree(obj, _ancestors: frozenset = frozenset()) -> xr.DataTree: for name, child in single_children.items(): if id(child) in ancestors: continue - children[name] = attrs_to_datatree(child, ancestors) + children[name] = dataclass_to_datatree(child, ancestors) for field_name, collection in collection_children.items(): items = collection.items() if isinstance(collection, dict) else enumerate(collection) for key, child in items: if id(child) in ancestors: continue node_name = str(key) if isinstance(collection, dict) else f"{field_name}{key}" - children[node_name] = attrs_to_datatree(child, ancestors) - return xr.DataTree(dataset=attrs_to_dataset(obj), children=children) + children[node_name] = dataclass_to_datatree(child, ancestors) + return xr.DataTree(dataset=dataclass_to_dataset(obj), children=children) def _unwrap_optional(tp): @@ -195,96 +213,99 @@ def _unwrap_optional(tp): return tp -def _child_field_spec(field: attrs.Attribute) -> "tuple[str, type] | None": - """If `field`'s declared type holds attrs-typed child/children, return - ``(kind, element_type)`` where `kind` is ``"one"``, ``"list"``, or - ``"dict"``. Returns `None` for plain scalar/array fields, an - unresolvable (string/forward-ref) annotation, or a collection whose - element type isn't a plain attrs class (e.g. a `Union` of arms -- - see module docstring). +def _child_field_spec(finfo: Any) -> "tuple[str, type] | None": + """If `finfo`'s declared type holds dataclass-typed child/children, + return ``(kind, element_type)`` where `kind` is ``"one"``, ``"list"``, + or ``"dict"``. Returns `None` for plain scalar/array fields, an + unresolvable annotation, or a collection whose element type isn't a + plain dataclass (e.g. a `Union` of arms -- see module docstring). """ - tp = field.type - if tp is None or isinstance(tp, str): + tp = finfo.annotation + if tp is None: return None tp = _unwrap_optional(tp) origin = get_origin(tp) if origin is None: - return ("one", tp) if attrs.has(tp) else None + return ("one", tp) if isinstance(tp, type) and _is_pydantic_type(tp) else None args = get_args(tp) - if origin in (list, tuple) and len(args) >= 1 and attrs.has(args[0]): + if origin in (list, tuple) and len(args) >= 1 and _is_pydantic_type(args[0]): return ("list", args[0]) - if origin is dict and len(args) == 2 and attrs.has(args[1]): + if origin is dict and len(args) == 2 and _is_pydantic_type(args[1]): return ("dict", args[1]) return None -def child_field_candidates(field: attrs.Attribute) -> "tuple[str, tuple[type, ...]] | None": +def _is_pydantic_type(tp: Any) -> bool: + return isinstance(tp, type) and is_pydantic_dataclass(tp) + + +def child_field_candidates(finfo: Any) -> "tuple[str, tuple[type, ...]] | None": """Like `_child_field_spec`, but resolves *every* concrete - attrs-decorated candidate class for the field, including each arm of - a `Union` of attrs classes in the collection-element (or bare "only") + dataclass-decorated candidate class for the field, including each arm + of a `Union` of dataclasses in the collection-element (or bare "only") position -- e.g. `list[Union[Chd, Chdg]]`, needed to disambiguate an MF6 base/grid-array package pair sharing one namefile ftype (see `converter/binding.py`'s `component_ftype()`). Returns `None` under the same conditions as `_child_field_spec`: an unresolvable annotation, or a type/collection-element that resolves to no - attrs-decorated candidate at all. + dataclass-decorated candidate at all. Kind is `"only"`, `"list"`, or `"dict"`, matching the child-collection vocabulary used throughout `flopy4/mf6/converter/`. """ - tp = field.type - if tp is None or isinstance(tp, str): + tp = finfo.annotation + if tp is None: return None tp = _unwrap_optional(tp) origin = get_origin(tp) if origin in (Union, types.UnionType): # Optional[Union[A, B]] -- _unwrap_optional only collapses a # single non-None arm, so a genuine multi-arm Union survives here. - candidates = tuple(a for a in get_args(tp) if a is not type(None) and attrs.has(a)) + candidates = tuple(a for a in get_args(tp) if a is not type(None) and _is_pydantic_type(a)) return ("only", candidates) if candidates else None if origin is None: - return ("only", (tp,)) if attrs.has(tp) else None + return ("only", (tp,)) if _is_pydantic_type(tp) else None args = get_args(tp) if origin in (list, tuple) and len(args) >= 1: elem = args[0] if get_origin(elem) in (Union, types.UnionType): - candidates = tuple(a for a in get_args(elem) if attrs.has(a)) + candidates = tuple(a for a in get_args(elem) if _is_pydantic_type(a)) return ("list", candidates) if candidates else None - return ("list", (elem,)) if attrs.has(elem) else None - if origin is dict and len(args) == 2 and attrs.has(args[1]): + return ("list", (elem,)) if _is_pydantic_type(elem) else None + if origin is dict and len(args) == 2 and _is_pydantic_type(args[1]): return ("dict", (args[1],)) return None -def datatree_to_attrs(cls: type, tree: xr.DataTree): +def datatree_to_dataclass(cls: type, tree: xr.DataTree): """Construct a `cls` instance from an `xr.DataTree` produced by - `attrs_to_datatree()`. + `dataclass_to_datatree()`. The root dataset supplies `cls`'s own scalar/array field kwargs (see - `dataset_to_attrs()`). Each attrs-typed child field is matched to + `dataset_to_dataclass()`). Each dataclass-typed child field is matched to child node(s) by name and recursively reconstructed against the field's own declared element type. See the module docstring for the "dict"-kind and `Union`-element limitations. """ kwargs = _leaf_kwargs_from_dataset(cls, tree.dataset) - for field in attrs.fields(cls): - if field.init is False: + for name, finfo in pydantic_fields(cls).items(): + if finfo.init is False: continue - spec = _child_field_spec(field) + spec = _child_field_spec(finfo) if spec is None: continue kind, elem_type = spec if kind == "one": - if field.name in tree.children: - kwargs[field.name] = datatree_to_attrs(elem_type, tree.children[field.name]) + if name in tree.children: + kwargs[name] = datatree_to_dataclass(elem_type, tree.children[name]) elif kind == "list": items = [] i = 0 - while f"{field.name}{i}" in tree.children: - items.append(datatree_to_attrs(elem_type, tree.children[f"{field.name}{i}"])) + while f"{name}{i}" in tree.children: + items.append(datatree_to_dataclass(elem_type, tree.children[f"{name}{i}"])) i += 1 if items: - kwargs[field.name] = items + kwargs[name] = items # "dict"-kind: not reconstructable from node name alone -- see # module docstring. Left unset; a caller with namefile binding # rows can fill it in separately. diff --git a/flopy4/dimensions.py b/flopy4/dimensions.py index c055d6b4..be40f1e3 100644 --- a/flopy4/dimensions.py +++ b/flopy4/dimensions.py @@ -2,7 +2,9 @@ from typing import Protocol, runtime_checkable -import attrs +from pydantic.dataclasses import is_pydantic_dataclass + +from flopy4.spec import field_meta, pydantic_fields @runtime_checkable @@ -97,21 +99,21 @@ class DimensionResolverMixin: Attributes ---------- _dimension_cache : dict[str, int] - Cache of resolved dimensions (stored as instance variable, not attrs field) + Cache of resolved dimensions (stored as instance variable, not a dataclass field) """ @property def _dimension_cache(self) -> dict: - # Lazily initialize in __dict__ directly rather than as a real attrs - # field: avoids needing a mutable-default Factory, and doesn't - # depend on __attrs_post_init__ chaining order across mixins. + # Lazily initialize in __dict__ directly rather than as a real + # dataclass field: avoids a mutable default, and doesn't depend on + # __post_init__ chaining order across mixins. if "_dimension_cache" not in self.__dict__: self.__dict__["_dimension_cache"] = {} return self.__dict__["_dimension_cache"] - def __attrs_post_init__(self) -> None: - if hasattr(super(), "__attrs_post_init__"): - super().__attrs_post_init__() # type: ignore[misc] + def __post_init__(self) -> None: + if hasattr(super(), "__post_init__"): + super().__post_init__() # type: ignore[misc] def resolve_dims(self, *dims: str) -> dict[str, int]: """ @@ -191,19 +193,19 @@ def _find_dimension_in_children(self, dim_name: str) -> int | None: def _walk_providers(self): """Yield (source_label, dims_dict) for each DimensionProvider in child fields.""" - for field_obj in attrs.fields(type(self)): # type: ignore[arg-type] - if (value := getattr(self, field_obj.name, None)) is None: + for name in pydantic_fields(type(self)): + if (value := getattr(self, name, None)) is None: continue if isinstance(value, DimensionProvider): - yield field_obj.name, value.get_dims() + yield name, value.get_dims() elif isinstance(value, dict): for child_key, child in value.items(): if isinstance(child, DimensionProvider): - yield f"{field_obj.name}[{child_key}]", child.get_dims() + yield f"{name}[{child_key}]", child.get_dims() elif isinstance(value, list): for idx, child in enumerate(value): if isinstance(child, DimensionProvider): - yield f"{field_obj.name}[{idx}]", child.get_dims() + yield f"{name}[{idx}]", child.get_dims() def _get_all_dimensions(self) -> dict[str, int]: """Get all dimensions from children and parent. Children take precedence.""" @@ -271,10 +273,11 @@ def validate_dimension_resolution(component) -> list[str]: errors = [] # Check all array fields on this component - for field in attrs.fields(type(component)): + for name, finfo in pydantic_fields(type(component)).items(): # Check if field has dimension metadata - if hasattr(field, "metadata") and field.metadata and "dims" in field.metadata: - dims_needed = field.metadata["dims"] + meta = field_meta(finfo) + if isinstance(meta, dict) and "dims" in meta: + dims_needed = meta["dims"] # Check if this component has a parent and can resolve dimensions if hasattr(component, "_parent") and component._parent: if hasattr(component._parent, "resolve_dims"): @@ -282,40 +285,26 @@ def validate_dimension_resolution(component) -> list[str]: result = component._parent.resolve_dims(dim) if dim not in result: errors.append( - f"{type(component).__name__}.{field.name} needs dimension '{dim}' " + f"{type(component).__name__}.{name} needs dimension '{dim}' " f"but it's not available in parent hierarchy" ) # Recursively validate children - for field in attrs.fields(type(component)): - value = getattr(component, field.name, None) + for name in pydantic_fields(type(component)): + value = getattr(component, name, None) if value is None: continue - # Check if child is a component with attrs fields - if hasattr(value, "__class__") and hasattr(attrs, "fields"): - try: - attrs.fields(type(value)) - # It's an attrs class, validate it - errors.extend(validate_dimension_resolution(value)) - except Exception: - # Not an attrs class, skip - pass + # Check if child is a pydantic dataclass instance + if is_pydantic_dataclass(type(value)): + errors.extend(validate_dimension_resolution(value)) elif isinstance(value, dict): for child in value.values(): - if hasattr(child, "__class__") and hasattr(attrs, "fields"): - try: - attrs.fields(type(child)) - errors.extend(validate_dimension_resolution(child)) - except Exception: - pass + if is_pydantic_dataclass(type(child)): + errors.extend(validate_dimension_resolution(child)) elif isinstance(value, list): for child in value: - if hasattr(child, "__class__") and hasattr(attrs, "fields"): - try: - attrs.fields(type(child)) - errors.extend(validate_dimension_resolution(child)) - except Exception: - pass + if is_pydantic_dataclass(type(child)): + errors.extend(validate_dimension_resolution(child)) return errors diff --git a/flopy4/mf6/__init__.py b/flopy4/mf6/__init__.py index e008756d..b9300d9a 100644 --- a/flopy4/mf6/__init__.py +++ b/flopy4/mf6/__init__.py @@ -63,7 +63,7 @@ def _load_mf6(cls, path: Path, name: "str | None" = None) -> Component: instance = structure_component(raw, cls, workspace=path.parent, name=name) if isinstance(instance, Context): instance.workspace = path.parent - instance.filename = path.name + instance.filename = Path(path.name) return instance diff --git a/flopy4/mf6/_types.py b/flopy4/mf6/_types.py index 8eb97884..f47ceaf2 100644 --- a/flopy4/mf6/_types.py +++ b/flopy4/mf6/_types.py @@ -2,7 +2,7 @@ from datetime import datetime from pathlib import Path -from typing import Protocol, TypeAlias, TypeVar +from typing import Protocol, TypeAlias, TypeVar, runtime_checkable import numpy as np @@ -15,6 +15,7 @@ unions support this natively).""" +@runtime_checkable class _ArrayLike(Protocol[_DT]): """Structural stand-in for "ndarray or duck array of this dtype". @@ -24,6 +25,15 @@ class _ArrayLike(Protocol[_DT]): Used for griddata and READARRAY period fields, which may be dask-backed (see `codec/writer/filters.py`'s `array2chunks`, which streams dask-backed arrays without materializing them). + + `@runtime_checkable` is required: under + `arbitrary_types_allowed=True`, pydantic builds an `isinstance()`-based + validator for any type it doesn't otherwise understand, which requires + the protocol to support `isinstance()` at all -- confirmed empirically + that schema-building itself fails with a `SchemaError` (not even a + runtime `ValidationError`) without this decorator, even though the + generic type parameter (`_DT`) is itself ignored by the resulting + isinstance check either way, same as plain Python `Protocol` semantics. """ @property @@ -37,7 +47,7 @@ def shape(self) -> tuple[int, ...]: ... def _optional_path(v): - """Converter for Optional[Path] attrs fields. + """Converter for Optional[Path] fields. Accepts None, str, or Path; returns None or Path. """ diff --git a/flopy4/mf6/adapters.py b/flopy4/mf6/adapters.py index 3a3bc0af..7aa9ad71 100644 --- a/flopy4/mf6/adapters.py +++ b/flopy4/mf6/adapters.py @@ -3,7 +3,6 @@ from typing import Optional, Union, get_args, get_origin from warnings import warn -import attrs import numpy as np from flopy.datbase import DataInterface, DataListInterface, DataType from flopy.discretization.grid import Grid @@ -12,10 +11,12 @@ from flopy.mbase import ModelInterface from flopy.pakbase import PackageInterface from flopy.plot.plotutil import PlotUtilities +from pydantic.dataclasses import is_pydantic_dataclass -from flopy4.attrs_xarray import attrs_to_dataset +from flopy4.dataclass_xarray import dataclass_to_dataset from flopy4.mf6.model import Model from flopy4.mf6.package import Package +from flopy4.spec import field_meta, pydantic_fields def _to_numpy(val): @@ -35,7 +36,7 @@ def _resolve_leaf_type(annotation) -> "type | None": `NDArray[np.float64]`) down to the concrete runtime type `Flopy3Data` dispatches on (`bool`/`int`/`float`/`str`/`np.ndarray`). - Returns `None` for anything else (a nested attrs/Component type, + Returns `None` for anything else (a nested Component type, `Path`, `datetime`, `Record`, a bare `dict`/`list` period field, ...) -- those aren't representable as a single flopy3 `Data` leaf. """ @@ -207,7 +208,7 @@ def __init__( ): self._model = model self._package = package - self._dataset = attrs_to_dataset(package) + self._dataset = dataclass_to_dataset(package) if modelgrid: self._grid = modelgrid elif model: @@ -217,13 +218,13 @@ def __init__( self._time = modeltime self._dlist = list() - field_by_name = {f.name: f for f in attrs.fields(type(package))} + field_by_name = dict(pydantic_fields(type(package))) for a, value in self._dataset.attrs.items(): field = field_by_name.get(a) if field is None or value is None: continue - leaf_type = _resolve_leaf_type(field.type) + leaf_type = _resolve_leaf_type(field.annotation) if leaf_type is None: continue d_fp3 = Flopy3Data( @@ -238,7 +239,7 @@ def __init__( self._dlist.append(d_fp3) for v, data_array in self._dataset.data_vars.items(): - field = field_by_name.get(v) + field = field_by_name.get(str(v)) if field is None: continue d_fp3 = Flopy3Data( @@ -291,14 +292,13 @@ def has_stress_period_data(self): # Any other fill-forward (period) field (covers OC's own # _stress_period_data too, redundantly with the check above -- kept # as a generic fallback for any period field shape). - try: - for f in attrs.fields(type(self._package)): - if f.metadata.get("fill_forward"): - attr_name = f.alias if (f.alias and f.name.startswith("_")) else f.name + if is_pydantic_dataclass(type(self._package)): + for name, f in pydantic_fields(type(self._package)).items(): + meta = field_meta(f) + if isinstance(meta, dict) and meta.get("fill_forward"): + attr_name = f.alias if (f.alias and name.startswith("_")) else name if getattr(self._package, attr_name, None) is not None: return True - except attrs.exceptions.NotAnAttrsClassError: - pass return "nper" in self._dataset.dims def check(self, f=None, verbose=True, level=1, checktype=None): diff --git a/flopy4/mf6/codec/writer/filters.py b/flopy4/mf6/codec/writer/filters.py index 3a0f545f..7f832a54 100644 --- a/flopy4/mf6/codec/writer/filters.py +++ b/flopy4/mf6/codec/writer/filters.py @@ -1,8 +1,9 @@ +import dataclasses from collections.abc import Hashable, Mapping from io import StringIO +from pathlib import PurePath from typing import Any, Literal -import attrs import numpy as np import xarray as xr from numpy.typing import NDArray @@ -150,7 +151,7 @@ def array2string(value: NDArray, precision: int = 9) -> str: return buffer.getvalue().strip() -def quote_if_needed(value: str) -> str: +def quote_if_needed(value: Any) -> str: """ Wrap a string in single quotes if it contains double-quotes. @@ -158,7 +159,12 @@ def quote_if_needed(value: str) -> str: contain double-quotes and must be single-quoted for MF6 to parse them. MF6 keyword sequences like 'STEPS 1 5' or 'all' are left as-is even if they contain spaces, because they are not string literals. + + Paths are written with POSIX separators so input files are portable + across platforms. """ + if isinstance(value, PurePath): + return value.as_posix() if isinstance(value, str) and '"' in value: return f"'{value}'" return str(value) @@ -267,7 +273,7 @@ def dataset2list(value: xr.Dataset): if name == "perioddata": val = value[name] val = val.item() if val.shape == () else val - yield attrs.astuple(val, recurse=True) # type: ignore + yield dataclasses.astuple(val) # type: ignore continue val = value[name] val = val.item() if val.shape == () else val diff --git a/flopy4/mf6/codec/writer/templates/macros.jinja b/flopy4/mf6/codec/writer/templates/macros.jinja index d4fe8e78..1dcfb0e1 100644 --- a/flopy4/mf6/codec/writer/templates/macros.jinja +++ b/flopy4/mf6/codec/writer/templates/macros.jinja @@ -53,7 +53,7 @@ {{ (2 * inset) ~ chunk|array2string(context.float_precision) }} {%- endfor %} {% elif how == "external" %} -OPEN/CLOSE {{ value }} +OPEN/CLOSE {{ value|quote_if_needed }} {% endif %} {% endif %} {% endmacro %} diff --git a/flopy4/mf6/component.py b/flopy4/mf6/component.py index 53409fa3..d4709a45 100644 --- a/flopy4/mf6/component.py +++ b/flopy4/mf6/component.py @@ -1,18 +1,29 @@ +import dataclasses from abc import ABC from collections.abc import MutableMapping from os import PathLike from pathlib import Path from typing import Any, ClassVar, Optional -import attrs -from attrs import fields +from pydantic import ConfigDict, Field, field_validator +from pydantic.dataclasses import dataclass, is_pydantic_dataclass from flopy4.dimensions import DimensionResolverMixin from flopy4.mf6.constants import MF6 -from flopy4.mf6.spec import field, fields_dict +from flopy4.mf6.spec import fields_dict from flopy4.mf6.write_context import WriteContext +from flopy4.spec import field_meta, pydantic_fields from flopy4.uio import IO, Loader, Writer +# Shared config for every Component/Package (sub)class. Pydantic doesn't +# inherit dataclass config, so each class passes it explicitly at its own +# `@dataclass(config=CFG, ...)` decoration site, as codegen does. +CFG = ConfigDict( + arbitrary_types_allowed=True, + validate_assignment=True, + extra="forbid", +) + FNAMES: "dict[str, type[Component]]" = {} """MF6 component name (e.g. 'gwf-dis') -> component class.""" @@ -67,7 +78,7 @@ def _is_default_child_name(child: "Component") -> bool: """Whether `child`'s current `.name` is still at its class-name default (see `Component.name`'s own field docstring), i.e. no explicit name was ever given.""" - return child.name == type(child).__name__.lower() # type: ignore[attr-defined] + return child.name == type(child).__name__.lower() def _resolve_child_name(used: "set[str]", kind: str, field_name: str, child: "Component") -> str: @@ -88,12 +99,11 @@ def _resolve_child_name(used: "set[str]", kind: str, field_name: str, child: "Co if kind not in ("only", "list"): raise TypeError(f"Bad child collection kind '{kind}'") if not _is_default_child_name(child): - if child.name in used: # type: ignore[attr-defined] + if child.name in used: raise ValueError( - f"Child name '{child.name}' collides with an existing child " # type: ignore[attr-defined] - "on the same parent." + f"Child name '{child.name}' collides with an existing child on the same parent." ) - return child.name # type: ignore[attr-defined] + return child.name if kind == "only": return field_name i = 0 @@ -106,33 +116,34 @@ def _find_child_field(parent_cls: type, child_cls: type) -> "tuple[Any, str] | N """Find the single field on `parent_cls` that accepts `child_cls` as a child, by type annotation (`child_field_candidates()`). - Returns `(field, kind)`, or `None` if no field matches. Raises + Returns `(finfo, kind)`, or `None` if no field matches. Raises `TypeError` if more than one field matches (ambiguous). """ - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates matches = [] - for f in fields(parent_cls): # type: ignore[arg-type] - spec = child_field_candidates(f) + for name, finfo in pydantic_fields(parent_cls).items(): + spec = child_field_candidates(finfo) if spec is None: continue kind, candidates = spec if any(issubclass(child_cls, c) for c in candidates): - matches.append((f, kind)) + matches.append((name, finfo, kind)) if not matches: return None if len(matches) > 1: - names = ", ".join(f.name for f, _ in matches) + names = ", ".join(name for name, _, _ in matches) raise TypeError( f"Class '{parent_cls.__name__}' has multiple fields of type " f"'{child_cls.__name__}' ({names}); can't bind." ) - return matches[0] + name, _finfo, kind = matches[0] + return name, kind # kw_only=True necessary so we can define optional fields here -# and required fields in subclasses. attrs complains otherwise -@attrs.define(kw_only=True, slots=False) +# and required fields in subclasses. +@dataclass(config=CFG, kw_only=True) class Component(DimensionResolverMixin, ABC, MutableMapping): """ Base class for MF6 components. @@ -150,41 +161,65 @@ class Component(DimensionResolverMixin, ABC, MutableMapping): _load = IO(Loader) # type: ignore _write = IO(Writer) # type: ignore - filename: str | None = field(default=None) - """The name of the component's input file.""" - - name: str = field( - default=attrs.Factory(lambda self: type(self).__name__.lower(), takes_self=True) + filename: Optional[Path] = Field(default=None) + """The component's input file, relative to the workspace (a `str` is + accepted and converted). Written to name files with POSIX separators.""" + + name: str = Field(default="", validate_default=True) + """The component's own identity/tag name. Defaults to the *actual* + runtime class's lowercased name -- not whichever class in the hierarchy + happens to declare this field -- so a `Package` leaf (e.g. `Ic`, never + separately subclassed for this field) still gets "ic", not "package". + See `_default_name`. Overridden explicitly by + `_resolve_child_name()`/`_attach_to_parent_field()` when a component is + attached as a named child; otherwise this default stands.""" + + _parent: Any = dataclasses.field( + default=Field(default=None, alias="parent", repr=False), compare=False ) - """The component's own identity/tag name. Computed per-instance from - the *actual* runtime class (`takes_self=True`), not whichever class in - the hierarchy happens to declare this field -- so a `Package` leaf - (e.g. `Ic`, never separately subclassed for this field) still gets - "ic", not "package". Overridden explicitly by `_resolve_child_name()`/ - `_attach_to_parent_field()` when a component is attached as a named - child; otherwise this default stands.""" - - _parent: Any = field(default=None, repr=False, eq=False) """Parent back-reference -- source of truth for "who is this component's parent", top-down (`Gwf(dis=Dis(...))`) and bottom-up - (`Dis(parent=gwf)`) alike. Leading underscore triggers attrs' private- - attribute convention, so the constructor keyword stays `parent=` even - though the field is `_parent`. Typed `Any` so `child_field_candidates()` - (type-annotation based) doesn't mistake it for a real child field. + (`Dis(parent=gwf)`) alike. The alias keeps the constructor keyword + `parent=` even though the field is `_parent`. Typed `Any` so + `child_field_candidates()` (type-annotation based) doesn't mistake it + for a real child field, and `compare=False` because comparing a live + `.parent` would recurse: comparing a component's parent compares the + parent's own children, including this component again. Populated by `_set_child_parents()` (top-down) and `_attach_to_parent_field()` (bottom-up), and kept current by - `parent`'s setter below. Excluded by name from `to_dict()`'s - `attrs.asdict()` recursion (alongside the unrelated `Output.parent`) - since a live `.parent` would otherwise be a reference cycle. + `parent`'s setter below. Excluded by name from `to_dict()`'s recursion + (alongside the unrelated `Output.parent`) since a live `.parent` would + otherwise be a reference cycle. """ - dims: dict = field(default=attrs.Factory(dict), repr=False, eq=False) - """Accepts `dims=` at construction (e.g. `Ic(dims={"nodes": 900})`) - for API-compatibility with existing call sites. Read directly via - `self.__dict__.get("dims")` by `Package.__attrs_post_init__` for - griddata broadcasting -- not resolved/consumed by anything at the - `Component` level itself.""" + dims: dataclasses.InitVar[Optional[dict]] = None + """Dimension sizes to size griddata with at construction (e.g. + `Ic(dims={"nodes": 900})`). Construction-only: passed through the + `__post_init__` chain to `Package.__post_init__`, which broadcasts + scalar griddata to full shape, and not stored on the instance.""" + + @field_validator("name", mode="before") + @classmethod + def _default_name(cls, v: Any) -> str: + """Default `name` to the runtime class's lowercased name. `cls` is + the class actually being constructed, so this is subclass-aware; + `validate_default=True` makes it run when `name` isn't given.""" + return v or cls.__name__.lower() + + @field_validator("*", mode="before") + @classmethod + def _apply_converter(cls, v: Any, info) -> Any: + """Apply each field's `converter=`, if any. Pydantic has no + per-field converter hook, so `flopy4.mf6.spec.field()`/`path()` + stash the callable in `json_schema_extra["converter"]` and this + single validator applies it for every field of every class.""" + finfo = pydantic_fields(cls).get(info.field_name) + if finfo is None or v is None: + return v + meta = field_meta(finfo) + conv = meta.get("converter") if isinstance(meta, dict) else None + return conv(v) if conv is not None else v @property def parent(self) -> "Component | None": @@ -216,7 +251,7 @@ def parent(self, value: "Component | None") -> None: if old is value: return if old is not None: - del old[self.name] # type: ignore[attr-defined] + del old[self.name] self._parent = None if value is not None: self._parent = value @@ -242,18 +277,18 @@ def _children(self) -> "dict[str, Component]": (`write()`, `NetCDFModel.from_model()`, ...), so nothing can read a stale name depending on call order. - Detects child fields via `flopy4.attrs_xarray.child_field_candidates()`. + Detects child fields via `flopy4.dataclass_xarray.child_field_candidates()`. """ - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates self._set_child_parents() result: "dict[str, Component]" = {} - for f in fields(type(self)): - spec = child_field_candidates(f) + for name, finfo in pydantic_fields(type(self)).items(): + spec = child_field_candidates(finfo) if spec is None: continue - value = getattr(self, f.name, None) + value = getattr(self, name, None) if value is None: continue kind, _ = spec @@ -280,14 +315,14 @@ def _set_child_parents(self) -> None: Detects child fields via `child_field_candidates()` (type- annotation based). """ - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates used: "set[str]" = set() - for f in fields(type(self)): - spec = child_field_candidates(f) + for name, finfo in pydantic_fields(type(self)).items(): + spec = child_field_candidates(finfo) if spec is None: continue - value = getattr(self, f.name, None) + value = getattr(self, name, None) if value is None: continue kind, _ = spec @@ -295,14 +330,14 @@ def _set_child_parents(self) -> None: if kind == "only": if isinstance(value, Component): value.__dict__["_parent"] = self - value.name = _resolve_child_name(used, kind, f.name, value) # type: ignore[attr-defined] - used.add(value.name) # type: ignore[attr-defined] + value.name = _resolve_child_name(used, kind, name, value) + used.add(value.name) elif kind == "list": for child in value: if isinstance(child, Component): child.__dict__["_parent"] = self - child.name = _resolve_child_name(used, kind, f.name, child) # type: ignore[attr-defined] - used.add(child.name) # type: ignore[attr-defined] + child.name = _resolve_child_name(used, kind, name, child) + used.add(child.name) elif kind == "dict": for key, child in value.items(): if isinstance(child, Component): @@ -312,13 +347,13 @@ def _set_child_parents(self) -> None: f"Child name '{key}' collides with an " "existing child on the same parent." ) - child.name = key # type: ignore[attr-defined] - used.add(child.name) # type: ignore[attr-defined] + child.name = key + used.add(child.name) @property def path(self) -> Path: """The path to the component's input file.""" - self.filename = self.filename or self.default_filename() + self.filename = self.filename or Path(self.default_filename()) return Path.cwd() / self.filename def default_filename(self) -> str: @@ -334,22 +369,24 @@ def default_filename(self) -> str: cls_name = self.__class__.__name__.lower() return f"{name}.{cls_name}" - def __attrs_post_init__(self): + def __post_init__(self, dims: Optional[dict] = None): """ Post-initialization hook for all components. Chains to parent class post-init hooks (including DimensionRegistryMixin). - Also runs the two `_parent`-tracking hooks (see `_parent`'s - docstring): stamps `_parent` on this component's own already- - populated children (top-down construction), and -- if this - component's own `_parent` was given directly as `parent=`, i.e. - bottom-up construction -- attaches `self` into the matching field - on it and resolves its `.name` -- see `_attach_to_parent_field()`'s - docstring. + `dims` is the `dims` InitVar; only `Package` uses it, so it isn't + passed further up the chain. + + Runs the two `_parent`-tracking hooks (see `_parent`'s docstring): + stamps `_parent` on this component's own already-populated children + (top-down construction), and -- if this component's own `_parent` + was given directly as `parent=`, i.e. bottom-up construction -- + attaches `self` into the matching field on it and resolves its + `.name` -- see `_attach_to_parent_field()`'s docstring. """ # Chain to parent classes (including DimensionRegistryMixin) - if hasattr(super(), "__attrs_post_init__"): - super().__attrs_post_init__() # type: ignore[misc] + if hasattr(super(), "__post_init__"): + super().__post_init__() # type: ignore[misc] if self._parent is not None: self._attach_to_parent_field(self._parent) self._set_child_parents() @@ -369,28 +406,28 @@ def _attach_to_parent_field(self, parent: "Component") -> None: match = _find_child_field(type(parent), type(self)) if match is None: return - target_field, kind = match - used = {c.name for c in parent._children.values()} # type: ignore[attr-defined] + target_name, kind = match + used = {c.name for c in parent._children.values()} if kind == "only": - self.name = _resolve_child_name(used, kind, target_field.name, self) # type: ignore[attr-defined] - setattr(parent, target_field.name, self) + self.name = _resolve_child_name(used, kind, target_name, self) + setattr(parent, target_name, self) elif kind == "list": - self.name = _resolve_child_name(used, kind, target_field.name, self) # type: ignore[attr-defined] - getattr(parent, target_field.name).append(self) + self.name = _resolve_child_name(used, kind, target_name, self) + getattr(parent, target_name).append(self) elif kind == "dict": # No positional auto-key to fall back on for an unnamed child, # unlike "only"/"list" -- see `_set_child_parents`'s "dict" # branch: the child's own `.name` (explicit, or its # class-name default) is the key. - key = self.name # type: ignore[attr-defined] + key = self.name if key in used: raise ValueError( f"Child name '{key}' collides with an existing child on the same parent." ) - getattr(parent, target_field.name)[key] = self + getattr(parent, target_name)[key] = self - @classmethod - def __attrs_init_subclass__(cls): + def __init_subclass__(cls, **kwargs): + super().__init_subclass__(**kwargs) # Only register classes that declare their own `dfn_name`. # Abstract bases (Package, Context, Model, Exchange, Solution, # DisBase, ...) have no `dfn_name` of their own and are silently @@ -416,30 +453,30 @@ def __setitem__(self, key, value): if not isinstance(value, Component): raise TypeError(f"Expected a Component, got {type(value).__name__}") - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates - for f in fields(type(self)): - spec = child_field_candidates(f) + for name, finfo in pydantic_fields(type(self)).items(): + spec = child_field_candidates(finfo) if spec is None: continue kind, _ = spec - current = getattr(self, f.name, None) + current = getattr(self, name, None) if kind == "only": - if isinstance(current, Component) and current.name == key: # type: ignore[attr-defined] - value.name = key # type: ignore[attr-defined] + if isinstance(current, Component) and current.name == key: + value.name = key value.__dict__["_parent"] = self - setattr(self, f.name, value) + setattr(self, name, value) return elif kind == "list": for i, child in enumerate(current or []): - if isinstance(child, Component) and child.name == key: # type: ignore[attr-defined] - value.name = key # type: ignore[attr-defined] + if isinstance(child, Component) and child.name == key: + value.name = key value.__dict__["_parent"] = self current[i] = value return elif kind == "dict": if current and key in current: - value.name = key # type: ignore[attr-defined] + value.name = key value.__dict__["_parent"] = self current[key] = value return @@ -447,34 +484,34 @@ def __setitem__(self, key, value): match = _find_child_field(type(self), type(value)) if match is None: raise TypeError(f"No field on {type(self).__name__} accepts a {type(value).__name__}") - target_field, kind = match + target_name, kind = match value.__dict__["_parent"] = self - value.name = key # type: ignore[attr-defined] + value.name = key if kind == "only": - setattr(self, target_field.name, value) + setattr(self, target_name, value) elif kind == "list": - getattr(self, target_field.name).append(value) + getattr(self, target_name).append(value) elif kind == "dict": - getattr(self, target_field.name)[key] = value + getattr(self, target_name)[key] = value def __delitem__(self, key): """Detach the child named `key`, from whatever field/slot currently holds it.""" - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates - for f in fields(type(self)): - spec = child_field_candidates(f) + for name, finfo in pydantic_fields(type(self)).items(): + spec = child_field_candidates(finfo) if spec is None: continue kind, _ = spec - value = getattr(self, f.name, None) + value = getattr(self, name, None) if kind == "only": - if isinstance(value, Component) and value.name == key: # type: ignore[attr-defined] - setattr(self, f.name, None) + if isinstance(value, Component) and value.name == key: + setattr(self, name, None) return elif kind == "list": for i, child in enumerate(value or []): - if isinstance(child, Component) and child.name == key: # type: ignore[attr-defined] + if isinstance(child, Component) and child.name == key: del value[i] return elif kind == "dict": @@ -518,7 +555,7 @@ def write(self, format: str = MF6, context: Optional[WriteContext] = None) -> No # name as this component's filename stem, if it has one. an # actual solution is to auto-set the filename when children # are attached to parents. - self.filename = self.filename or self.default_filename() + self.filename = self.filename or Path(self.default_filename()) # Determine active context: provided > current > default active_context = context or WriteContext.current() @@ -527,6 +564,30 @@ def write(self, format: str = MF6, context: Optional[WriteContext] = None) -> No for child in self._children.values(): child.write(format=format, context=context) + def _asdict_filtered(self) -> dict[str, Any]: + """Recursive `dataclasses.asdict()`, excluding any field literally + named "parent" or "_parent" at every recursion level, not just + this component's own: e.g. `Gwf.Output.parent` is a genuine + back-reference to the owning `Gwf`, unrelated to `Component. + _parent`, but recursing into it the same way would infinitely + loop (output -> parent -> output -> ...). `dataclasses.asdict()` + has no filter hook, so this walks by hand instead.""" + + def _convert(value: Any) -> Any: + if is_pydantic_dataclass(type(value)): + return { + name: _convert(getattr(value, name)) + for name in pydantic_fields(type(value)) + if name not in ("parent", "_parent") + } + if isinstance(value, dict): + return {k: _convert(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return type(value)(_convert(v) for v in value) + return value + + return _convert(self) + def to_dict(self, blocks: bool = False, strict: bool = False) -> dict[str, Any]: """ Convert the component to a dictionary representation. @@ -545,14 +606,7 @@ def to_dict(self, blocks: bool = False, strict: bool = False) -> dict[str, Any]: Dictionary containing component data, either in terms of fields (flat) or blocks (nested). """ - # Exclude any field literally named "parent" or "_parent" at every - # recursion level, not just this component's own: e.g. - # Gwf.Output.parent is a genuine back-reference to the owning Gwf, - # unrelated to Component._parent, but recursing into it the same - # way would infinitely loop (output -> parent -> output -> ...). - data = attrs.asdict( - self, recurse=True, filter=lambda attr, value: attr.name not in ("parent", "_parent") - ) + data = self._asdict_filtered() spec = fields_dict(self.__class__) if strict: @@ -561,9 +615,10 @@ def to_dict(self, blocks: bool = False, strict: bool = False) -> dict[str, Any]: if blocks: blocks_ = {} # type: ignore - for field_name, field_attr in spec.items(): + for field_name, finfo in spec.items(): field_value = data[field_name] - block_name = field_attr.metadata.get("block") + meta = field_meta(finfo) + block_name = meta.get("block") if isinstance(meta, dict) else None if strict and block_name is None: continue if block_name not in blocks_: @@ -573,22 +628,22 @@ def to_dict(self, blocks: bool = False, strict: bool = False) -> dict[str, Any]: else: return { field_name: data[field_name] - for field_name, field_attr in spec.items() - if field_attr.metadata.get("block") or not strict + for field_name, finfo in spec.items() + if field_meta(finfo).get("block") or not strict } def to_xarray(self): """Flat xr.Dataset of this component's own scalar/array fields, merged with any child packages that have griddata fields. - Built directly from live attribute values via flopy4.attrs_xarray's - attrs_to_dataset. + Built directly from live attribute values via flopy4.dataclass_xarray's + dataclass_to_dataset. """ import xarray as _xr - from flopy4.attrs_xarray import attrs_to_dataset + from flopy4.dataclass_xarray import dataclass_to_dataset - base = attrs_to_dataset(self) + base = dataclass_to_dataset(self) extra = list(self._collect_child_griddata_datasets().values()) if not extra: return base @@ -604,11 +659,10 @@ def _collect_child_griddata_datasets(self) -> dict: result: dict = {} try: for name, child in self._children.items(): - try: - _fields = attrs.fields(type(child)) - except attrs.exceptions.NotAnAttrsClassError: + if not is_pydantic_dataclass(type(child)): continue - if not any(f.metadata.get("block") == "griddata" for f in _fields): + _fields = pydantic_fields(type(child)) + if not any(field_meta(f).get("block") == "griddata" for f in _fields.values()): continue try: ds = child.to_xarray() diff --git a/flopy4/mf6/context.py b/flopy4/mf6/context.py index d085dcb4..5a4111d7 100644 --- a/flopy4/mf6/context.py +++ b/flopy4/mf6/context.py @@ -1,55 +1,56 @@ from abc import ABC from pathlib import Path +from typing import Any, Optional -import attrs from modflow_devtools.misc import cd +from pydantic import Field +from pydantic.dataclasses import dataclass -from flopy4.mf6.component import Component +from flopy4.mf6.component import CFG, Component from flopy4.mf6.constants import MF6 -from flopy4.mf6.spec import field from flopy4.utils import to_path -def update_child_attr(instance, attribute, new_value): - """ - Generalized function to update child attribute (e.g. workspace). - - Args: - instance: The model instance - attribute: The attribute being set (from attrs on_setattr) - new_value: The new value being set - - Returns: - The new_value (unchanged) - """ - - for child in instance._children.values(): - if hasattr(child, attribute.name): - setattr(child, attribute.name, new_value) - - return new_value - - -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Context(Component, ABC): - workspace: Path = field(default=None, converter=to_path, on_setattr=update_child_attr) + # `_workspace`/`workspace` mirrors `Component._parent`/`.parent`'s + # private-field-plus-property pattern: pydantic has no per-field + # on-setattr hook, so propagating to children happens in the property + # setter below. + _workspace: Any = Field(default=None, alias="workspace", repr=False) - def __attrs_post_init__(self): - super().__attrs_post_init__() - # By the time this runs, `super().__attrs_post_init__()` - # (Component's) has already resolved `_parent`/`.parent` for both - # top-down and bottom-up construction (see `Component._parent`'s - # docstring). - if self.workspace is None: + @property + def workspace(self) -> Path: + """The directory this context's files live in. Given at + construction, or resolved by `__post_init__` from the parent's + workspace or the current directory.""" + if self._workspace is None: + raise RuntimeError(f"{type(self).__name__}.workspace is not resolved yet") + return self._workspace + + @workspace.setter + def workspace(self, value) -> None: + """Coerce `value` to a `Path`, then propagate it to every child + `Context`.""" + value = to_path(value) + self._workspace = value + for child in self._children.values(): + if isinstance(child, Context): + child.workspace = value + + def __post_init__(self, dims: Optional[dict] = None): + super().__post_init__(dims) + # By the time this runs, `super().__post_init__()` (Component's) + # has already resolved `_parent`/`.parent` for both top-down and + # bottom-up construction (see `Component._parent`'s docstring). + if self._workspace is None: self.workspace = ( - self._parent.workspace - if self._parent and hasattr(self._parent, "workspace") - else Path.cwd() + self._parent.workspace if isinstance(self._parent, Context) else Path.cwd() ) @property def path(self) -> Path: - self.filename = self.filename or self.default_filename() + self.filename = self.filename or Path(self.default_filename()) return self.workspace / self.filename @classmethod @@ -69,10 +70,10 @@ def write(self, format=MF6, context=None): def to_xarray(self): """DataTree for this context and its full child hierarchy. - Built directly from live attribute values via flopy4.attrs_xarray's - attrs_to_datatree. Each child node, leaf packages included, is + Built directly from live attribute values via flopy4.dataclass_xarray's + dataclass_to_datatree. Each child node, leaf packages included, is built from its own fields directly, so griddata appears natively. """ - from flopy4.attrs_xarray import attrs_to_datatree + from flopy4.dataclass_xarray import dataclass_to_datatree - return attrs_to_datatree(self) + return dataclass_to_datatree(self) diff --git a/flopy4/mf6/converter/__init__.py b/flopy4/mf6/converter/__init__.py index 9c0803bf..8f0b6b6d 100644 --- a/flopy4/mf6/converter/__init__.py +++ b/flopy4/mf6/converter/__init__.py @@ -34,7 +34,7 @@ def structure(data: dict[str, Any], path: Path) -> Component: component = COMPONENT_CONVERTER.structure(data, Component) if isinstance(component, Context): component.workspace = path.parent - component.filename = path.name + component.filename = Path(path.name) return component diff --git a/flopy4/mf6/converter/binding.py b/flopy4/mf6/converter/binding.py index ab67408e..c350276e 100644 --- a/flopy4/mf6/converter/binding.py +++ b/flopy4/mf6/converter/binding.py @@ -1,4 +1,7 @@ -from attrs import define +from pathlib import Path + +from pydantic import ConfigDict +from pydantic.dataclasses import dataclass from flopy4.mf6.component import Component from flopy4.mf6.exchange import Exchange @@ -19,7 +22,10 @@ def component_ftype(cls: type) -> str: return f"{cls_name.upper()}6" -@define +_CFG = ConfigDict(arbitrary_types_allowed=True, extra="forbid") + + +@dataclass(config=_CFG) class Binding: """A serializable representation of a component.""" @@ -46,6 +52,6 @@ def _get_binding_terms(component: Component) -> tuple[str, ...] | None: return cls( type=component_ftype(type(component)), - fname=component.filename or component.default_filename(), + fname=Path(component.filename or component.default_filename()).as_posix(), terms=_get_binding_terms(component), ) diff --git a/flopy4/mf6/converter/egress/unstructure.py b/flopy4/mf6/converter/egress/unstructure.py index 7657b067..f7b02838 100644 --- a/flopy4/mf6/converter/egress/unstructure.py +++ b/flopy4/mf6/converter/egress/unstructure.py @@ -1,13 +1,13 @@ from collections.abc import Iterable, Mapping from datetime import datetime from pathlib import Path -from typing import Any +from typing import Any, cast -import attrs import numpy as np import xarray as xr +from pydantic.fields import FieldInfo -from flopy4.attrs_xarray import child_field_candidates +from flopy4.dataclass_xarray import child_field_candidates from flopy4.mf6.component import Component from flopy4.mf6.constants import FILL_DNODATA from flopy4.mf6.context import Context @@ -16,19 +16,22 @@ from flopy4.mf6.package import Package from flopy4.mf6.record import Record from flopy4.mf6.spec import FileDirection, block_sort_key, blocks_dict, to_field_type +from flopy4.spec import field_meta, pydantic_fields -def _path_to_tuple(field: attrs.Attribute, value: Path) -> tuple[str, ...]: +def _path_to_tuple(name: str, field: FieldInfo, value: Path) -> tuple[str, ...]: """A block-level file record's ``KEYWORD FILEIN|FILEOUT `` row. The keyword is the field's own ``_keyword`` metadata (see spec.path).""" - keyword = field.metadata.get("_keyword") + meta = field_meta(field) + assert isinstance(meta, dict) + keyword = meta.get("_keyword") if not keyword: - raise ValueError(f"file field {field.name!r} has no _keyword metadata") - direction: FileDirection | None = field.metadata.get("direction") - t = [keyword.upper()] + raise ValueError(f"file field {name!r} has no _keyword metadata") + direction = cast(FileDirection | None, meta.get("direction")) + t = [str(keyword).upper()] if direction: t.append("FILEOUT" if direction == "out" else "FILEIN") - t.append(str(value)) + t.append(value.as_posix()) return tuple(t) @@ -38,13 +41,13 @@ def _make_binding_blocks(value: Component) -> dict[str, dict[str, list[tuple[str blocks = {} # type: ignore - for f in attrs.fields(type(value)): # type: ignore[arg-type] + for child_name, f in pydantic_fields(type(value)).items(): if child_field_candidates(f) is None: continue - child_name = f.name if (child := getattr(value, child_name, None)) is None: continue - block_name = f.metadata.get("block") + meta = field_meta(f) + block_name = meta.get("block") if isinstance(meta, dict) else None if block_name is None: continue if block_name not in blocks: @@ -129,9 +132,9 @@ def _unstructure_package(value: Package) -> dict[str, Any]: except ImportError: _DaskArray = type(None) # type: ignore[misc,assignment] - for f in attrs.fields(cls): # type: ignore[arg-type] - meta = f.metadata - block_name = meta.get("block") + for name, f in pydantic_fields(cls).items(): + meta = field_meta(f) + block_name = meta.get("block") if isinstance(meta, dict) else None if not block_name: continue @@ -139,12 +142,12 @@ def _unstructure_package(value: Package) -> dict[str, Any]: write_if_empty_set.add(block_name) blocks.setdefault(block_name, {}) - attr_name = f.alias if (f.alias and f.name.startswith("_")) else f.name + attr_name = f.alias if (f.alias and name.startswith("_")) else name field_value = getattr(value, attr_name, None) if field_value is None: continue - dfn_type = to_field_type(f.type) + dfn_type = to_field_type(f.annotation) # fill-forward block if meta.get("fill_forward"): @@ -154,7 +157,7 @@ def _unstructure_package(value: Package) -> dict[str, Any]: is_layered = meta.get("layered", False) nper = field_value.shape[0] # aux field: shape (nper, ncpl, naux) - if f.name == "aux" and field_value.ndim == 3: + if name == "aux" and field_value.ndim == 3: aux_names: list[str] = list(getattr(value, "auxiliary", None) or []) naux = field_value.shape[2] for kper in range(nper): @@ -170,7 +173,7 @@ def _unstructure_package(value: Package) -> dict[str, Any]: da = xr.DataArray(layer_slice, dims=("nlay",) + extra_dims) else: da = xr.DataArray(layer_slice) - readarray_period.setdefault(kper, {})[f.name] = da + readarray_period.setdefault(kper, {})[name] = da continue # list: dict[int, list[Item]] if not isinstance(field_value, dict): @@ -188,7 +191,7 @@ def _unstructure_package(value: Package) -> dict[str, Any]: continue if isinstance(field_value, dict): - array_key = f.name if meta.get("tagged") else "" + array_key = name if meta.get("tagged") else "" for tval, arr in field_value.items(): blocks[f"{block_name} {tval}"] = {array_key: _wrap_array(arr)} continue @@ -199,17 +202,17 @@ def _unstructure_package(value: Package) -> dict[str, Any]: if dfn_type == "keyword": if field_value: - blocks[block_name][f.name] = field_value + blocks[block_name][name] = field_value elif meta.get("direction") and isinstance(field_value, Path): - t = _path_to_tuple(f, field_value) + t = _path_to_tuple(name, f, field_value) blocks[block_name][t[0].lower()] = t elif isinstance(field_value, list) and field_value and isinstance(field_value[0], Item): - blocks[block_name][f.name] = _rows_to_tuples(field_value) + blocks[block_name][name] = _rows_to_tuples(field_value) elif isinstance(field_value, list) and field_value and isinstance(field_value[0], tuple): - blocks[block_name][f.name] = field_value + blocks[block_name][name] = field_value elif meta.get("shape") and not isinstance(field_value, bool): if meta["shape"]: @@ -225,26 +228,26 @@ def _unstructure_package(value: Package) -> dict[str, Any]: _nlay = _dims_d.get("nlay", 0) _ncpl = _dims_d.get("ncpl", 0) if _nlay > 1 and _ncpl > 0 and field_value.size == _nlay * _ncpl: - blocks[block_name][f.name] = xr.DataArray( + blocks[block_name][name] = xr.DataArray( field_value.reshape(_nlay, _ncpl), dims=("nlay", "ncpl"), ) continue - blocks[block_name][f.name] = _wrap_array(field_value) + blocks[block_name][name] = _wrap_array(field_value) - elif f.name == "auxiliary" and isinstance(field_value, list): - blocks[block_name][f.name] = ("AUXILIARY",) + tuple(field_value) + elif name == "auxiliary" and isinstance(field_value, list): + blocks[block_name][name] = ("AUXILIARY",) + tuple(field_value) elif isinstance(field_value, Record): - blocks[block_name][f.name] = field_value.to_tokens() + blocks[block_name][name] = field_value.to_tokens() elif dfn_type in ("integer", "double", "double precision"): if field_value == 0 and meta.get("auto_from"): continue - blocks[block_name][f.name] = field_value + blocks[block_name][name] = field_value elif dfn_type == "string" and field_value: - blocks[block_name][f.name] = field_value + blocks[block_name][name] = field_value # `maxbound` is a computed property on some classes, not a real field if isinstance(getattr(cls, "maxbound", None), property): @@ -283,10 +286,10 @@ def unstructure_component(value: Component) -> dict[str, Any]: def _unstructure_component(value: Component) -> dict[str, Any]: """Unstructure an internal-node component (Gwf, Simulation, etc.) with - attrs-typed child fields, including its child binding blocks.""" + pydantic-typed child fields, including its child binding blocks.""" blockspec = blocks_dict(type(value)) blocks: dict[str, dict[str, Any]] = {} - fields_by_name = {f.name: f for f in attrs.fields(type(value))} # type: ignore[arg-type] + fields_by_name = dict(pydantic_fields(type(value))) # create child component binding blocks blocks.update(_make_binding_blocks(value)) @@ -298,11 +301,13 @@ def _unstructure_component(value: Component) -> dict[str, Any]: for field_name in block.keys(): # Skip child components already processed as bindings field = fields_by_name.get(field_name) + fmeta = field_meta(field) if field is not None else {} if ( isinstance(value, Context) and field is not None and child_field_candidates(field) is not None - and field.metadata.get("block") == block_name + and isinstance(fmeta, dict) + and fmeta.get("block") == block_name ): continue @@ -322,7 +327,7 @@ def _unstructure_component(value: Component) -> dict[str, Any]: blocks[block_name][field_name] = field_value case Path(): assert field is not None # field_name comes from blocks_dict(type(value)) - t = _path_to_tuple(field, field_value) + t = _path_to_tuple(field_name, field, field_value) blocks[block_name][t[0]] = t case datetime(): blocks[block_name][field_name] = field_value.isoformat() diff --git a/flopy4/mf6/converter/ingress/structure.py b/flopy4/mf6/converter/ingress/structure.py index 93bc1807..12135712 100644 --- a/flopy4/mf6/converter/ingress/structure.py +++ b/flopy4/mf6/converter/ingress/structure.py @@ -3,7 +3,6 @@ from pathlib import Path from typing import Any, get_args -import attrs import numpy as np from flopy4.dimensions import DimensionProvider @@ -13,10 +12,11 @@ from flopy4.mf6.package import Package from flopy4.mf6.record import Record from flopy4.mf6.spec import repeating_array_key_type, to_field_type +from flopy4.spec import field_meta, pydantic_fields def _inner_class_type(field_type) -> type[Record] | None: - """If field_type is Optional[C] where C is an attrs inner-record class, return C.""" + """If field_type is Optional[C] where C is a pydantic inner-record class, return C.""" args = get_args(field_type) if not args: return None @@ -213,7 +213,7 @@ def _griddata_flat_length(f, dims: dict, default: int) -> int: back to `default` (the grid's total node count) when the field's declared shape dimension isn't resolvable from `dims`. """ - shape_meta = f.metadata.get("shape") + shape_meta = field_meta(f).get("shape") if shape_meta: dim_name = shape_meta[-1] if isinstance(shape_meta, (tuple, list)) else shape_meta if dim_name in dims: @@ -245,11 +245,12 @@ def _parse_griddata_block( continue key = str(row[0]).lower() f = fields_by_name.get(key) - if f is None or f.metadata.get("block") != "griddata": + meta = field_meta(f) if f is not None else {} + if f is None or meta.get("block") != "griddata": i += 1 continue - dtype = np.int64 if to_field_type(f.type) == "integer" else np.float64 + dtype = np.int64 if to_field_type(f.annotation) == "integer" else np.float64 layered = any(str(t).upper() == "LAYERED" for t in row[1:]) i += 1 @@ -287,13 +288,13 @@ def _parse_griddata_block( stacked = np.concatenate(layers).astype(dtype) if stacked.size < target and stacked.size and target % stacked.size == 0: stacked = np.tile(stacked, target // stacked.size) - result[f.name] = stacked + result[key] = stacked else: if i >= len(rows): break length = _griddata_flat_length(f, dims, nodes) value, i = _read_control_record(rows, i, workspace, dtype, length) - result[f.name] = value + result[key] = value return result @@ -366,10 +367,9 @@ def _parse_readarray_period_block( _, i = _read_control_record(rows, i, workspace, np.float64, ncpl) continue - is_int = to_field_type(f.type) == "integer" - is_layered = f.metadata.get("layered", False) or any( - str(t).upper() == "LAYERED" for t in row[1:] - ) + is_int = to_field_type(f.annotation) == "integer" + meta = field_meta(f) + is_layered = meta.get("layered", False) or any(str(t).upper() == "LAYERED" for t in row[1:]) dtype = np.int64 if is_int else np.float64 if is_layered: @@ -383,12 +383,12 @@ def _parse_readarray_period_block( if layers: if len(layers) < nlay: layers = (layers * nlay)[:nlay] - result[f.name] = np.stack(layers) # (nlay, ncpl) + result[key] = np.stack(layers) # (nlay, ncpl) else: if i >= len(rows): break value, i = _read_control_record(rows, i, workspace, dtype, ncpl) - result[f.name] = value + result[key] = value return result @@ -437,7 +437,7 @@ def _resolve_bindings(cls: type, raw_lower: dict, workspace: Path) -> dict[str, dims=dims)` calls -- `dimensions.py`'s object-graph walk only helps once a child is already attached, not while its siblings are still loading. """ - from flopy4.attrs_xarray import child_field_candidates + from flopy4.dataclass_xarray import child_field_candidates from flopy4.mf6.converter.binding import component_ftype from flopy4.mf6.exchange import Exchange from flopy4.mf6.model import Model @@ -454,15 +454,16 @@ def _resolve_bindings(cls: type, raw_lower: dict, workspace: Path) -> dict[str, model_prefix = cls.__name__.lower() if issubclass(cls, Model) else None fields_by_block: dict[str, list] = {} - for f in attrs.fields(cls): # type: ignore[arg-type] + for name, f in pydantic_fields(cls).items(): spec = child_field_candidates(f) if spec is None: continue - block_name = f.metadata.get("block") + meta = field_meta(f) + block_name = meta.get("block") if isinstance(meta, dict) else None if block_name is None: continue kind, candidates = spec - fields_by_block.setdefault(block_name, []).append((f.name, kind, candidates)) + fields_by_block.setdefault(block_name, []).append((name, kind, candidates)) if not fields_by_block: return {} @@ -538,7 +539,7 @@ def _resolve_bindings(cls: type, raw_lower: dict, workspace: Path) -> dict[str, if issubclass(target_cls, Package) else target_cls.load(workspace / fname, name=pname) ) - child.filename = fname + child.filename = Path(fname) _apply_binding_terms(child, row[2:]) if isinstance(child, DimensionProvider): dims = {**dims, **child.get_dims()} @@ -625,56 +626,72 @@ def structure_component( # so a block keyword can't shadow them (e.g. utl-tas's "NAME ..." row # is for time_series_name, not Package.name). all_fields = { - f.name: f for f in attrs.fields(cls) if f.init is not False and "block" in f.metadata + fname: f + for fname, f in pydantic_fields(cls).items() + if f.init is not False and "block" in field_meta(f) } alias_map: dict[str, str] = {} # alias → name - for f in attrs.fields(cls): - if f.alias and f.alias != f.name: - alias_map[f.alias] = f.name + for fname, f in pydantic_fields(cls).items(): + if f.alias and f.alias != fname: + alias_map[f.alias] = fname # Block-level file records ("TS6 FILEIN ") are keyed by their # trigger keyword, the path field's _keyword metadata (see spec.path), # not by the field's py name (ts_file). - file_fields: dict[str, Any] = { - f.metadata["_keyword"]: f for f in all_fields.values() if f.metadata.get("_keyword") + file_fields: dict[str, tuple[str, Any]] = { + kw: (fname, f) for fname, f in all_fields.items() if (kw := field_meta(f).get("_keyword")) } # Index Optional[InnerClass] fields by the inner class's _keyword (lowercase). # Covers options-block compound records like Npf.Cvoptions, Ims.Rclose, etc. + # + # NOTE: every field-iteration loop below in this function deliberately + # uses `fname`, never `name` -- this function's own `name` KEYWORD + # PARAMETER (the explicit override name, used at the very bottom) would + # otherwise be silently shadowed by the loop variable (a plain `for` + # doesn't get its own scope in Python, unlike a comprehension). Confirmed + # as a real bug by running the real test suite: every load-a-real-model + # test failed, with kwargs["name"] ending up as the *last field name + # iterated* (e.g. "output") instead of the caller's actual override. inner_class_fields: dict[str, tuple] = {} - for f in attrs.fields(cls): + for fname, f in pydantic_fields(cls).items(): if f.init is False: continue - inner_cls = _inner_class_type(f.type) + inner_cls = _inner_class_type(f.annotation) if inner_cls is None: continue kw = vars(inner_cls).get("_keyword", "") if kw: - inner_class_fields[kw.lower()] = (f, inner_cls) + inner_class_fields[kw.lower()] = (fname, f, inner_cls) # Identify Item-list fields (packagedata, connectiondata, partitions …) -- # the field's own type annotation (Optional[list[ItemClass]] or # Optional[dict[int, list[ItemClass]]]) is the schema. - block_item_fields: dict[str, tuple] = {} # block_name → (field, item_cls) + block_item_fields: dict[str, tuple] = {} # block_name → (name, field, item_cls) period_field = None # field for the period Item-list + period_field_name: "str | None" = None period_item_cls: "type[Item] | tuple[type[Item], ...] | None" = None # Fill-forward repeating blocks (period), per the fields' own # fill_forward metadata (from the DFN's BlockHeader.fill_forward). fill_forward_blocks = { - f.metadata["block"] for f in attrs.fields(cls) if f.metadata.get("fill_forward") + field_meta(f)["block"] + for f in pydantic_fields(cls).values() + if field_meta(f).get("fill_forward") } - for f in attrs.fields(cls): - block = f.metadata.get("block", "") - item_cls = item_list_type(f.type) + for fname, f in pydantic_fields(cls).items(): + meta = field_meta(f) + block = meta.get("block", "") if isinstance(meta, dict) else "" + item_cls = item_list_type(f.annotation) if item_cls is None: continue - if f.metadata.get("fill_forward"): + if meta.get("fill_forward"): period_field = f + period_field_name = fname period_item_cls = item_cls else: - block_item_fields[block] = (f, item_cls) + block_item_fields[block] = (fname, f, item_cls) # Array fields whose own block repeats per header value -- e.g. # utl-tas's tas_array, dict[float, ndarray]. Detected structurally from @@ -682,11 +699,13 @@ def structure_component( # way block_item_fields above is detected via item_list_type -- not via # a metadata flag. Mirrors egress/unstructure.py's write path. repeating_array_fields = { - f.name: f - for f in attrs.fields(cls) - if repeating_array_key_type(f.type) is not None and f.init is not False + fname: f + for fname, f in pydantic_fields(cls).items() + if repeating_array_key_type(f.annotation) is not None and f.init is not False + } + repeating_array_block_prefixes = { + field_meta(f)["block"] for f in repeating_array_fields.values() } - repeating_array_block_prefixes = {f.metadata["block"] for f in repeating_array_fields.values()} # ── Pass 1: scalar blocks (options, dimensions, etc.) ──────────────────── kwargs: dict[str, Any] = {} @@ -707,36 +726,79 @@ def structure_component( if (ff := file_fields.get(key)) is not None: # KEYWORD FILEIN|FILEOUT : the path follows the # keyword and the direction token. + ff_name, ff_f = ff tokens = row[1:] if tokens and str(tokens[0]).upper() in ("FILEIN", "FILEOUT"): tokens = tokens[1:] if tokens: - kwargs[ff.alias or ff.name] = Path(_strip_quotes(str(tokens[0]))) + kwargs[ff_f.alias or ff_name] = Path(_strip_quotes(str(tokens[0]))) continue - f = all_fields.get(key) or all_fields.get(alias_map.get(key, "")) - if f is None or f.init is False: + found_name = key if key in all_fields else alias_map.get(key) + found = all_fields.get(found_name) if found_name is not None else None + if found_name is None or found is None or found.init is False: if key in inner_class_fields: - cand_f, inner_cls = inner_class_fields[key] - cand_init = cand_f.alias if cand_f.alias else cand_f.name + cand_name, cand_f, inner_cls = inner_class_fields[key] + cand_init = cand_f.alias if cand_f.alias else cand_name kwargs[cand_init] = inner_cls.from_tokens(row) continue - init_key = f.alias if f.alias else f.name - # A Record-typed field must go through from_tokens(), even when - # the matched token is the field's own name rather than the - # record's separate trigger keyword (e.g. sfacrecord's outer - # field is itself named "sfac"). - inner_cls = _inner_class_type(f.type) - if inner_cls is not None: - kwargs[init_key] = inner_cls.from_tokens(row) + f = found + # A field whose OWN name happens to equal its inner Record's + # _keyword (e.g. Npf.rewet: Optional[Rewet], Rewet._keyword == + # "rewet") matches `all_fields` above before `inner_class_fields` + # (a fallback keyed by the row's *keyword*, not the field's own + # name) is ever consulted, so without this it would get a raw + # token list instead of a Rewet instance. Route it through the + # same from_tokens() path here + # too, rather than falling into the generic scalar/list branch + # below, which is wrong for any Record-typed field. + direct_inner_cls = _inner_class_type(f.annotation) + if direct_inner_cls is not None: + init_key = f.alias if f.alias else found_name + kwargs[init_key] = direct_inner_cls.from_tokens(row) continue - if len(row) == 1: + init_key = f.alias if f.alias else found_name + if len(row) == 1 or to_field_type(f.annotation) == "keyword": + # A bare bool/keyword flag -- true just from the keyword's + # presence on the row, regardless of trailing tokens. A + # compound record this keyword is also meant to trigger + # (e.g. Gwf's `newton`/`newtonoptions` pair) isn't resolved + # here if it isn't reachable via inner_class_fields (that + # needs the inner class's own _keyword ClassVar). A + # trailing modifier token (e.g. "NEWTON UNDER_RELAXATION") + # isn't a valid bool, so the keyword's presence is the + # field's value. kwargs[init_key] = True else: # List-valued options (auxiliary, etc.) have shape metadata; - # always keep them as a list so __attrs_post_init__ can use len(). - is_list_opt = isinstance(f.metadata.get("shape"), tuple) + # always keep them as a list so __post_init__ can use len(). + # `auxiliary` itself doesn't actually carry shape= metadata + # in the real DFN corpus (confirmed empirically) despite + # being declared Optional[list[str]], so it's also matched + # by name directly -- otherwise a row naming exactly one + # aux variable (e.g. "AUXILIARY MULT", row length 2) fell + # through to the plain-scalar branch below and stored a + # bare string, which fails list[str] validation. + _f_meta = field_meta(f) + is_list_opt = ( + isinstance(_f_meta, dict) and isinstance(_f_meta.get("shape"), tuple) + ) or found_name == "auxiliary" if is_list_opt: kwargs[init_key] = list(row[1:]) + elif to_field_type(f.annotation) in ( + "integer", + "double", + "double precision", + "string", + ): + # A genuine scalar field (int/float/str, not + # shape-tagged) with extra trailing tokens on its row -- + # seen in some older, migrated fixtures (e.g. IMS + # "OUTER_MAXIMUM 100 500" or "REORDERING_METHOD NONE + # RKM"), the second value belonging to a field the + # current DFN schema no longer declares. A list fails + # scalar validation, so take just the first value and + # drop the stale extra tokens. + kwargs[init_key] = row[1] else: kwargs[init_key] = list(row[1:]) if len(row) > 2 else row[1] @@ -753,7 +815,7 @@ def structure_component( effective_dims = dims or _self_dims_from_kwargs(kwargs) # ── Pass 2: block Item-list fields (packagedata, partitions …) ────────── - for block_name, (f, item_cls) in block_item_fields.items(): + for block_name, (name, f, item_cls) in block_item_fields.items(): rows = raw_lower.get(block_name, []) if not rows: continue @@ -762,7 +824,7 @@ def structure_component( rows, item_cls, naux=naux, boundnames=boundnames, dims=effective_dims ) if row_list is not None: - init_key = f.alias if (f.alias and not f.alias.startswith("_")) else f.name + init_key = f.alias if (f.alias and not f.alias.startswith("_")) else name kwargs[init_key] = row_list # ── Pass 3: period blocks ──────────────────────────────────────────────── @@ -776,7 +838,9 @@ def structure_component( if kper_rows: if period_field is not None: - assert period_item_cls is not None # set together with period_field above + assert ( + period_item_cls is not None and period_field_name is not None + ) # set with period_field spd: dict[int, list] = {} for kper, rows in sorted(kper_rows.items()): if not rows: @@ -789,7 +853,7 @@ def structure_component( spd[kper] = row_list if spd: # Use the alias (stress_period_data) as the init kwarg - init_key = period_field.alias if period_field.alias else period_field.name + init_key = period_field.alias if period_field.alias else period_field_name kwargs[init_key] = spd else: @@ -800,11 +864,11 @@ def structure_component( # set above) and not dict-wrapped (else it'd be a # repeating_array_field instead, see below). ra_fields = { - f.name: f - for f in attrs.fields(cls) - if f.metadata.get("fill_forward") - and f.name not in repeating_array_fields - and to_field_type(f.type) in ("integer", "double") + name: f + for name, f in pydantic_fields(cls).items() + if field_meta(f).get("fill_forward") + and name not in repeating_array_fields + and to_field_type(f.annotation) in ("integer", "double") and f.init is not False } if ra_fields and dims: @@ -816,7 +880,8 @@ def structure_component( # fill-forward semantics (egress skips all-FILL_DNODATA periods). accum: dict[str, np.ndarray] = {} for fname, f in ra_fields.items(): - shape = (nper, nlay, ncpl) if f.metadata.get("layered", False) else (nper, ncpl) + _meta = field_meta(f) + shape = (nper, nlay, ncpl) if _meta.get("layered", False) else (nper, ncpl) accum[fname] = np.full(shape, FILL_DNODATA) for kper, rows in sorted(kper_rows.items()): if not rows: @@ -844,12 +909,12 @@ def structure_component( nlay = effective_dims.get("nlay", 1) ncpl = nodes // nlay if nlay > 1 else nodes for fname, f in repeating_array_fields.items(): - block = f.metadata["block"] - key_type = repeating_array_key_type(f.type) + block = field_meta(f)["block"] + key_type = repeating_array_key_type(f.annotation) series_rows = _group_repeating_rows(raw_lower, {block}, key_type).get(block) if not series_rows: continue - dtype = np.int64 if to_field_type(f.type) == "integer" else np.float64 + dtype = np.int64 if to_field_type(f.annotation) == "integer" else np.float64 length = _griddata_flat_length(f, effective_dims, ncpl) series: dict[Any, np.ndarray] = {} for header, rows in sorted(series_rows.items()): @@ -873,9 +938,9 @@ def structure_component( effective_dims = dims or _self_dims_from_kwargs(kwargs) if effective_dims: gd_fields = { - f.name: f - for f in attrs.fields(cls) - if f.metadata.get("block") == "griddata" and f.init is not False + name: f + for name, f in pydantic_fields(cls).items() + if field_meta(f).get("block") == "griddata" and f.init is not False } parsed = _parse_griddata_block(griddata_rows, gd_fields, effective_dims, workspace) kwargs.update(parsed) diff --git a/flopy4/mf6/ems.py b/flopy4/mf6/ems.py index ce44218b..5844473b 100644 --- a/flopy4/mf6/ems.py +++ b/flopy4/mf6/ems.py @@ -1,12 +1,12 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.solution import Solution +from flopy4.mf6.solution import CFG, Solution -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ems(Solution): dfn_name: ClassVar[str] = "sln-ems" diff --git a/flopy4/mf6/exchange.py b/flopy4/mf6/exchange.py index 8c96554d..69efe04d 100644 --- a/flopy4/mf6/exchange.py +++ b/flopy4/mf6/exchange.py @@ -2,13 +2,13 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Exchange(Package, ABC): exgtype: Optional[type] = field(default=None) # type: ignore exgfile: Optional[Path] = field(default=None) # type: ignore @@ -19,14 +19,14 @@ def default_filename(self) -> str: return f"{self.name}.exg" # type: ignore -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class GwfGwt(Exchange): """GWF-GWT flow-transport exchange (declares coupling in mfsim.nam).""" dfn_name: ClassVar[str] = "exg-gwfgwt" -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class GwfGwe(Exchange): """GWF-GWE flow-energy exchange (declares coupling in mfsim.nam).""" diff --git a/flopy4/mf6/exg/gwfgwe.py b/flopy4/mf6/exg/gwfgwe.py index f1415b11..bb318939 100644 --- a/flopy4/mf6/exg/gwfgwe.py +++ b/flopy4/mf6/exg/gwfgwe.py @@ -1,11 +1,11 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwfgwe(Package): dfn_name: ClassVar[str] = "exg-gwfgwe" diff --git a/flopy4/mf6/exg/gwfgwt.py b/flopy4/mf6/exg/gwfgwt.py index a4156ac6..7b8538b6 100644 --- a/flopy4/mf6/exg/gwfgwt.py +++ b/flopy4/mf6/exg/gwfgwt.py @@ -1,11 +1,11 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwfgwt(Package): dfn_name: ClassVar[str] = "exg-gwfgwt" diff --git a/flopy4/mf6/exg/gwfprt.py b/flopy4/mf6/exg/gwfprt.py index d2de62e5..2b88f0d8 100644 --- a/flopy4/mf6/exg/gwfprt.py +++ b/flopy4/mf6/exg/gwfprt.py @@ -1,11 +1,11 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwfprt(Package): dfn_name: ClassVar[str] = "exg-gwfprt" diff --git a/flopy4/mf6/gwe/__init__.py b/flopy4/mf6/gwe/__init__.py index 7f0ad7a6..99c9f54d 100644 --- a/flopy4/mf6/gwe/__init__.py +++ b/flopy4/mf6/gwe/__init__.py @@ -1,9 +1,9 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs from flopy.discretization.structuredgrid import StructuredGrid from flopy.discretization.vertexgrid import VertexGrid +from pydantic.dataclasses import dataclass from flopy4.mf6.gwe.adv import Adv from flopy4.mf6.gwe.cnd import Cnd @@ -18,7 +18,7 @@ from flopy4.mf6.gwe.oc import Oc from flopy4.mf6.gwe.ssm import Ssm from flopy4.mf6.gwf.disbase import DisBase -from flopy4.mf6.model import Model +from flopy4.mf6.model import CFG, Model from flopy4.mf6.spec import field, path from flopy4.utils import to_path @@ -50,7 +50,7 @@ def convert_grid(value): ] -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwe(Model): dfn_name: ClassVar[str] = "gwe-nam" @@ -86,9 +86,9 @@ class Gwe(Model): adv: Adv | None = field(block="packages", default=None) cnd: Cnd | None = field(block="packages", default=None) est: Est | None = field(block="packages", default=None) - ctp: list[Ctp] = field(block="packages", default=attrs.Factory(list)) - esl: list[Esl] = field(block="packages", default=attrs.Factory(list)) - lke: list[Lke] = field(block="packages", default=attrs.Factory(list)) + ctp: list[Ctp] = field(block="packages", default_factory=list) + esl: list[Esl] = field(block="packages", default_factory=list) + lke: list[Lke] = field(block="packages", default_factory=list) ssm: Ssm | None = field(block="packages", default=None) mve: Mve | None = field(block="packages", default=None) diff --git a/flopy4/mf6/gwe/adv.py b/flopy4/mf6/gwe/adv.py index 1b9ae233..ed74f2b8 100644 --- a/flopy4/mf6/gwe/adv.py +++ b/flopy4/mf6/gwe/adv.py @@ -1,13 +1,13 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Adv(Package): dfn_name: ClassVar[str] = "gwe-adv" diff --git a/flopy4/mf6/gwe/cnd.py b/flopy4/mf6/gwe/cnd.py index e5a04a1f..f4c8f1d9 100644 --- a/flopy4/mf6/gwe/cnd.py +++ b/flopy4/mf6/gwe/cnd.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Cnd(Package): dfn_name: ClassVar[str] = "gwe-cnd" diff --git a/flopy4/mf6/gwe/ctp.py b/flopy4/mf6/gwe/ctp.py index 97f1be9e..a4386466 100644 --- a/flopy4/mf6/gwe/ctp.py +++ b/flopy4/mf6/gwe/ctp.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ctp(Package): dfn_name: ClassVar[str] = "gwe-ctp" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) temp: Union[float, str] = field(time_series=True) @@ -74,7 +75,7 @@ class StressPeriodData(Item): direction="in", keyword="obs6", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwe/dis.py b/flopy4/mf6/gwe/dis.py index 03146527..013efe77 100644 --- a/flopy4/mf6/gwe/dis.py +++ b/flopy4/mf6/gwe/dis.py @@ -1,18 +1,19 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import StructuredGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Dis(DisBase): dfn_name: ClassVar[str] = "gwe-dis" @@ -31,11 +32,11 @@ class Dis(DisBase): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=1, block="dimensions") ncol: int = field(default=2, block="dimensions") nrow: int = field(default=2, block="dimensions") - delr: NDArray[np.float64] = field( + delr: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a row", block="griddata", @@ -43,7 +44,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - delc: NDArray[np.float64] = field( + delc: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a column", block="griddata", @@ -51,7 +52,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=1.0, longname="cell top elevation", block="griddata", @@ -59,7 +60,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=0.0, longname="cell bottom elevation", block="griddata", @@ -76,12 +77,12 @@ class Dis(DisBase): longname="idomain existence array", ) - def __attrs_post_init__(self): + def __post_init__(self, dims: Optional[dict] = None): self.nodes = self.ncol * self.nrow * self.nlay self.ncpl = self.ncol * self.nrow self.nvert = (self.ncol + 1) * (self.nrow + 1) self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwe/disv.py b/flopy4/mf6/gwe/disv.py index 3fc3845c..b10ec6de 100644 --- a/flopy4/mf6/gwe/disv.py +++ b/flopy4/mf6/gwe/disv.py @@ -1,31 +1,32 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field, field_validator +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.item import Item from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import VertexGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Disv(DisBase): dfn_name: ClassVar[str] = "gwe-disv" - @attrs.define(slots=False) + @dataclass(config=CFG) class Cell2dRecord: - icell2d: int = attrs.field() - xc: float = attrs.field() - yc: float = attrs.field() - ncvert: int = attrs.field() - icvert: tuple[int, ...] = attrs.field() + icell2d: int = Field() + xc: float = Field() + yc: float = Field() + ncvert: int = Field() + icvert: tuple[int, ...] = Field() - @attrs.define + @dataclass(config=CFG) class Vertices(Item): iv: int xv: float @@ -46,11 +47,11 @@ class Vertices(Item): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=0, block="dimensions") ncpl: int = field(default=0, block="dimensions") nvert: int = field(default=0, block="dimensions") - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=None, longname="model top elevation", block="griddata", @@ -58,7 +59,7 @@ class Vertices(Item): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=None, longname="model bottom elevation", block="griddata", @@ -74,20 +75,31 @@ class Vertices(Item): layered=True, netcdf=True, ) - iv: Optional[NDArray[np.int64]] = attrs.field(default=None) - xv: Optional[NDArray[np.float64]] = attrs.field(default=None) - yv: Optional[NDArray[np.float64]] = attrs.field(default=None) + iv: Optional[NDArray[np.int64]] = Field(default=None) + xv: Optional[NDArray[np.float64]] = Field(default=None) + yv: Optional[NDArray[np.float64]] = Field(default=None) + + # iv/xv/yv are declared NDArray-typed but commonly constructed from a + # plain list/tuple (see from_grid() below) -- unlike Package's own + # griddata fields, these carry no block="griddata"/shape= metadata, so + # Package._coerce_arrays' shape-driven check doesn't reach them, so + # they get their own mode="before" coercion (as do Tdis.perlen/nstp/ + # tsmult). + @field_validator("iv", mode="before") + @classmethod + def _coerce_iv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.int64) + + @field_validator("xv", "yv", mode="before") + @classmethod + def _coerce_xv_yv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.float64) + vertices: Optional[list[Vertices]] = field(default=None, block="vertices") - cell2ddata: Optional[list] = attrs.field(default=None) + cell2ddata: Optional[list] = Field(default=None) cell2d: Optional[list] = field(default=None, init=False, block="cell2d") - def __attrs_post_init__(self): - if self.iv is not None and (not isinstance(self.iv, np.ndarray)): - object.__setattr__(self, "iv", np.asarray(self.iv, dtype=np.int64)) - if self.xv is not None and (not isinstance(self.xv, np.ndarray)): - object.__setattr__(self, "xv", np.asarray(self.xv, dtype=np.float64)) - if self.yv is not None and (not isinstance(self.yv, np.ndarray)): - object.__setattr__(self, "yv", np.asarray(self.yv, dtype=np.float64)) + def __post_init__(self, dims: Optional[dict] = None): if self.iv is not None and self.xv is not None and (self.yv is not None): rows = [ self.Vertices(iv=int(iv) + 1, xv=float(xv), yv=float(yv)) @@ -95,18 +107,18 @@ def __attrs_post_init__(self): ] object.__setattr__(self, "vertices", rows) if self.cell2ddata is not None: - rows = [] + cell_rows = [] for rec in self.cell2ddata: row = (rec.icell2d + 1, rec.xc, rec.yc, rec.ncvert) + tuple( (v + 1 for v in rec.icvert) ) - rows.append(row) - object.__setattr__(self, "cell2d", rows) + cell_rows.append(row) + object.__setattr__(self, "cell2d", cell_rows) self.nodes = self.ncpl * self.nlay self.nrow = 0 self.ncol = 0 self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwe/esl.py b/flopy4/mf6/gwe/esl.py index f9132764..b7a115a3 100644 --- a/flopy4/mf6/gwe/esl.py +++ b/flopy4/mf6/gwe/esl.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Esl(Package): dfn_name: ClassVar[str] = "gwe-esl" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) senerrate: Union[float, str] = field(time_series=True) @@ -74,7 +75,7 @@ class StressPeriodData(Item): direction="in", keyword="obs6", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwe/est.py b/flopy4/mf6/gwe/est.py index a8189674..2293cdbf 100644 --- a/flopy4/mf6/gwe/est.py +++ b/flopy4/mf6/gwe/est.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Est(Package): dfn_name: ClassVar[str] = "gwe-est" diff --git a/flopy4/mf6/gwe/fmi.py b/flopy4/mf6/gwe/fmi.py index e08198c5..e25766a0 100644 --- a/flopy4/mf6/gwe/fmi.py +++ b/flopy4/mf6/gwe/fmi.py @@ -2,18 +2,19 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Fmi(Package): dfn_name: ClassVar[str] = "gwe-fmi" - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): flowtype: Union[float, str] = field() fname: Path = path(converter=Path, direction="in") @@ -31,7 +32,7 @@ class Packagedata(Item): optional=True, longname="correct for flow imbalance", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", ) diff --git a/flopy4/mf6/gwe/ic.py b/flopy4/mf6/gwe/ic.py index 2e67c8c4..c0a5b475 100644 --- a/flopy4/mf6/gwe/ic.py +++ b/flopy4/mf6/gwe/ic.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ic(Package): dfn_name: ClassVar[str] = "gwe-ic" diff --git a/flopy4/mf6/gwe/lke.py b/flopy4/mf6/gwe/lke.py index 1d155a15..ff852e31 100644 --- a/flopy4/mf6/gwe/lke.py +++ b/flopy4/mf6/gwe/lke.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Lke(Package): dfn_name: ClassVar[str] = "gwe-lke" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): lakeno: int = field(index=True, pk=True) strt: float = field() @@ -25,43 +26,43 @@ class Packagedata(Item): aux: tuple = () boundname: Optional[str] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class Status(Item): _keyword: ClassVar[str] = "status" lakeno: int = field(index=True, fk="packagedata.lakeno") status: Union[float, str] = field() - @attrs.define + @dataclass(config=CFG) class Temperature(Item): _keyword: ClassVar[str] = "temperature" lakeno: int = field(index=True, fk="packagedata.lakeno") temperature: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Rainfall(Item): _keyword: ClassVar[str] = "rainfall" lakeno: int = field(index=True, fk="packagedata.lakeno") rainfall: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Evaporation(Item): _keyword: ClassVar[str] = "evaporation" lakeno: int = field(index=True, fk="packagedata.lakeno") evaporation: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Runoff(Item): _keyword: ClassVar[str] = "runoff" lakeno: int = field(index=True, fk="packagedata.lakeno") runoff: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class ExtInflow(Item): _keyword: ClassVar[str] = "ext-inflow" lakeno: int = field(index=True, fk="packagedata.lakeno") ext_inflow: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Auxiliary(Item): _keyword: ClassVar[str] = "auxiliary" lakeno: int = field(index=True, fk="packagedata.lakeno") @@ -159,12 +160,12 @@ class Auxiliary(Item): direction="in", keyword="obs6", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", write_if_empty=True, ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwe/mve.py b/flopy4/mf6/gwe/mve.py index 6ce9f6eb..3f7aa7d6 100644 --- a/flopy4/mf6/gwe/mve.py +++ b/flopy4/mf6/gwe/mve.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Mve(Package): dfn_name: ClassVar[str] = "gwe-mve" diff --git a/flopy4/mf6/gwe/oc.py b/flopy4/mf6/gwe/oc.py index 4f09e060..718404e5 100644 --- a/flopy4/mf6/gwe/oc.py +++ b/flopy4/mf6/gwe/oc.py @@ -2,62 +2,63 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import Field, SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Oc(Package): dfn_name: ClassVar[str] = "gwe-oc" - @attrs.define + @dataclass(config=CFG) class Format(Record): _keyword: ClassVar[str] = "" - format_: str = attrs.field() - columns: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - width: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - digits: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) + format_: str = Field() + columns: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + width: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + digits: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Temperatureprint(Record): _keyword: ClassVar[str] = "temperature" _extra_tokens: ClassVar[tuple[str, ...]] = ("PRINT_FORMAT",) - formatrecord: "Oc.Format" = attrs.field() + formatrecord: "Oc.Format" = Field() - @attrs.define + @dataclass(config=CFG) class All(Item): _keyword: ClassVar[str] = "all" - @attrs.define + @dataclass(config=CFG) class First(Item): _keyword: ClassVar[str] = "first" - @attrs.define + @dataclass(config=CFG) class Last(Item): _keyword: ClassVar[str] = "last" - @attrs.define + @dataclass(config=CFG) class Frequency(Item): _keyword: ClassVar[str] = "frequency" frequency: int = field() - @attrs.define + @dataclass(config=CFG) class Steps(Item): _keyword: ClassVar[str] = "steps" steps: tuple = field(default=(), array=True) - @attrs.define + @dataclass(config=CFG) class Save(Item): _keyword: ClassVar[str] = "save" rtype: Union[float, str] = field() ocsetting: "Oc.All | Oc.First | Oc.Last | Oc.Frequency | Oc.Steps" = field() - @attrs.define + @dataclass(config=CFG) class Print(Item): _keyword: ClassVar[str] = "print" rtype: Union[float, str] = field() @@ -93,7 +94,7 @@ class Print(Item): default=None, block="options", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwe/ssm.py b/flopy4/mf6/gwe/ssm.py index 68d5d51f..9a3543fa 100644 --- a/flopy4/mf6/gwe/ssm.py +++ b/flopy4/mf6/gwe/ssm.py @@ -2,24 +2,25 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ssm(Package): dfn_name: ClassVar[str] = "gwe-ssm" - @attrs.define + @dataclass(config=CFG) class Sources(Item): pname: Union[float, str] = field() srctype: Union[float, str] = field() auxname: Union[float, str] = field() - @attrs.define + @dataclass(config=CFG) class Fileinput(Item): pname: Union[float, str] = field() spc6_filename: Path = path(converter=Path, direction="in", keyword="spc6") @@ -37,12 +38,12 @@ class Fileinput(Item): optional=True, longname="save calculated flows to budget file", ) - sources: Optional[list[Sources]] = field( + sources: Optional[SkipValidation[list[Sources]]] = field( default=None, block="sources", write_if_empty=True, ) - fileinput: Optional[list[Fileinput]] = field( + fileinput: Optional[SkipValidation[list[Fileinput]]] = field( default=None, block="fileinput", ) diff --git a/flopy4/mf6/gwf/__init__.py b/flopy4/mf6/gwf/__init__.py index 037ad392..a9786caa 100644 --- a/flopy4/mf6/gwf/__init__.py +++ b/flopy4/mf6/gwf/__init__.py @@ -1,13 +1,13 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs import xarray as xr import xugrid as xu -from attrs import define from flopy.discretization.grid import Grid from flopy.discretization.structuredgrid import StructuredGrid from flopy.discretization.vertexgrid import VertexGrid +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6.gwf.buy import Buy from flopy4.mf6.gwf.chd import Chd @@ -35,7 +35,7 @@ from flopy4.mf6.gwf.vsc import Vsc from flopy4.mf6.gwf.wel import Wel from flopy4.mf6.gwf.welg import Welg -from flopy4.mf6.model import Model +from flopy4.mf6.model import CFG, Model from flopy4.mf6.spec import field, path from flopy4.mf6.utils import open_cbc, open_hds from flopy4.utils import to_path @@ -85,18 +85,18 @@ def convert_grid(value): raise TypeError(f"Expected Grid or Dis/Disv, got {type(value)}") -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwf(Model): dfn_name: ClassVar[str] = "gwf-nam" - @define + @dataclass(config=CFG) class NewtonOptions: newton: bool = field() under_relaxation: bool = field() - @define + @dataclass(config=CFG) class Output: - parent: "Gwf" = attrs.field(repr=False) + parent: "Gwf" = Field(repr=False) @property def head(self) -> xr.DataArray | xu.UgridDataArray: @@ -179,20 +179,26 @@ def budget(self) -> xr.Dataset | xu.UgridDataset: npf: Npf | None = field(block="packages", default=None) sto: Sto | None = field(block="packages", default=None) buy: Buy | None = field(block="packages", default=None) - chd: list[Union[Chd, Chdg]] = field(block="packages", default=attrs.Factory(list)) - drn: list[Union[Drn, Drng]] = field(block="packages", default=attrs.Factory(list)) - evt: list[Union[Evt, Evta]] = field(block="packages", default=attrs.Factory(list)) - ghb: list[Union[Ghb, Ghbg]] = field(block="packages", default=attrs.Factory(list)) - rch: list[Union[Rch, Rcha]] = field(block="packages", default=attrs.Factory(list)) - riv: list[Union[Riv, Rivg]] = field(block="packages", default=attrs.Factory(list)) - csub: list[Csub] = field(block="packages", default=attrs.Factory(list)) - lak: list[Lak] = field(block="packages", default=attrs.Factory(list)) + chd: list[Union[Chd, Chdg]] = field(block="packages", default_factory=list) + drn: list[Union[Drn, Drng]] = field(block="packages", default_factory=list) + evt: list[Union[Evt, Evta]] = field(block="packages", default_factory=list) + ghb: list[Union[Ghb, Ghbg]] = field(block="packages", default_factory=list) + rch: list[Union[Rch, Rcha]] = field(block="packages", default_factory=list) + riv: list[Union[Riv, Rivg]] = field(block="packages", default_factory=list) + csub: list[Csub] = field(block="packages", default_factory=list) + lak: list[Lak] = field(block="packages", default_factory=list) mvr: Mvr | None = field(block="packages", default=None) vsc: Vsc | None = field(block="packages", default=None) - wel: list[Union[Wel, Welg]] = field(block="packages", default=attrs.Factory(list)) - output: Output = attrs.field( - default=attrs.Factory(lambda self: Gwf.Output(self), takes_self=True) - ) + wel: list[Union[Wel, Welg]] = field(block="packages", default_factory=list) + # Needs the instance to build, which default_factory can't see, so + # it's Optional and filled in by __post_init__ below (same pattern as + # Component.name). + output: Optional[Output] = Field(default=None, repr=False) + + def __post_init__(self, dims: Optional[dict] = None): + super().__post_init__(dims) + if self.output is None: + self.output = Gwf.Output(self) @property def grid(self) -> Grid: diff --git a/flopy4/mf6/gwf/api.py b/flopy4/mf6/gwf/api.py index 4eaadeca..80e9e397 100644 --- a/flopy4/mf6/gwf/api.py +++ b/flopy4/mf6/gwf/api.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Api(Package): dfn_name: ClassVar[str] = "gwf-api" diff --git a/flopy4/mf6/gwf/buy.py b/flopy4/mf6/gwf/buy.py index 17555763..f91d6c60 100644 --- a/flopy4/mf6/gwf/buy.py +++ b/flopy4/mf6/gwf/buy.py @@ -2,19 +2,20 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Buy(Package): dfn_name: ClassVar[str] = "gwf-buy" - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): irhospec: int = field(index=True, pk=True) drhodc: float = field() @@ -48,7 +49,7 @@ class Packagedata(Item): block="dimensions", longname="number of species used in density equation of state", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", auto_from="packagedata", diff --git a/flopy4/mf6/gwf/chd.py b/flopy4/mf6/gwf/chd.py index 4090d6ab..83cf4ffb 100644 --- a/flopy4/mf6/gwf/chd.py +++ b/flopy4/mf6/gwf/chd.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Chd(Package): dfn_name: ClassVar[str] = "gwf-chd" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) head: Union[float, str] = field(time_series=True) @@ -74,7 +75,7 @@ class StressPeriodData(Item): direction="in", keyword="obs6", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/chdg.py b/flopy4/mf6/gwf/chdg.py index 4ecb8975..393eb0ac 100644 --- a/flopy4/mf6/gwf/chdg.py +++ b/flopy4/mf6/gwf/chdg.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Chdg(Package): dfn_name: ClassVar[str] = "gwf-chdg" diff --git a/flopy4/mf6/gwf/csub.py b/flopy4/mf6/gwf/csub.py index f8ac854f..2aa6b79b 100644 --- a/flopy4/mf6/gwf/csub.py +++ b/flopy4/mf6/gwf/csub.py @@ -2,19 +2,20 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Csub(Package): dfn_name: ClassVar[str] = "gwf-csub" - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): icsubno: int = field(index=True, pk=True) cellid: tuple = field(cellid=True) @@ -30,7 +31,7 @@ class Packagedata(Item): aux: tuple = () boundname: Optional[str] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) sig0: Union[float, str] = field(time_series=True) @@ -242,7 +243,7 @@ class StressPeriodData(Item): optional=True, longname="maximum number of stress offset cells", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", auto_from="packagedata", @@ -277,7 +278,7 @@ class StressPeriodData(Item): optional=True, longname="specific gravity of saturated sediments", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/dis.py b/flopy4/mf6/gwf/dis.py index feb7ce7e..a1c61db1 100644 --- a/flopy4/mf6/gwf/dis.py +++ b/flopy4/mf6/gwf/dis.py @@ -1,18 +1,19 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import StructuredGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Dis(DisBase): dfn_name: ClassVar[str] = "gwf-dis" @@ -31,11 +32,11 @@ class Dis(DisBase): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=1, block="dimensions") ncol: int = field(default=2, block="dimensions") nrow: int = field(default=2, block="dimensions") - delr: NDArray[np.float64] = field( + delr: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a row", block="griddata", @@ -43,7 +44,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - delc: NDArray[np.float64] = field( + delc: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a column", block="griddata", @@ -51,7 +52,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=1.0, longname="cell top elevation", block="griddata", @@ -59,7 +60,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=0.0, longname="cell bottom elevation", block="griddata", @@ -76,12 +77,12 @@ class Dis(DisBase): longname="idomain existence array", ) - def __attrs_post_init__(self): + def __post_init__(self, dims: Optional[dict] = None): self.nodes = self.ncol * self.nrow * self.nlay self.ncpl = self.ncol * self.nrow self.nvert = (self.ncol + 1) * (self.nrow + 1) self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwf/disbase.py b/flopy4/mf6/gwf/disbase.py index 00352e89..6503b8cd 100644 --- a/flopy4/mf6/gwf/disbase.py +++ b/flopy4/mf6/gwf/disbase.py @@ -1,56 +1,65 @@ from pathlib import Path from typing import Optional -import attrs import numpy as np from flopy.discretization.grid import Grid as LegacyGrid +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6.constants import MF6 from flopy4.mf6.package import _DTYPE_MAP as _PKG_DTYPE_MAP -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package +from flopy4.mf6.spec import to_field_type from flopy4.mf6.write_context import WriteContext +from flopy4.spec import field_meta, pydantic_fields -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class DisBase(Package): # Derived dimensions — not read/written by the codec, set by subclass post_init. - nlay: Optional[int] = attrs.field(default=None, init=False) - nrow: Optional[int] = attrs.field(default=None, init=False) - ncol: Optional[int] = attrs.field(default=None, init=False) - ncpl: Optional[int] = attrs.field(default=None, init=False) - nvert: Optional[int] = attrs.field(default=None, init=False) - nodes: Optional[int] = attrs.field(default=None, init=False) + nlay: Optional[int] = Field(default=None, init=False) + nrow: Optional[int] = Field(default=None, init=False) + ncol: Optional[int] = Field(default=None, init=False) + ncpl: Optional[int] = Field(default=None, init=False) + nvert: Optional[int] = Field(default=None, init=False) + nodes: Optional[int] = Field(default=None, init=False) - def __attrs_post_init__(self): - super().__attrs_post_init__() + def __post_init__(self, dims: Optional[dict] = None): + super().__post_init__(dims) def _coerce_griddata(self) -> None: """Coerce griddata fields: list→ndarray, per-layer expansion, flatten. Must be called after derived dimensions (nodes, ncpl) are set and - before _broadcast_griddata / super().__attrs_post_init__(). + before _broadcast_griddata / super().__post_init__(). By this + point, any scalar/list griddata value has already passed through + Package._coerce_arrays (inherited -- a field_validator("*", + mode="before") runs during construction, ahead of __post_init__), + so it already arrives here as a real ndarray of the right dtype; + the list/tuple branch below only matters for a caller that + bypasses construction-time validation via a direct __dict__ write, + as several places in this codebase do. """ - import attrs as _attrs - - fields = _attrs.fields(type(self)) + fields = pydantic_fields(type(self)) dims = self.get_dims() ncpl = dims.get("ncpl", 0) nlay = dims.get("nlay", 1) - for f in fields: - if f.metadata.get("block") != "griddata": + for name, f in fields.items(): + meta = field_meta(f) + if meta.get("block") != "griddata": continue - val = self.__dict__.get(f.name) + val = self.__dict__.get(name) if val is None: continue - dtype = _PKG_DTYPE_MAP.get(f.metadata.get("dfn_type", "double"), np.float64) + dtype = _PKG_DTYPE_MAP.get(to_field_type(f.annotation), np.float64) if isinstance(val, (list, tuple)): val = np.asarray(val, dtype=dtype) - self.__dict__[f.name] = val + self.__dict__[name] = val if isinstance(val, np.ndarray): - if f.metadata.get("layered") and val.size == nlay and nlay > 0 and ncpl > 0: - self.__dict__[f.name] = np.repeat(val, ncpl).astype(dtype) + if meta.get("layered") and val.size == nlay and nlay > 0 and ncpl > 0: + self.__dict__[name] = np.repeat(val, ncpl).astype(dtype) elif val.ndim > 1: - self.__dict__[f.name] = val.ravel() + self.__dict__[name] = val.ravel() self._broadcast_griddata(fields, dims) def write(self, format: str = MF6, context: Optional[WriteContext] = None) -> None: @@ -59,7 +68,7 @@ def write(self, format: str = MF6, context: Optional[WriteContext] = None) -> No ncf = getattr(self, "ncf", None) if ncf is not None: if getattr(self, "ncf6_filerecord", None) is None and ncf.filename is not None: - setattr(self, "ncf6_filerecord", Path(Path(ncf.filename).name)) + setattr(self, "ncf6_filerecord", Path(ncf.filename.name)) super().write(format=format, context=context) if ncf is not None: # NCF lat/lon coordinate arrays require full float64 precision. diff --git a/flopy4/mf6/gwf/disv.py b/flopy4/mf6/gwf/disv.py index 3cbdc025..f565b3f3 100644 --- a/flopy4/mf6/gwf/disv.py +++ b/flopy4/mf6/gwf/disv.py @@ -1,31 +1,32 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field, field_validator +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.item import Item from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import VertexGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Disv(DisBase): dfn_name: ClassVar[str] = "gwf-disv" - @attrs.define(slots=False) + @dataclass(config=CFG) class Cell2dRecord: - icell2d: int = attrs.field() - xc: float = attrs.field() - yc: float = attrs.field() - ncvert: int = attrs.field() - icvert: tuple[int, ...] = attrs.field() + icell2d: int = Field() + xc: float = Field() + yc: float = Field() + ncvert: int = Field() + icvert: tuple[int, ...] = Field() - @attrs.define + @dataclass(config=CFG) class Vertices(Item): iv: int xv: float @@ -46,11 +47,11 @@ class Vertices(Item): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=0, block="dimensions") ncpl: int = field(default=0, block="dimensions") nvert: int = field(default=0, block="dimensions") - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=None, longname="model top elevation", block="griddata", @@ -58,7 +59,7 @@ class Vertices(Item): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=None, longname="model bottom elevation", block="griddata", @@ -74,20 +75,31 @@ class Vertices(Item): layered=True, netcdf=True, ) - iv: Optional[NDArray[np.int64]] = attrs.field(default=None) - xv: Optional[NDArray[np.float64]] = attrs.field(default=None) - yv: Optional[NDArray[np.float64]] = attrs.field(default=None) + iv: Optional[NDArray[np.int64]] = Field(default=None) + xv: Optional[NDArray[np.float64]] = Field(default=None) + yv: Optional[NDArray[np.float64]] = Field(default=None) + + # iv/xv/yv are declared NDArray-typed but commonly constructed from a + # plain list/tuple (see from_grid() below) -- unlike Package's own + # griddata fields, these carry no block="griddata"/shape= metadata, so + # Package._coerce_arrays' shape-driven check doesn't reach them, so + # they get their own mode="before" coercion (as do Tdis.perlen/nstp/ + # tsmult). + @field_validator("iv", mode="before") + @classmethod + def _coerce_iv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.int64) + + @field_validator("xv", "yv", mode="before") + @classmethod + def _coerce_xv_yv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.float64) + vertices: Optional[list[Vertices]] = field(default=None, block="vertices") - cell2ddata: Optional[list] = attrs.field(default=None) + cell2ddata: Optional[list] = Field(default=None) cell2d: Optional[list] = field(default=None, init=False, block="cell2d") - def __attrs_post_init__(self): - if self.iv is not None and (not isinstance(self.iv, np.ndarray)): - object.__setattr__(self, "iv", np.asarray(self.iv, dtype=np.int64)) - if self.xv is not None and (not isinstance(self.xv, np.ndarray)): - object.__setattr__(self, "xv", np.asarray(self.xv, dtype=np.float64)) - if self.yv is not None and (not isinstance(self.yv, np.ndarray)): - object.__setattr__(self, "yv", np.asarray(self.yv, dtype=np.float64)) + def __post_init__(self, dims: Optional[dict] = None): if self.iv is not None and self.xv is not None and (self.yv is not None): rows = [ self.Vertices(iv=int(iv) + 1, xv=float(xv), yv=float(yv)) @@ -95,18 +107,18 @@ def __attrs_post_init__(self): ] object.__setattr__(self, "vertices", rows) if self.cell2ddata is not None: - rows = [] + cell_rows = [] for rec in self.cell2ddata: row = (rec.icell2d + 1, rec.xc, rec.yc, rec.ncvert) + tuple( (v + 1 for v in rec.icvert) ) - rows.append(row) - object.__setattr__(self, "cell2d", rows) + cell_rows.append(row) + object.__setattr__(self, "cell2d", cell_rows) self.nodes = self.ncpl * self.nlay self.nrow = 0 self.ncol = 0 self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwf/drn.py b/flopy4/mf6/gwf/drn.py index 3847a511..23725547 100644 --- a/flopy4/mf6/gwf/drn.py +++ b/flopy4/mf6/gwf/drn.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Drn(Package): dfn_name: ClassVar[str] = "gwf-drn" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) elev: Union[float, str] = field(time_series=True) @@ -86,7 +87,7 @@ class StressPeriodData(Item): block="options", optional=True, ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/drng.py b/flopy4/mf6/gwf/drng.py index c50de963..98e6720e 100644 --- a/flopy4/mf6/gwf/drng.py +++ b/flopy4/mf6/gwf/drng.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Drng(Package): dfn_name: ClassVar[str] = "gwf-drng" diff --git a/flopy4/mf6/gwf/evt.py b/flopy4/mf6/gwf/evt.py index dadcd33e..7b3e4164 100644 --- a/flopy4/mf6/gwf/evt.py +++ b/flopy4/mf6/gwf/evt.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Evt(Package): dfn_name: ClassVar[str] = "gwf-evt" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) surface: Union[float, str] = field(time_series=True) @@ -96,7 +97,7 @@ class StressPeriodData(Item): block="dimensions", longname="number of ET segments", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/evta.py b/flopy4/mf6/gwf/evta.py index bf3f5d19..dabb4085 100644 --- a/flopy4/mf6/gwf/evta.py +++ b/flopy4/mf6/gwf/evta.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, IntArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Evta(Package): dfn_name: ClassVar[str] = "gwf-evta" diff --git a/flopy4/mf6/gwf/ghb.py b/flopy4/mf6/gwf/ghb.py index c76c741f..019698c3 100644 --- a/flopy4/mf6/gwf/ghb.py +++ b/flopy4/mf6/gwf/ghb.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ghb(Package): dfn_name: ClassVar[str] = "gwf-ghb" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) bhead: Union[float, str] = field(time_series=True) @@ -80,7 +81,7 @@ class StressPeriodData(Item): block="options", optional=True, ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/ghbg.py b/flopy4/mf6/gwf/ghbg.py index 8a0714ff..9e623f9c 100644 --- a/flopy4/mf6/gwf/ghbg.py +++ b/flopy4/mf6/gwf/ghbg.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ghbg(Package): dfn_name: ClassVar[str] = "gwf-ghbg" diff --git a/flopy4/mf6/gwf/ic.py b/flopy4/mf6/gwf/ic.py index f0d9d779..4665459f 100644 --- a/flopy4/mf6/gwf/ic.py +++ b/flopy4/mf6/gwf/ic.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ic(Package): dfn_name: ClassVar[str] = "gwf-ic" diff --git a/flopy4/mf6/gwf/lak.py b/flopy4/mf6/gwf/lak.py index e8ab7424..ec53f592 100644 --- a/flopy4/mf6/gwf/lak.py +++ b/flopy4/mf6/gwf/lak.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Lak(Package): dfn_name: ClassVar[str] = "gwf-lak" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): ifno: int = field(index=True, pk=True) strt: float = field() @@ -24,7 +25,7 @@ class Packagedata(Item): aux: tuple = () boundname: Optional[str] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class Connectiondata(Item): ifno: int = field(index=True, fk="packagedata.ifno") iconn: int = field(index=True) @@ -36,12 +37,12 @@ class Connectiondata(Item): connlen: float = field() connwidth: float = field() - @attrs.define + @dataclass(config=CFG) class Tables(Item): ifno: int = field(index=True, fk="packagedata.ifno") tab6_filename: Path = path(converter=Path, direction="in", keyword="tab6") - @attrs.define + @dataclass(config=CFG) class Outlets(Item): outletno: int = field(index=True, pk=True) lakein: int = field(index=True, fk="packagedata.ifno") @@ -52,79 +53,79 @@ class Outlets(Item): rough: Union[float, str] = field(time_series=True) slope: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Status(Item): _keyword: ClassVar[str] = "status" lakeno: int = field(index=True, fk="packagedata.ifno") status: Union[float, str] = field() - @attrs.define + @dataclass(config=CFG) class Stage(Item): _keyword: ClassVar[str] = "stage" lakeno: int = field(index=True, fk="packagedata.ifno") stage: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Rainfall(Item): _keyword: ClassVar[str] = "rainfall" lakeno: int = field(index=True, fk="packagedata.ifno") rainfall: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Evaporation(Item): _keyword: ClassVar[str] = "evaporation" lakeno: int = field(index=True, fk="packagedata.ifno") evaporation: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Runoff(Item): _keyword: ClassVar[str] = "runoff" lakeno: int = field(index=True, fk="packagedata.ifno") runoff: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Inflow(Item): _keyword: ClassVar[str] = "inflow" lakeno: int = field(index=True, fk="packagedata.ifno") inflow: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Withdrawal(Item): _keyword: ClassVar[str] = "withdrawal" lakeno: int = field(index=True, fk="packagedata.ifno") withdrawal: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Rate(Item): _keyword: ClassVar[str] = "rate" outletno: int = field(index=True, fk="outlets.outletno") rate: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Invert(Item): _keyword: ClassVar[str] = "invert" outletno: int = field(index=True, fk="outlets.outletno") invert: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Width(Item): _keyword: ClassVar[str] = "width" outletno: int = field(index=True, fk="outlets.outletno") width: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Slope(Item): _keyword: ClassVar[str] = "slope" outletno: int = field(index=True, fk="outlets.outletno") slope: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Rough(Item): _keyword: ClassVar[str] = "rough" outletno: int = field(index=True, fk="outlets.outletno") rough: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Auxiliary(Item): _keyword: ClassVar[str] = "auxiliary" lakeno: int = field(index=True, fk="packagedata.ifno") @@ -286,27 +287,27 @@ class Auxiliary(Item): block="dimensions", longname="number of tables", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", auto_from="packagedata", ) - connectiondata: Optional[list[Connectiondata]] = field( + connectiondata: Optional[SkipValidation[list[Connectiondata]]] = field( default=None, block="connectiondata", write_if_empty=True, ) - tables: Optional[list[Tables]] = field( + tables: Optional[SkipValidation[list[Tables]]] = field( default=None, block="tables", auto_from="tables", ) - outlets: Optional[list[Outlets]] = field( + outlets: Optional[SkipValidation[list[Outlets]]] = field( default=None, block="outlets", auto_from="outlets", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/mvr.py b/flopy4/mf6/gwf/mvr.py index 65d4a195..e4cd2a72 100644 --- a/flopy4/mf6/gwf/mvr.py +++ b/flopy4/mf6/gwf/mvr.py @@ -2,24 +2,25 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Mvr(Package): dfn_name: ClassVar[str] = "gwf-mvr" - @attrs.define + @dataclass(config=CFG) class Packages(Item): pname: Union[float, str] = field() mname: Optional[Union[float, str]] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): pname1: Union[float, str] = field() id1: int = field(index=True) @@ -75,12 +76,12 @@ class StressPeriodData(Item): block="dimensions", longname="number of packages to be used with the mover", ) - packages: Optional[list[Packages]] = field( + packages: Optional[SkipValidation[list[Packages]]] = field( default=None, block="packages", auto_from="packages", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/npf.py b/flopy4/mf6/gwf/npf.py index 493d5c33..29805daf 100644 --- a/flopy4/mf6/gwf/npf.py +++ b/flopy4/mf6/gwf/npf.py @@ -2,34 +2,35 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, IntArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Npf(Package): dfn_name: ClassVar[str] = "gwf-npf" - @attrs.define + @dataclass(config=CFG) class Cvoptions(Record): _keyword: ClassVar[str] = "variablecv" - dewatered: Optional[bool] = attrs.field(default=None, metadata={"tagged": True}) + dewatered: Optional[bool] = Field(default=None, json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Rewet(Record): _keyword: ClassVar[str] = "rewet" - wetfct: float = attrs.field(metadata={"tagged": True}) - iwetit: int = attrs.field(metadata={"tagged": True}) - ihdwet: int = attrs.field(metadata={"tagged": True}) + wetfct: float = Field(json_schema_extra={"tagged": True}) + iwetit: int = Field(json_schema_extra={"tagged": True}) + ihdwet: int = Field(json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Xt3doptions(Record): _keyword: ClassVar[str] = "xt3d" - rhs: Optional[bool] = attrs.field(default=None, metadata={"tagged": True}) + rhs: Optional[bool] = Field(default=None, json_schema_extra={"tagged": True}) save_flows: bool = field( default=False, diff --git a/flopy4/mf6/gwf/oc.py b/flopy4/mf6/gwf/oc.py index bb74d9b7..3ad47903 100644 --- a/flopy4/mf6/gwf/oc.py +++ b/flopy4/mf6/gwf/oc.py @@ -2,62 +2,63 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import Field, SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Oc(Package): dfn_name: ClassVar[str] = "gwf-oc" - @attrs.define + @dataclass(config=CFG) class Format(Record): _keyword: ClassVar[str] = "" - format_: str = attrs.field() - columns: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - width: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - digits: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) + format_: str = Field() + columns: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + width: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + digits: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Headprint(Record): _keyword: ClassVar[str] = "head" _extra_tokens: ClassVar[tuple[str, ...]] = ("PRINT_FORMAT",) - formatrecord: "Oc.Format" = attrs.field() + formatrecord: "Oc.Format" = Field() - @attrs.define + @dataclass(config=CFG) class All(Item): _keyword: ClassVar[str] = "all" - @attrs.define + @dataclass(config=CFG) class First(Item): _keyword: ClassVar[str] = "first" - @attrs.define + @dataclass(config=CFG) class Last(Item): _keyword: ClassVar[str] = "last" - @attrs.define + @dataclass(config=CFG) class Frequency(Item): _keyword: ClassVar[str] = "frequency" frequency: int = field() - @attrs.define + @dataclass(config=CFG) class Steps(Item): _keyword: ClassVar[str] = "steps" steps: tuple = field(default=(), array=True) - @attrs.define + @dataclass(config=CFG) class Save(Item): _keyword: ClassVar[str] = "save" rtype: Union[float, str] = field() ocsetting: "Oc.All | Oc.First | Oc.Last | Oc.Frequency | Oc.Steps" = field() - @attrs.define + @dataclass(config=CFG) class Print(Item): _keyword: ClassVar[str] = "print" rtype: Union[float, str] = field() @@ -93,7 +94,7 @@ class Print(Item): default=None, block="options", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/rch.py b/flopy4/mf6/gwf/rch.py index 9b4aa851..9ad7445e 100644 --- a/flopy4/mf6/gwf/rch.py +++ b/flopy4/mf6/gwf/rch.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Rch(Package): dfn_name: ClassVar[str] = "gwf-rch" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) recharge: Union[float, str] = field(time_series=True) @@ -80,7 +81,7 @@ class StressPeriodData(Item): direction="in", keyword="obs6", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/rcha.py b/flopy4/mf6/gwf/rcha.py index 85ce1baa..b4044fd8 100644 --- a/flopy4/mf6/gwf/rcha.py +++ b/flopy4/mf6/gwf/rcha.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, IntArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Rcha(Package): dfn_name: ClassVar[str] = "gwf-rcha" diff --git a/flopy4/mf6/gwf/riv.py b/flopy4/mf6/gwf/riv.py index e7e32a17..3c06295e 100644 --- a/flopy4/mf6/gwf/riv.py +++ b/flopy4/mf6/gwf/riv.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Riv(Package): dfn_name: ClassVar[str] = "gwf-riv" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) stage: Union[float, str] = field(time_series=True) @@ -81,7 +82,7 @@ class StressPeriodData(Item): block="options", optional=True, ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/rivg.py b/flopy4/mf6/gwf/rivg.py index 91531256..b13b6935 100644 --- a/flopy4/mf6/gwf/rivg.py +++ b/flopy4/mf6/gwf/rivg.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Rivg(Package): dfn_name: ClassVar[str] = "gwf-rivg" diff --git a/flopy4/mf6/gwf/sto.py b/flopy4/mf6/gwf/sto.py index c837579c..018ebbc9 100644 --- a/flopy4/mf6/gwf/sto.py +++ b/flopy4/mf6/gwf/sto.py @@ -2,19 +2,20 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, IntArrayLike, _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Sto(Package): dfn_name: ClassVar[str] = "gwf-sto" - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): storage: str = field() @@ -77,7 +78,7 @@ class StressPeriodData(Item): netcdf=True, longname="specific yield", ) # type: ignore[assignment] - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/vsc.py b/flopy4/mf6/gwf/vsc.py index e76ffaa0..ada93fc6 100644 --- a/flopy4/mf6/gwf/vsc.py +++ b/flopy4/mf6/gwf/vsc.py @@ -2,19 +2,20 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Vsc(Package): dfn_name: ClassVar[str] = "gwf-vsc" - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): iviscspec: int = field(index=True, pk=True) dviscdc: float = field() @@ -72,7 +73,7 @@ class Packagedata(Item): block="dimensions", longname="number of species used in viscosity equation of state", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", auto_from="packagedata", diff --git a/flopy4/mf6/gwf/wel.py b/flopy4/mf6/gwf/wel.py index 5803aaaf..4123a8bc 100644 --- a/flopy4/mf6/gwf/wel.py +++ b/flopy4/mf6/gwf/wel.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Wel(Package): dfn_name: ClassVar[str] = "gwf-wel" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) q: Union[float, str] = field(time_series=True) @@ -105,7 +106,7 @@ class StressPeriodData(Item): block="options", optional=True, ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwf/welg.py b/flopy4/mf6/gwf/welg.py index 86c51423..05c040b9 100644 --- a/flopy4/mf6/gwf/welg.py +++ b/flopy4/mf6/gwf/welg.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Welg(Package): dfn_name: ClassVar[str] = "gwf-welg" diff --git a/flopy4/mf6/gwt/__init__.py b/flopy4/mf6/gwt/__init__.py index d329104b..8b4cdda3 100644 --- a/flopy4/mf6/gwt/__init__.py +++ b/flopy4/mf6/gwt/__init__.py @@ -1,9 +1,9 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs from flopy.discretization.structuredgrid import StructuredGrid from flopy.discretization.vertexgrid import VertexGrid +from pydantic.dataclasses import dataclass from flopy4.mf6.gwf.disbase import DisBase from flopy4.mf6.gwt.adv import Adv @@ -19,7 +19,7 @@ from flopy4.mf6.gwt.oc import Oc from flopy4.mf6.gwt.src import Src from flopy4.mf6.gwt.ssm import Ssm -from flopy4.mf6.model import Model +from flopy4.mf6.model import CFG, Model from flopy4.mf6.spec import field, path from flopy4.utils import to_path @@ -52,7 +52,7 @@ def convert_grid(value): ] -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Gwt(Model): dfn_name: ClassVar[str] = "gwt-nam" @@ -88,9 +88,9 @@ class Gwt(Model): adv: Adv | None = field(block="packages", default=None) dsp: Dsp | None = field(block="packages", default=None) mst: Mst | None = field(block="packages", default=None) - cnc: list[Cnc] = field(block="packages", default=attrs.Factory(list)) - src: list[Src] = field(block="packages", default=attrs.Factory(list)) - lkt: list[Lkt] = field(block="packages", default=attrs.Factory(list)) + cnc: list[Cnc] = field(block="packages", default_factory=list) + src: list[Src] = field(block="packages", default_factory=list) + lkt: list[Lkt] = field(block="packages", default_factory=list) ssm: Ssm | None = field(block="packages", default=None) mvt: Mvt | None = field(block="packages", default=None) api: Api | None = field(block="packages", default=None) diff --git a/flopy4/mf6/gwt/adv.py b/flopy4/mf6/gwt/adv.py index 3fbeb8fc..54953143 100644 --- a/flopy4/mf6/gwt/adv.py +++ b/flopy4/mf6/gwt/adv.py @@ -1,13 +1,13 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Adv(Package): dfn_name: ClassVar[str] = "gwt-adv" diff --git a/flopy4/mf6/gwt/api.py b/flopy4/mf6/gwt/api.py index 6bbfc5c5..e2b360ca 100644 --- a/flopy4/mf6/gwt/api.py +++ b/flopy4/mf6/gwt/api.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Api(Package): dfn_name: ClassVar[str] = "gwt-api" diff --git a/flopy4/mf6/gwt/cnc.py b/flopy4/mf6/gwt/cnc.py index c0f45f8e..9d301b64 100644 --- a/flopy4/mf6/gwt/cnc.py +++ b/flopy4/mf6/gwt/cnc.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Cnc(Package): dfn_name: ClassVar[str] = "gwt-cnc" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) conc: Union[float, str] = field(time_series=True) @@ -74,7 +75,7 @@ class StressPeriodData(Item): direction="in", keyword="obs6", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwt/dis.py b/flopy4/mf6/gwt/dis.py index e46e5f62..9bf9c720 100644 --- a/flopy4/mf6/gwt/dis.py +++ b/flopy4/mf6/gwt/dis.py @@ -1,18 +1,19 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import StructuredGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Dis(DisBase): dfn_name: ClassVar[str] = "gwt-dis" @@ -31,11 +32,11 @@ class Dis(DisBase): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=1, block="dimensions") ncol: int = field(default=2, block="dimensions") nrow: int = field(default=2, block="dimensions") - delr: NDArray[np.float64] = field( + delr: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a row", block="griddata", @@ -43,7 +44,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - delc: NDArray[np.float64] = field( + delc: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a column", block="griddata", @@ -51,7 +52,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=1.0, longname="cell top elevation", block="griddata", @@ -59,7 +60,7 @@ class Dis(DisBase): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=0.0, longname="cell bottom elevation", block="griddata", @@ -76,12 +77,12 @@ class Dis(DisBase): longname="idomain existence array", ) - def __attrs_post_init__(self): + def __post_init__(self, dims: Optional[dict] = None): self.nodes = self.ncol * self.nrow * self.nlay self.ncpl = self.ncol * self.nrow self.nvert = (self.ncol + 1) * (self.nrow + 1) self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwt/disv.py b/flopy4/mf6/gwt/disv.py index f231c710..c0d7d823 100644 --- a/flopy4/mf6/gwt/disv.py +++ b/flopy4/mf6/gwt/disv.py @@ -1,31 +1,32 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field, field_validator +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.item import Item from flopy4.mf6.spec import field, path from flopy4.mf6.utils.grid import VertexGrid from flopy4.mf6.utl.ncf import Ncf -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Disv(DisBase): dfn_name: ClassVar[str] = "gwt-disv" - @attrs.define(slots=False) + @dataclass(config=CFG) class Cell2dRecord: - icell2d: int = attrs.field() - xc: float = attrs.field() - yc: float = attrs.field() - ncvert: int = attrs.field() - icvert: tuple[int, ...] = attrs.field() + icell2d: int = Field() + xc: float = Field() + yc: float = Field() + ncvert: int = Field() + icvert: tuple[int, ...] = Field() - @attrs.define + @dataclass(config=CFG) class Vertices(Item): iv: int xv: float @@ -46,11 +47,11 @@ class Vertices(Item): direction="in", keyword="ncf6", ) - ncf: Optional[Ncf] = attrs.field(default=None) + ncf: Optional[Ncf] = Field(default=None) nlay: int = field(default=0, block="dimensions") ncpl: int = field(default=0, block="dimensions") nvert: int = field(default=0, block="dimensions") - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=None, longname="model top elevation", block="griddata", @@ -58,7 +59,7 @@ class Vertices(Item): layered=False, netcdf=True, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=None, longname="model bottom elevation", block="griddata", @@ -74,20 +75,31 @@ class Vertices(Item): layered=True, netcdf=True, ) - iv: Optional[NDArray[np.int64]] = attrs.field(default=None) - xv: Optional[NDArray[np.float64]] = attrs.field(default=None) - yv: Optional[NDArray[np.float64]] = attrs.field(default=None) + iv: Optional[NDArray[np.int64]] = Field(default=None) + xv: Optional[NDArray[np.float64]] = Field(default=None) + yv: Optional[NDArray[np.float64]] = Field(default=None) + + # iv/xv/yv are declared NDArray-typed but commonly constructed from a + # plain list/tuple (see from_grid() below) -- unlike Package's own + # griddata fields, these carry no block="griddata"/shape= metadata, so + # Package._coerce_arrays' shape-driven check doesn't reach them, so + # they get their own mode="before" coercion (as do Tdis.perlen/nstp/ + # tsmult). + @field_validator("iv", mode="before") + @classmethod + def _coerce_iv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.int64) + + @field_validator("xv", "yv", mode="before") + @classmethod + def _coerce_xv_yv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.float64) + vertices: Optional[list[Vertices]] = field(default=None, block="vertices") - cell2ddata: Optional[list] = attrs.field(default=None) + cell2ddata: Optional[list] = Field(default=None) cell2d: Optional[list] = field(default=None, init=False, block="cell2d") - def __attrs_post_init__(self): - if self.iv is not None and (not isinstance(self.iv, np.ndarray)): - object.__setattr__(self, "iv", np.asarray(self.iv, dtype=np.int64)) - if self.xv is not None and (not isinstance(self.xv, np.ndarray)): - object.__setattr__(self, "xv", np.asarray(self.xv, dtype=np.float64)) - if self.yv is not None and (not isinstance(self.yv, np.ndarray)): - object.__setattr__(self, "yv", np.asarray(self.yv, dtype=np.float64)) + def __post_init__(self, dims: Optional[dict] = None): if self.iv is not None and self.xv is not None and (self.yv is not None): rows = [ self.Vertices(iv=int(iv) + 1, xv=float(xv), yv=float(yv)) @@ -95,18 +107,18 @@ def __attrs_post_init__(self): ] object.__setattr__(self, "vertices", rows) if self.cell2ddata is not None: - rows = [] + cell_rows = [] for rec in self.cell2ddata: row = (rec.icell2d + 1, rec.xc, rec.yc, rec.ncvert) + tuple( (v + 1 for v in rec.icvert) ) - rows.append(row) - object.__setattr__(self, "cell2d", rows) + cell_rows.append(row) + object.__setattr__(self, "cell2d", cell_rows) self.nodes = self.ncpl * self.nlay self.nrow = 0 self.ncol = 0 self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/gwt/dsp.py b/flopy4/mf6/gwt/dsp.py index 09a99e61..bc118bb7 100644 --- a/flopy4/mf6/gwt/dsp.py +++ b/flopy4/mf6/gwt/dsp.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Dsp(Package): dfn_name: ClassVar[str] = "gwt-dsp" diff --git a/flopy4/mf6/gwt/fmi.py b/flopy4/mf6/gwt/fmi.py index 8165de69..dc3ce4db 100644 --- a/flopy4/mf6/gwt/fmi.py +++ b/flopy4/mf6/gwt/fmi.py @@ -2,18 +2,19 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Fmi(Package): dfn_name: ClassVar[str] = "gwt-fmi" - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): flowtype: Union[float, str] = field() fname: Path = path(converter=Path, direction="in") @@ -31,7 +32,7 @@ class Packagedata(Item): optional=True, longname="correct for flow imbalance", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", ) diff --git a/flopy4/mf6/gwt/ic.py b/flopy4/mf6/gwt/ic.py index 0908aa1b..e327069a 100644 --- a/flopy4/mf6/gwt/ic.py +++ b/flopy4/mf6/gwt/ic.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ic(Package): dfn_name: ClassVar[str] = "gwt-ic" diff --git a/flopy4/mf6/gwt/ist.py b/flopy4/mf6/gwt/ist.py index a363ce08..95799194 100644 --- a/flopy4/mf6/gwt/ist.py +++ b/flopy4/mf6/gwt/ist.py @@ -2,33 +2,34 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ist(Package): dfn_name: ClassVar[str] = "gwt-ist" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Format(Record): _keyword: ClassVar[str] = "" - format_: str = attrs.field() - columns: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - width: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - digits: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) + format_: str = Field() + columns: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + width: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + digits: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Cimprint(Record): _keyword: ClassVar[str] = "cim" _extra_tokens: ClassVar[tuple[str, ...]] = ("PRINT_FORMAT",) - formatrecord: "Ist.Format" = attrs.field() + formatrecord: "Ist.Format" = Field() save_flows: bool = field( default=False, diff --git a/flopy4/mf6/gwt/lkt.py b/flopy4/mf6/gwt/lkt.py index 066e9628..20500074 100644 --- a/flopy4/mf6/gwt/lkt.py +++ b/flopy4/mf6/gwt/lkt.py @@ -2,64 +2,65 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Lkt(Package): dfn_name: ClassVar[str] = "gwt-lkt" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): ifno: int = field(index=True, pk=True) strt: float = field() aux: tuple = () boundname: Optional[str] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class Status(Item): _keyword: ClassVar[str] = "status" ifno: int = field(index=True, fk="packagedata.ifno") status: Union[float, str] = field() - @attrs.define + @dataclass(config=CFG) class Concentration(Item): _keyword: ClassVar[str] = "concentration" ifno: int = field(index=True, fk="packagedata.ifno") concentration: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Rainfall(Item): _keyword: ClassVar[str] = "rainfall" ifno: int = field(index=True, fk="packagedata.ifno") rainfall: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Evaporation(Item): _keyword: ClassVar[str] = "evaporation" ifno: int = field(index=True, fk="packagedata.ifno") evaporation: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Runoff(Item): _keyword: ClassVar[str] = "runoff" ifno: int = field(index=True, fk="packagedata.ifno") runoff: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class ExtInflow(Item): _keyword: ClassVar[str] = "ext-inflow" ifno: int = field(index=True, fk="packagedata.ifno") ext_inflow: Union[float, str] = field(time_series=True) - @attrs.define + @dataclass(config=CFG) class Auxiliary(Item): _keyword: ClassVar[str] = "auxiliary" ifno: int = field(index=True, fk="packagedata.ifno") @@ -157,12 +158,12 @@ class Auxiliary(Item): direction="in", keyword="obs6", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", write_if_empty=True, ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwt/mst.py b/flopy4/mf6/gwt/mst.py index 46ce85f5..7e4541ce 100644 --- a/flopy4/mf6/gwt/mst.py +++ b/flopy4/mf6/gwt/mst.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Mst(Package): dfn_name: ClassVar[str] = "gwt-mst" diff --git a/flopy4/mf6/gwt/mvt.py b/flopy4/mf6/gwt/mvt.py index f5ff622a..d323239a 100644 --- a/flopy4/mf6/gwt/mvt.py +++ b/flopy4/mf6/gwt/mvt.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Mvt(Package): dfn_name: ClassVar[str] = "gwt-mvt" diff --git a/flopy4/mf6/gwt/oc.py b/flopy4/mf6/gwt/oc.py index 24ad976a..e3560079 100644 --- a/flopy4/mf6/gwt/oc.py +++ b/flopy4/mf6/gwt/oc.py @@ -2,62 +2,63 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import Field, SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Oc(Package): dfn_name: ClassVar[str] = "gwt-oc" - @attrs.define + @dataclass(config=CFG) class Format(Record): _keyword: ClassVar[str] = "" - format_: str = attrs.field() - columns: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - width: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) - digits: Optional[int] = attrs.field(default=None, metadata={"tagged": True}) + format_: str = Field() + columns: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + width: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) + digits: Optional[int] = Field(default=None, json_schema_extra={"tagged": True}) - @attrs.define + @dataclass(config=CFG) class Concentrationprint(Record): _keyword: ClassVar[str] = "concentration" _extra_tokens: ClassVar[tuple[str, ...]] = ("PRINT_FORMAT",) - formatrecord: "Oc.Format" = attrs.field() + formatrecord: "Oc.Format" = Field() - @attrs.define + @dataclass(config=CFG) class All(Item): _keyword: ClassVar[str] = "all" - @attrs.define + @dataclass(config=CFG) class First(Item): _keyword: ClassVar[str] = "first" - @attrs.define + @dataclass(config=CFG) class Last(Item): _keyword: ClassVar[str] = "last" - @attrs.define + @dataclass(config=CFG) class Frequency(Item): _keyword: ClassVar[str] = "frequency" frequency: int = field() - @attrs.define + @dataclass(config=CFG) class Steps(Item): _keyword: ClassVar[str] = "steps" steps: tuple = field(default=(), array=True) - @attrs.define + @dataclass(config=CFG) class Save(Item): _keyword: ClassVar[str] = "save" rtype: Union[float, str] = field() ocsetting: "Oc.All | Oc.First | Oc.Last | Oc.Frequency | Oc.Steps" = field() - @attrs.define + @dataclass(config=CFG) class Print(Item): _keyword: ClassVar[str] = "print" rtype: Union[float, str] = field() @@ -93,7 +94,7 @@ class Print(Item): default=None, block="options", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwt/src.py b/flopy4/mf6/gwt/src.py index af6e5aa8..562622d1 100644 --- a/flopy4/mf6/gwt/src.py +++ b/flopy4/mf6/gwt/src.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Src(Package): dfn_name: ClassVar[str] = "gwt-src" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class StressPeriodData(Item): cellid: tuple = field(cellid=True) smassrate: Union[float, str] = field(time_series=True) @@ -80,7 +81,7 @@ class StressPeriodData(Item): optional=True, longname="apply source to highest saturated cell", ) - _stress_period_data: Optional[dict[int, list[StressPeriodData]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[StressPeriodData]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/gwt/ssm.py b/flopy4/mf6/gwt/ssm.py index 17ba9bea..c5c8fe0d 100644 --- a/flopy4/mf6/gwt/ssm.py +++ b/flopy4/mf6/gwt/ssm.py @@ -2,24 +2,25 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ssm(Package): dfn_name: ClassVar[str] = "gwt-ssm" - @attrs.define + @dataclass(config=CFG) class Sources(Item): pname: Union[float, str] = field() srctype: Union[float, str] = field() auxname: Union[float, str] = field() - @attrs.define + @dataclass(config=CFG) class Fileinput(Item): pname: Union[float, str] = field() spc6_filename: Path = path(converter=Path, direction="in", keyword="spc6") @@ -37,12 +38,12 @@ class Fileinput(Item): optional=True, longname="save calculated flows to budget file", ) - sources: Optional[list[Sources]] = field( + sources: Optional[SkipValidation[list[Sources]]] = field( default=None, block="sources", write_if_empty=True, ) - fileinput: Optional[list[Fileinput]] = field( + fileinput: Optional[SkipValidation[list[Fileinput]]] = field( default=None, block="fileinput", ) diff --git a/flopy4/mf6/ims.py b/flopy4/mf6/ims.py index 318e4a99..a996ea6b 100644 --- a/flopy4/mf6/ims.py +++ b/flopy4/mf6/ims.py @@ -2,30 +2,31 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.record import Record -from flopy4.mf6.solution import Solution +from flopy4.mf6.solution import CFG, Solution from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ims(Solution): dfn_name: ClassVar[str] = "sln-ims" slntype: ClassVar[str] = "ims" - @attrs.define + @dataclass(config=CFG) class NoPtc(Record): _keyword: ClassVar[str] = "no_ptc" - no_ptc_option: Optional[str] = attrs.field(default=None) + no_ptc_option: Optional[str] = Field(default=None) - @attrs.define + @dataclass(config=CFG) class Rclose(Record): _keyword: ClassVar[str] = "" - inner_rclose: float = attrs.field(metadata={"tagged": True}) - rclose_option: Optional[str] = attrs.field(default=None) + inner_rclose: float = Field(json_schema_extra={"tagged": True}) + rclose_option: Optional[str] = Field(default=None) print_option: Optional[str] = field( default=None, diff --git a/flopy4/mf6/item.py b/flopy4/mf6/item.py index 1f491f44..1f1e69a5 100644 --- a/flopy4/mf6/item.py +++ b/flopy4/mf6/item.py @@ -13,14 +13,14 @@ by its own leading keyword token (STATUS/STAGE/RATE/...). """ +from __future__ import annotations + import re -from functools import lru_cache from pathlib import Path -from typing import Any, Union, cast, get_args, get_origin - -import attrs +from typing import Annotated, Any, Union, cast, get_args, get_origin -from flopy4.mf6.record import Record, _coerce, _resolve_sibling_class +from flopy4.mf6.record import Record, _coerce +from flopy4.spec import field_meta _AUX_KEY_RE = re.compile(r"^aux(\d+)$") @@ -41,44 +41,35 @@ def normalize_aux_keys(item: dict) -> dict: return rest -def _cellid_field(cls: type) -> attrs.Attribute | None: - return next((f for f in cast(type[Record], cls).fields() if f.metadata.get("cellid")), None) +def _cellid_field(cls: type) -> Any | None: + return next( + (f for f in cast(type[Record], cls).fields().values() if field_meta(f).get("cellid")), + None, + ) def _has_aux_field(cls: type) -> bool: - return any(f.name == "aux" for f in cast(type[Record], cls).fields()) + return "aux" in cast(type[Record], cls).fields() def _has_boundname_field(cls: type) -> bool: - return any(f.name == "boundname" for f in cast(type[Record], cls).fields()) + return "boundname" in cast(type[Record], cls).fields() -@lru_cache(maxsize=None) -def _nested_union_classes(cls: type, type_str: str) -> "tuple[type[Item], ...] | None": - """If a field's raw type annotation is a `` | ``-joined forward - reference to sibling Item classes (e.g. ``"Oc.All | Oc.First | ..."``, - see make.py's _build_arm_specs_from_union), return the tuple of - resolved classes; else None. Resolvability is itself the signal, same - as record.py's _nested_class. +def _is_item_union(annotation: Any) -> "tuple[type[Item], ...] | None": + """If a field's (already-resolved) annotation is a Union of sibling + Item classes (e.g. OC's ``All | First | Last | Frequency | Steps``), + return the tuple of arm classes; else None. - Cached since to_tokens/from_tokens call this per field, often - repeatedly while parsing many rows. + By the time `Record.fields()` has run, `annotation` (a pydantic + `FieldInfo.annotation`) is already the resolved union of classes, not + a forward-ref string. """ - names = [part.strip().rsplit(".", 1)[-1] for part in type_str.split(" | ")] - if len(names) < 2: - return None - resolved = [_resolve_sibling_class(cls, name) for name in names] - if any(not (isinstance(r, type) and issubclass(r, Item)) for r in resolved): - return None - return tuple(cast("list[type[Item]]", resolved)) - - -def _field_type_str(f: attrs.Attribute) -> "str | None": - """attrs stubs type `Attribute.type` as `type | None`, but attrs - actually stores the raw annotation there -- a string for a forward - reference (every nested-sibling-union field). None otherwise.""" - t: Any = f.type - return t if isinstance(t, str) else None + origin = get_origin(annotation) + if origin is Union or origin is type(int | str): + arms = tuple(a for a in get_args(annotation) if isinstance(a, type) and issubclass(a, Item)) + return arms or None + return None def construct_item(item_cls: type, values) -> "Item": @@ -89,34 +80,26 @@ def construct_item(item_cls: type, values) -> "Item": trailing string when the class also has boundname (always declared last) -- a string there unambiguously isn't a numeric aux value. """ - fields = cast(type[Record], item_cls).fields() + fields = list(cast(type[Record], item_cls).fields().items()) tuple_idx = next( ( i - for i, f in enumerate(fields) - if f.name == "aux" - or f.metadata.get("array") - or ( - (t := _field_type_str(f)) is not None - and _nested_union_classes(item_cls, t) is not None - ) + for i, (name, f) in enumerate(fields) + if name == "aux" + or field_meta(f).get("array") + or _is_item_union(f.annotation) is not None ), None, ) values = list(values) if tuple_idx is None: return cast("Item", item_cls(*values)) - nested_field = fields[tuple_idx] - nested_field_type = _field_type_str(nested_field) - arm_classes = ( - _nested_union_classes(item_cls, nested_field_type) - if nested_field_type is not None - else None - ) + _, nested_finfo = fields[tuple_idx] + arm_classes = _is_item_union(nested_finfo.annotation) boundname_val = None if ( fields - and fields[-1].name == "boundname" + and fields[-1][0] == "boundname" and len(values) > tuple_idx and isinstance(values[-1], str) and arm_classes is None @@ -138,13 +121,14 @@ def _n_fixed_tokens(cls: type) -> int: infer a variable-width cellid's element count from total token length.""" cls = cast(type[Record], cls) n = 1 if cls.keyword() else 0 - for f in cls.fields(): - if f.metadata.get("cellid") or f.name in ("aux", "boundname"): + for name, f in cls.fields().items(): + meta = field_meta(f) + if meta.get("cellid") or name in ("aux", "boundname"): continue - if f.metadata.get("optional"): + if meta.get("optional"): continue - n += 1 + (1 if f.metadata.get("_keyword") else 0) - if f.metadata.get("direction"): + n += 1 + (1 if meta.get("_keyword") else 0) + if meta.get("direction"): n += 1 return n @@ -216,27 +200,28 @@ def to_tokens(self) -> tuple: keyword = cls.keyword() row: list[Any] = [] keyword_emitted = not keyword - for f in fields: - if f.name in ("aux", "boundname"): + for name, f in fields.items(): + if name in ("aux", "boundname"): continue - val = getattr(self, f.name) + val = getattr(self, name) if val is None: continue - if f.metadata.get("cellid"): + meta = field_meta(f) + if meta.get("cellid"): row.extend(int(c) + 1 for c in val) - elif f.metadata.get("index"): + elif meta.get("index"): row.append(int(val) + 1) - elif f.metadata.get("array"): + elif meta.get("array"): if not keyword_emitted: row.append(keyword.upper()) keyword_emitted = True row.extend(val) - elif f.metadata.get("tagged"): + elif meta.get("tagged"): if not keyword_emitted: row.append(keyword.upper()) keyword_emitted = True if val: - row.append(f.name.upper()) + row.append(name.upper()) elif isinstance(val, Record): # Nested keystring-union field (OC's ocsetting) -- val is # already the resolved arm instance and knows how to @@ -249,11 +234,11 @@ def to_tokens(self) -> tuple: if not keyword_emitted: row.append(keyword.upper()) keyword_emitted = True - if file_kw := f.metadata.get("_keyword"): + if file_kw := meta.get("_keyword"): row.append(file_kw.upper()) - if direction := f.metadata.get("direction"): + if direction := meta.get("direction"): row.append("FILEOUT" if direction == "out" else "FILEIN") - row.append(str(val) if isinstance(val, Path) else val) + row.append(val.as_posix() if isinstance(val, Path) else val) if not keyword_emitted: row.append(keyword.upper()) aux = getattr(self, "aux", None) @@ -284,51 +269,51 @@ def from_tokens( # type: ignore[override] tok_idx = 0 n = len(tokens) - def consume(f: attrs.Attribute) -> None: + def consume(name: str, f: Any) -> None: nonlocal tok_idx, keyword_skipped - if f.metadata.get("cellid"): + meta = field_meta(f) + if meta.get("cellid"): cellid = tuple(int(tokens[tok_idx + j]) - 1 for j in range(ncelldim)) - kwargs[f.name] = cellid + kwargs[name] = cellid tok_idx += ncelldim return - if f.metadata.get("index"): - kwargs[f.name] = int(float(str(tokens[tok_idx]))) - 1 + if meta.get("index"): + kwargs[name] = int(float(str(tokens[tok_idx]))) - 1 tok_idx += 1 return if not keyword_skipped: tok_idx += 1 keyword_skipped = True - if f.metadata.get("_keyword"): + if meta.get("_keyword"): tok_idx += 1 - if f.metadata.get("direction"): + if meta.get("direction"): tok_idx += 1 if tok_idx >= n: return - kwargs[f.name] = _coerce(tokens[tok_idx], f) + kwargs[name] = _coerce(tokens[tok_idx], f) tok_idx += 1 - def width(f: attrs.Attribute) -> int: - w = 1 + (1 if f.metadata.get("_keyword") else 0) - if f.metadata.get("direction"): + def width(f: Any) -> int: + meta = field_meta(f) + w = 1 + (1 if meta.get("_keyword") else 0) + if meta.get("direction"): w += 1 return w - main_fields = [f for f in fields if f.name not in ("aux", "boundname")] + main_fields = [(name, f) for name, f in fields.items() if name not in ("aux", "boundname")] nested_union_fields = [ - f - for f in main_fields - if (t := _field_type_str(f)) is not None - # mypy false positive: type[Item] vs. Hashable (lru_cache arg) - and _nested_union_classes(cls, t) is not None # type: ignore[arg-type] + (name, f) for name, f in main_fields if _is_item_union(f.annotation) is not None ] - main_fields = [f for f in main_fields if f not in nested_union_fields] - array_fields = [f for f in main_fields if f.metadata.get("array")] - main_fields = [f for f in main_fields if not f.metadata.get("array")] - required_fields = [f for f in main_fields if not f.metadata.get("optional")] - optional_fields = [f for f in main_fields if f.metadata.get("optional")] + main_fields = [item for item in main_fields if item not in nested_union_fields] + array_fields = [(name, f) for name, f in main_fields if field_meta(f).get("array")] + main_fields = [item for item in main_fields if item not in array_fields] + required_fields = [ + (name, f) for name, f in main_fields if not field_meta(f).get("optional") + ] + optional_fields = [(name, f) for name, f in main_fields if field_meta(f).get("optional")] - for f in required_fields: - consume(f) + for name, f in required_fields: + consume(name, f) has_bn_token = False if has_boundname and n > tok_idx: @@ -336,10 +321,12 @@ def width(f: attrs.Attribute) -> int: has_bn_token = isinstance(last, str) and not _token_fits(last, float) remaining = n - tok_idx - (1 if has_bn_token else 0) - (naux if has_aux else 0) - budget_fields = [f for f in optional_fields if not f.metadata.get("tagged")] + budget_fields = [ + (name, f) for name, f in optional_fields if not field_meta(f).get("tagged") + ] n_opt_present = 0 used = 0 - for f in budget_fields: + for name, f in budget_fields: w = width(f) if used + w > remaining: break @@ -347,20 +334,21 @@ def width(f: attrs.Attribute) -> int: n_opt_present += 1 budget_idx = 0 - for f in optional_fields: - if f.metadata.get("tagged"): + for name, f in optional_fields: + meta = field_meta(f) + if meta.get("tagged"): if not keyword_skipped: tok_idx += 1 keyword_skipped = True - kw = f.name.upper() + kw = name.upper() if tok_idx < n and str(tokens[tok_idx]).upper() == kw: - kwargs[f.name] = str(tokens[tok_idx]) + kwargs[name] = str(tokens[tok_idx]) tok_idx += 1 continue present = budget_idx < n_opt_present budget_idx += 1 if present: - consume(f) + consume(name, f) if array_fields: # Consumes everything left up to aux/boundname's own reserved @@ -371,7 +359,7 @@ def width(f: attrs.Attribute) -> int: if not keyword_skipped: tok_idx += 1 keyword_skipped = True - f = array_fields[0] + name, _f = array_fields[0] end = n - (1 if has_bn_token else 0) - (naux if has_aux else 0) vals = [] while tok_idx < end: @@ -381,7 +369,7 @@ def width(f: attrs.Attribute) -> int: except (ValueError, TypeError): vals.append(tok) tok_idx += 1 - kwargs[f.name] = tuple(vals) + kwargs[name] = tuple(vals) elif nested_union_fields: # Nested keystring-union field (OC's ocsetting) -- same span # logic as array_fields above, but dispatches a typed arm @@ -389,20 +377,18 @@ def width(f: attrs.Attribute) -> int: if not keyword_skipped: tok_idx += 1 keyword_skipped = True - f = nested_union_fields[0] - nested_field_type = _field_type_str(f) - assert nested_field_type is not None - arm_classes = _nested_union_classes(cls, nested_field_type) # type: ignore[arg-type] + name, f = nested_union_fields[0] + arm_classes = _is_item_union(f.annotation) assert arm_classes is not None end = n - (1 if has_bn_token else 0) - (naux if has_aux else 0) nested_tokens = list(tokens[tok_idx:end]) arm_cls = dispatch_union_item(nested_tokens, arm_classes) if arm_cls is None: raise ValueError( - f"{cls.__name__}.{f.name}: no matching arm in {arm_classes} " + f"{cls.__name__}.{name}: no matching arm in {arm_classes} " f"for tokens {nested_tokens}" ) - kwargs[f.name] = arm_cls.from_tokens(nested_tokens) + kwargs[name] = arm_cls.from_tokens(nested_tokens) tok_idx = end elif not keyword_skipped: tok_idx += 1 @@ -437,18 +423,35 @@ def _unwrap_item(item) -> "type[Item] | tuple[type[Item], ...] | None": return None +def _unwrap_skip_validation(t: Any) -> Any: + """Strip one `Annotated[X, SkipValidation()]` layer, if present. + + Item-list fields are pydantic.SkipValidation-wrapped (codegen emits + this -- see Package._init_item_lists' own docstring for why: pydantic + would otherwise validate an Item-list field's raw tuple/dict input + eagerly). get_origin() on the raw annotation returns Annotated, not + dict/list, so the unwrapping below needs this extra step. + """ + if get_origin(t) is Annotated: + return get_args(t)[0] + return t + + def item_list_type(field_type) -> "type[Item] | tuple[type[Item], ...] | None": - """For Optional[list[C]] or Optional[dict[int, list[C]]], return C (or - the tuple of arm classes for a Union item type).""" + """For Optional[list[C]] or Optional[dict[int, list[C]]] (each + optionally SkipValidation-wrapped), return C (or the tuple of arm + classes for a Union item type).""" args = get_args(field_type) inner = next((a for a in args if a is not type(None)), None) if inner is None: return None + inner = _unwrap_skip_validation(inner) origin = get_origin(inner) if origin is list: return _unwrap_item(get_args(inner)[0]) if origin is dict: _, val = get_args(inner) + val = _unwrap_skip_validation(val) if get_origin(val) is list: return _unwrap_item(get_args(val)[0]) return None diff --git a/flopy4/mf6/model.py b/flopy4/mf6/model.py index 40d1091a..6ee069c8 100644 --- a/flopy4/mf6/model.py +++ b/flopy4/mf6/model.py @@ -1,11 +1,11 @@ from abc import ABC -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.context import Context +from flopy4.mf6.context import CFG, Context -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Model(Context, ABC): def default_filename(self) -> str: return f"{self.name}.nam" # type: ignore diff --git a/flopy4/mf6/netcdf.py b/flopy4/mf6/netcdf.py index 026b0bed..2232531f 100644 --- a/flopy4/mf6/netcdf.py +++ b/flopy4/mf6/netcdf.py @@ -20,6 +20,7 @@ from flopy4.mf6.spec import to_field_type from flopy4.mf6.utils.grid import StructuredGrid, VertexGrid from flopy4.mf6.utils.time import Time +from flopy4.spec import field_meta, pydantic_fields from flopy4.version import __version__ @@ -79,19 +80,18 @@ def get_spec(package_name: str): class _PackageSpec: """Summarizes a package's netcdf-relevant array fields (dtype, dims, - metadata) from its attrs field metadata (block, shape, layered, ...).""" + metadata) from its field metadata (block, shape, layered, ...).""" _DTYPE_MAP = _PKG_DTYPE_MAP def __init__(self, cls): - import attrs as _attrs - class _ArrayInfo: - def __init__(self, f): + def __init__(self, name, f): + meta = field_meta(f) # A fill-forward (period) field's value has a leading nper axis. - fill_forward = bool(f.metadata.get("fill_forward")) - is_layered = f.metadata.get("layered", True) - raw_shape = f.metadata.get("shape") or ("nodes",) + fill_forward = bool(meta.get("fill_forward")) + is_layered = meta.get("layered", True) + raw_shape = meta.get("shape") or ("nodes",) # normalize ncpl → nodes for layered fields only if is_layered: normalized = tuple("nodes" if d == "ncpl" else d for d in raw_shape) @@ -101,18 +101,20 @@ def __init__(self, f): normalized = ("nper",) + normalized self.dtype = np.dtype( - _PackageSpec._DTYPE_MAP.get(to_field_type(f.type), np.float64) + _PackageSpec._DTYPE_MAP.get(to_field_type(f.annotation), np.float64) ) self.dims = normalized self.metadata = { - "longname": f.metadata.get("longname", f.name), - "block": f.metadata.get("block", ""), + "longname": meta.get("longname", name), + "block": meta.get("block", ""), "fill_forward": fill_forward, "netcdf": True, } self.arrays = { - f.name: _ArrayInfo(f) for f in _attrs.fields(cls) if f.metadata.get("netcdf") + name: _ArrayInfo(name, f) + for name, f in pydantic_fields(cls).items() + if field_meta(f).get("netcdf") } @@ -220,8 +222,6 @@ def from_model( "params": [], } - import attrs as _attrs - # compute total nodes for broadcasting scalars to full grid _dis = getattr(model, "dis", None) d = _dis.get_dims() if _dis is not None else {} @@ -233,24 +233,25 @@ def from_model( else: _nodes = d.get("nodes", _nlay) - for f in _attrs.fields(type(package)): - if not f.metadata.get("netcdf"): + for name, f in pydantic_fields(type(package)).items(): + meta = field_meta(f) + if not meta.get("netcdf"): continue - if f.metadata.get("block") == "griddata": - val = getattr(package, f.name) + if meta.get("block") == "griddata": + val = getattr(package, name) if val is None: continue arr = np.asarray(val, dtype=np.float64) # Only broadcast scalars to full grid for nodes-shaped fields - shape_meta = f.metadata.get("shape", ()) + shape_meta = meta.get("shape", ()) if "nodes" in shape_meta and arr.size < _nodes: arr = np.full(_nodes, float(arr.ravel()[0])) - p["params"].append({"name": f.name, "data": arr}) + p["params"].append({"name": name, "data": arr}) else: - val = getattr(package, f.name) + val = getattr(package, name) if val is None: continue - p["params"].append({"name": f.name, "data": np.asarray(val, dtype=np.float64)}) + p["params"].append({"name": name, "data": np.asarray(val, dtype=np.float64)}) if len(p["params"]) > 0: packages.append(p) diff --git a/flopy4/mf6/package.py b/flopy4/mf6/package.py index 6b3e75b5..dfc7db1f 100644 --- a/flopy4/mf6/package.py +++ b/flopy4/mf6/package.py @@ -1,12 +1,14 @@ from abc import ABC from pathlib import Path +from typing import Any, Optional -import attrs import numpy as np import pandas as pd import xarray as xr +from pydantic import field_validator +from pydantic.dataclasses import dataclass -from flopy4.mf6.component import Component +from flopy4.mf6.component import CFG, Component from flopy4.mf6.item import ( Item, construct_item, @@ -15,6 +17,7 @@ normalize_aux_keys, ) from flopy4.mf6.spec import to_field_type +from flopy4.spec import field_meta, pydantic_fields # DFN type -> numpy dtype, for broadcasting a scalar griddata default to a # full array. @@ -27,9 +30,67 @@ } -@attrs.define(kw_only=True, slots=False) +def _is_dask_array(v: Any) -> bool: + try: + from dask.array import Array as _DaskArray + except ImportError: + return False + return isinstance(v, _DaskArray) + + +@dataclass(config=CFG, kw_only=True) class Package(Component, ABC): - def __attrs_post_init__(self) -> None: + # A griddata field's *declared* type is an array type (NDArray[...]/ + # FloatArrayLike/IntArrayLike), but its value is often a bare scalar + # (e.g. `strt: FloatArrayLike = field(default=1.0, ...)`) that only + # becomes a real array once dims are known, in __post_init__'s + # _broadcast_griddata. Pydantic doesn't validate defaults (no + # validate_default=True), but an explicit scalar override would fail + # the array type check without this coercion. One shared + # `field_validator("*", mode="before")`, + # driven by each field's own `json_schema_extra["shape"]` (which + # `spec.field()` already emits today), covers every griddata field on + # every subclass -- not one per field, not one per generated class. + @field_validator("*", mode="before") + @classmethod + def _coerce_arrays(cls, v: Any, info) -> Any: + finfo = pydantic_fields(cls).get(info.field_name) + if finfo is None or v is None: + return v + meta = field_meta(finfo) + if not (isinstance(meta, dict) and meta.get("block") == "griddata" and meta.get("shape")): + return v + if isinstance(v, np.ndarray) or _is_dask_array(v): + # Already a real ndarray, or specifically a dask.array.Array + # (the one duck array actually exercised here -- see + # codec/writer/filters.py's array2chunks). np.asarray() below + # would materialize a dask array into a real ndarray, losing + # its laziness (confirmed by running the real dask-array + # griddata test), so it's passed through untouched. Anything + # ELSE duck-array-shaped (an xr.DataArray, notably -- confirmed + # by running Disv.from_grid() with a DataArray-backed grid) is + # NOT preserved as-is: some hand-written fields (Dis/Disv's + # own top/botm/delr/delc/iv/xv/yv) declare the stricter + # NDArray[...] rather than the _ArrayLike Protocol most + # generated griddata fields use, and only real np.ndarray + # satisfies that -- so it still needs materializing below. + return v + dtype = _DTYPE_MAP.get(to_field_type(finfo.annotation), np.float64) + if isinstance(v, dict): + # An empty-dict griddata value (e.g. Chd(dims={}) with no + # explicit scalar override) is _broadcast_griddata's own + # "use the field's own scalar default" signal, but it can't + # pass the NDArray/_ArrayLike type check (and np.asarray({}) + # raises). Pre-resolve it into the same + # 0-d default-valued array a bare scalar default produces here + # -- _broadcast_griddata's existing size==1 branch (added for + # that scalar case) picks it up and broadcasts it exactly the + # same way. + default = finfo.default if isinstance(finfo.default, (int, float)) else 0 + return np.asarray(default, dtype=dtype) + return np.asarray(v, dtype=dtype) + + def __post_init__(self, dims: Optional[dict] = None) -> None: """Post-init for Package subclasses. Handles three concerns in order: @@ -37,75 +98,71 @@ def __attrs_post_init__(self) -> None: auto-set ns. 2. Broadcast scalar griddata values to their DFN shape when dims is supplied (e.g. IC(strt=1.0, dims={"nodes": 900})). - 3. Chain to Component.__attrs_post_init__() via super() -- LAST, - after 1-2, in every exit path (including the two early - returns below). Several of a package's own fields (e.g. - griddata arrays default to a bare scalar/dict until step 2 - broadcasts them) aren't in their final shape until steps 1-2 - finish, and Component.__attrs_post_init__() (via - DimensionResolverMixin's chain and _set_child_parents(), - which walks every attrs field) reads them -- so chaining - before they're finalized breaks griddata broadcasting and - dims resolution. Matches the ordering DisBase/Dis already use - for their own __attrs_post_init__ chaining (super() called - last, after their own field setup). + 3. Chain to Component.__post_init__() via super() -- LAST, after + 1-2, in every exit path (including the early return below). + Several of a package's own fields (e.g. griddata arrays + default to a bare scalar/dict until step 2 broadcasts them) + aren't in their final shape until steps 1-2 finish, and + Component.__post_init__() (via DimensionResolverMixin's chain + and _set_child_parents(), which walks every field) reads them + -- so chaining before they're finalized breaks griddata + broadcasting and dims resolution. Matches the ordering + DisBase/Dis already use for their own __post_init__ chaining + (super() called last, after their own field setup). """ - import attrs as _attrs - - # Detect schema-driven fields by presence of 'block' in field metadata. - # Package subclasses with no fields of their own (e.g. the - # Gwfgwe/Gwfgwt/Gwfprt exchange leaves in flopy4/mf6/exg/ -- just - # dfn_name, no declared fields) no-op through the rest of this method. - try: - fields = _attrs.fields(type(self)) # type: ignore[arg-type] - except _attrs.exceptions.NotAnAttrsClassError: - super().__attrs_post_init__() - return - if not any(f.metadata.get("block") is not None for f in fields): - super().__attrs_post_init__() + # Detect schema-driven fields by presence of 'block' in field + # json_schema_extra. Package subclasses with no fields of their + # own (e.g. the Gwfgwe/Gwfgwt/Gwfprt exchange leaves in + # flopy4/mf6/exg/ -- just dfn_name, no declared fields) still + # inherit Component's own fields (filename, name, ...), none of + # which carry block metadata, so this simply no-ops through the + # rest of this method for them. + fields = pydantic_fields(type(self)) + if not any(field_meta(f).get("block") is not None for f in fields.values()): + super().__post_init__(dims) return # 1. Item-list coercion. self._init_item_lists(fields) # 2. Griddata broadcasting. - dims: dict = self.__dict__.get("dims") or {} if dims: self._broadcast_griddata(fields, dims) # 3. Chain to Component's own post-init -- see docstring above for # why this must run last, not first. - super().__attrs_post_init__() + super().__post_init__(dims) def _init_item_lists(self, fields) -> None: """Coerce raw list/dict block+period data into Item-list fields; auto-set ns from the resulting list lengths. `maxbound` (where applicable) is a computed property instead, not set here. - Reads/writes the field's real attribute name (f.name) always -- - aliases (e.g. _stress_period_data's "stress_period_data") only name - the __init__ parameter; the instance attribute (and __dict__ key - object.__setattr__ writes to) is still the real name. + Reads/writes the field's real attribute name (the dict key) + always -- aliases (e.g. _stress_period_data's "stress_period_data") + only name the __init__ parameter; the instance attribute (and + __dict__ key object.__setattr__ writes to) is still the real name. """ - for f in fields: - block = f.metadata.get("block") + for name, f in fields.items(): + meta = field_meta(f) + block = meta.get("block") if not block: continue - item_cls = item_list_type(f.type) + item_cls = item_list_type(f.annotation) if item_cls is None: continue - raw = self.__dict__.get(f.name) + raw = self.__dict__.get(name) if raw is None: continue - if f.metadata.get("fill_forward"): + if meta.get("fill_forward"): coerced = { kper: self._coerce_item_list(rows, item_cls) for kper, rows in raw.items() } - object.__setattr__(self, f.name, coerced) + object.__setattr__(self, name, coerced) else: coerced_list = self._coerce_item_list(raw, item_cls) - object.__setattr__(self, f.name, coerced_list) + object.__setattr__(self, name, coerced_list) if getattr(self, f"n{block}s", 0) == 0: object.__setattr__(self, f"n{block}s", len(coerced_list)) @@ -173,16 +230,17 @@ def _broadcast_griddata(self, fields, dims: dict) -> None: and _par_data.dims.get("nrow", 0) == 0 ) or ("ncpl" in dims and "nrow" not in dims) - for f in fields: - if f.metadata.get("block") != "griddata": + for name, f in fields.items(): + meta = field_meta(f) + if meta.get("block") != "griddata": continue - shape_meta = f.metadata.get("shape") + shape_meta = meta.get("shape") if not shape_meta: continue - val = self.__dict__.get(f.name) + val = self.__dict__.get(name) if val is None: continue - _gd_dtype = _DTYPE_MAP.get(to_field_type(f.type), np.float64) + _gd_dtype = _DTYPE_MAP.get(to_field_type(f.annotation), np.float64) try: resolved = [] for d in shape_meta: @@ -194,15 +252,23 @@ def _broadcast_griddata(self, fields, dims: dict) -> None: except KeyError: continue if isinstance(val, (int, float)): - self.__dict__[f.name] = np.full(shape, val, dtype=_gd_dtype) + self.__dict__[name] = np.full(shape, val, dtype=_gd_dtype) + elif isinstance(val, np.ndarray) and val.size == 1 and val.shape != shape: + # A scalar that already passed through _coerce_arrays' + # mode="before" validator (needed so pydantic's own type + # check on an _ArrayLike-typed field accepts it at all -- + # see that validator's docstring) arrives here as a 0-d + # ndarray, not a bare int/float, so the branch above + # doesn't match. Same broadcast, just unwrapped first. + self.__dict__[name] = np.full(shape, val.item(), dtype=_gd_dtype) elif isinstance(val, np.ndarray) and val.shape != shape: try: - self.__dict__[f.name] = val.reshape(shape) + self.__dict__[name] = val.reshape(shape) except ValueError: pass elif isinstance(val, dict) and not val: default = f.default if isinstance(f.default, (int, float)) else 0 - self.__dict__[f.name] = np.full(shape, default, dtype=_gd_dtype) + self.__dict__[name] = np.full(shape, default, dtype=_gd_dtype) @classmethod def load( # type: ignore[override] @@ -255,27 +321,23 @@ def to_dict(self, blocks: bool = False, strict: bool = False) -> dict: strict : bool If True, only include fields with ``block`` metadata. """ - import attrs as _attrs - - try: - all_fields = _attrs.fields(type(self)) # type: ignore[arg-type] - except _attrs.exceptions.NotAnAttrsClassError: - return super().to_dict(blocks=blocks, strict=strict) + all_fields = pydantic_fields(type(self)) # Fall back for a Package subclass with no schema-driven fields of # its own (e.g. the exchange leaves in flopy4/mf6/exg/). - if not any(f.metadata.get("block") for f in all_fields): + if not any(field_meta(f).get("block") for f in all_fields.values()): return super().to_dict(blocks=blocks, strict=strict) - _exclude = {"name", "parent", "_parent", "dims", "filename", "workspace", "strict"} + _exclude = {"name", "parent", "_parent", "filename", "workspace", "strict"} result: dict = {} - for f in all_fields: - if f.name in _exclude or f.init is False: + for name, f in all_fields.items(): + if name in _exclude or f.init is False: continue - block = f.metadata.get("block") + meta = field_meta(f) + block = meta.get("block") if not block: continue - key = f.alias if (f.alias and f.name.startswith("_")) else f.name + key = f.alias if (f.alias and name.startswith("_")) else name val = getattr(self, key, None) if blocks: result.setdefault(block, {})[key] = val @@ -285,12 +347,14 @@ def to_dict(self, blocks: bool = False, strict: bool = False) -> dict: def to_dataframe(self) -> pd.DataFrame: """Return stress period data as a tidy DataFrame. Zero cost if not called.""" + import dataclasses as _dc + _spd = self.__dict__.get("_stress_period_data") if not _spd: return pd.DataFrame() frames = [] for kper in sorted(_spd): - df = pd.DataFrame([attrs.asdict(row) for row in _spd[kper]]) + df = pd.DataFrame([_dc.asdict(row) for row in _spd[kper]]) df.insert(0, "kper", kper) frames.append(df) return pd.concat(frames, ignore_index=True) if frames else pd.DataFrame() @@ -315,13 +379,21 @@ def from_dataframe(self, df: pd.DataFrame) -> None: spd: dict[int, list] = {} for kper, group in df.groupby("kper"): group = group.drop(columns=["kper"]) - spd[int(kper)] = [item_cls(**row) for row in group.to_dict("records")] + rows = [] + for row in group.to_dict("records"): + # A missing optional column (e.g. boundname) round-trips + # through pandas as NaN, not absent -- drop it so the + # Item class's own default applies instead of failing + # type validation (a float in a str-typed field). + rows.append(item_cls(**{k: v for k, v in row.items() if not pd.isna(v)})) + spd[int(kper)] = rows self.__dict__["_stress_period_data"] = spd def _period_item_cls(self) -> "type[Item] | tuple[type[Item], ...]": - for f in attrs.fields(type(self)): # type: ignore[arg-type] - if f.metadata.get("fill_forward"): - item_cls = item_list_type(f.type) + for f in pydantic_fields(type(self)).values(): + meta = field_meta(f) + if meta.get("fill_forward"): + item_cls = item_list_type(f.annotation) if item_cls is not None: return item_cls raise ValueError(f"{type(self).__name__} has no period Item-list field") @@ -362,17 +434,15 @@ def to_xarray(self) -> "xr.Dataset": # type: ignore[override] Stays lazy if dask-backed. For packages with no array fields this falls through to ``Component.to_xarray()``, which returns whatever - ``attrs_to_dataset()`` finds -- empty for a package with no + ``dataclass_to_dataset()`` finds -- empty for a package with no griddata fields of its own. """ - import attrs as _attrs - - fields = _attrs.fields(type(self)) # type: ignore[arg-type] + fields = pydantic_fields(type(self)) for _block in ("griddata", "period"): data_vars = { - a.name: self.to_dataarray(a.name) - for a in fields - if a.metadata.get("block") == _block and getattr(self, a.name) is not None + name: self.to_dataarray(name) + for name, f in fields.items() + if field_meta(f).get("block") == _block and getattr(self, name) is not None } if data_vars: return xr.Dataset(data_vars) diff --git a/flopy4/mf6/prt/__init__.py b/flopy4/mf6/prt/__init__.py index 589ae7ad..e3fc66be 100644 --- a/flopy4/mf6/prt/__init__.py +++ b/flopy4/mf6/prt/__init__.py @@ -1,11 +1,11 @@ from typing import ClassVar, Optional -import attrs from flopy.discretization.structuredgrid import StructuredGrid from flopy.discretization.vertexgrid import VertexGrid +from pydantic.dataclasses import dataclass from flopy4.mf6.gwf.disbase import DisBase -from flopy4.mf6.model import Model +from flopy4.mf6.model import CFG, Model from flopy4.mf6.prt.dis import Dis from flopy4.mf6.prt.disv import Disv from flopy4.mf6.prt.fmi import Fmi @@ -36,7 +36,7 @@ def convert_grid(value): ] -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Prt(Model): dfn_name: ClassVar[str] = "prt-nam" @@ -48,7 +48,7 @@ class Prt(Model): fmi: Fmi | None = field(block="packages", default=None) mip: Mip | None = field(block="packages", default=None) oc: Oc | None = field(block="packages", default=None) - prp: list[Prp] = field(block="packages", default=attrs.Factory(list)) + prp: list[Prp] = field(block="packages", default_factory=list) @property def grid(self): diff --git a/flopy4/mf6/prt/dis.py b/flopy4/mf6/prt/dis.py index 95174895..cf0e389f 100644 --- a/flopy4/mf6/prt/dis.py +++ b/flopy4/mf6/prt/dis.py @@ -1,15 +1,15 @@ from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic.dataclasses import dataclass -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.spec import field from flopy4.mf6.utils.grid import StructuredGrid -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Dis(DisBase): dfn_name: ClassVar[str] = "prt-dis" @@ -23,7 +23,7 @@ class Dis(DisBase): nlay: int = field(default=1, block="dimensions") ncol: int = field(default=2, block="dimensions") nrow: int = field(default=2, block="dimensions") - delr: NDArray[np.float64] = field( + delr: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a row", block="griddata", @@ -31,7 +31,7 @@ class Dis(DisBase): layered=False, netcdf=False, ) - delc: NDArray[np.float64] = field( + delc: Optional[NDArray[np.float64]] = field( default=1.0, longname="spacing along a column", block="griddata", @@ -39,7 +39,7 @@ class Dis(DisBase): layered=False, netcdf=False, ) - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=1.0, longname="cell top elevation", block="griddata", @@ -47,7 +47,7 @@ class Dis(DisBase): layered=False, netcdf=False, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=0.0, longname="cell bottom elevation", block="griddata", @@ -64,12 +64,12 @@ class Dis(DisBase): netcdf=False, ) - def __attrs_post_init__(self): + def __post_init__(self, dims: Optional[dict] = None): self.nodes = self.ncol * self.nrow * self.nlay self.ncpl = self.ncol * self.nrow self.nvert = (self.ncol + 1) * (self.nrow + 1) self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/prt/disv.py b/flopy4/mf6/prt/disv.py index 2cda2753..50563f58 100644 --- a/flopy4/mf6/prt/disv.py +++ b/flopy4/mf6/prt/disv.py @@ -1,28 +1,29 @@ from typing import ClassVar, Optional -import attrs import numpy as np from numpy.typing import NDArray +from pydantic import Field, field_validator +from pydantic.dataclasses import dataclass -from flopy4.mf6.gwf.disbase import DisBase +from flopy4.mf6.gwf.disbase import CFG, DisBase from flopy4.mf6.item import Item from flopy4.mf6.spec import field from flopy4.mf6.utils.grid import VertexGrid -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Disv(DisBase): dfn_name: ClassVar[str] = "prt-disv" - @attrs.define(slots=False) + @dataclass(config=CFG) class Cell2dRecord: - icell2d: int = attrs.field() - xc: float = attrs.field() - yc: float = attrs.field() - ncvert: int = attrs.field() - icvert: tuple[int, ...] = attrs.field() + icell2d: int = Field() + xc: float = Field() + yc: float = Field() + ncvert: int = Field() + icvert: tuple[int, ...] = Field() - @attrs.define + @dataclass(config=CFG) class Vertices(Item): iv: int xv: float @@ -38,7 +39,7 @@ class Vertices(Item): nlay: int = field(default=0, block="dimensions") ncpl: int = field(default=0, block="dimensions") nvert: int = field(default=0, block="dimensions") - top: NDArray[np.float64] = field( + top: Optional[NDArray[np.float64]] = field( default=None, longname="model top elevation", block="griddata", @@ -46,7 +47,7 @@ class Vertices(Item): layered=False, netcdf=False, ) - botm: NDArray[np.float64] = field( + botm: Optional[NDArray[np.float64]] = field( default=None, longname="model bottom elevation", block="griddata", @@ -62,20 +63,31 @@ class Vertices(Item): layered=True, netcdf=False, ) - iv: Optional[NDArray[np.int64]] = attrs.field(default=None) - xv: Optional[NDArray[np.float64]] = attrs.field(default=None) - yv: Optional[NDArray[np.float64]] = attrs.field(default=None) + iv: Optional[NDArray[np.int64]] = Field(default=None) + xv: Optional[NDArray[np.float64]] = Field(default=None) + yv: Optional[NDArray[np.float64]] = Field(default=None) + + # iv/xv/yv are declared NDArray-typed but commonly constructed from a + # plain list/tuple (see from_grid() below) -- unlike Package's own + # griddata fields, these carry no block="griddata"/shape= metadata, so + # Package._coerce_arrays' shape-driven check doesn't reach them, so + # they get their own mode="before" coercion (as do Tdis.perlen/nstp/ + # tsmult). + @field_validator("iv", mode="before") + @classmethod + def _coerce_iv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.int64) + + @field_validator("xv", "yv", mode="before") + @classmethod + def _coerce_xv_yv(cls, v): + return v if v is None or isinstance(v, np.ndarray) else np.asarray(v, dtype=np.float64) + vertices: Optional[list[Vertices]] = field(default=None, block="vertices") - cell2ddata: Optional[list] = attrs.field(default=None) + cell2ddata: Optional[list] = Field(default=None) cell2d: Optional[list] = field(default=None, init=False, block="cell2d") - def __attrs_post_init__(self): - if self.iv is not None and (not isinstance(self.iv, np.ndarray)): - object.__setattr__(self, "iv", np.asarray(self.iv, dtype=np.int64)) - if self.xv is not None and (not isinstance(self.xv, np.ndarray)): - object.__setattr__(self, "xv", np.asarray(self.xv, dtype=np.float64)) - if self.yv is not None and (not isinstance(self.yv, np.ndarray)): - object.__setattr__(self, "yv", np.asarray(self.yv, dtype=np.float64)) + def __post_init__(self, dims: Optional[dict] = None): if self.iv is not None and self.xv is not None and (self.yv is not None): rows = [ self.Vertices(iv=int(iv) + 1, xv=float(xv), yv=float(yv)) @@ -83,18 +95,18 @@ def __attrs_post_init__(self): ] object.__setattr__(self, "vertices", rows) if self.cell2ddata is not None: - rows = [] + cell_rows = [] for rec in self.cell2ddata: row = (rec.icell2d + 1, rec.xc, rec.yc, rec.ncvert) + tuple( (v + 1 for v in rec.icvert) ) - rows.append(row) - object.__setattr__(self, "cell2d", rows) + cell_rows.append(row) + object.__setattr__(self, "cell2d", cell_rows) self.nodes = self.ncpl * self.nlay self.nrow = 0 self.ncol = 0 self._coerce_griddata() - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/prt/fmi.py b/flopy4/mf6/prt/fmi.py index 449be91d..2758c363 100644 --- a/flopy4/mf6/prt/fmi.py +++ b/flopy4/mf6/prt/fmi.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Fmi(Package): dfn_name: ClassVar[str] = "prt-fmi" diff --git a/flopy4/mf6/prt/mip.py b/flopy4/mf6/prt/mip.py index 83566432..00289481 100644 --- a/flopy4/mf6/prt/mip.py +++ b/flopy4/mf6/prt/mip.py @@ -1,14 +1,14 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, IntArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Mip(Package): dfn_name: ClassVar[str] = "prt-mip" diff --git a/flopy4/mf6/prt/oc.py b/flopy4/mf6/prt/oc.py index 32c4ceeb..bbb57afe 100644 --- a/flopy4/mf6/prt/oc.py +++ b/flopy4/mf6/prt/oc.py @@ -2,51 +2,52 @@ from pathlib import Path from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Oc(Package): dfn_name: ClassVar[str] = "prt-oc" - @attrs.define + @dataclass(config=CFG) class Tracktimes(Item): time: float = field() - @attrs.define + @dataclass(config=CFG) class All(Item): _keyword: ClassVar[str] = "all" - @attrs.define + @dataclass(config=CFG) class First(Item): _keyword: ClassVar[str] = "first" - @attrs.define + @dataclass(config=CFG) class Last(Item): _keyword: ClassVar[str] = "last" - @attrs.define + @dataclass(config=CFG) class Frequency(Item): _keyword: ClassVar[str] = "frequency" frequency: int = field() - @attrs.define + @dataclass(config=CFG) class Steps(Item): _keyword: ClassVar[str] = "steps" steps: tuple = field(default=(), array=True) - @attrs.define + @dataclass(config=CFG) class Save(Item): _keyword: ClassVar[str] = "save" rtype: Union[float, str] = field() ocsetting: "Oc.All | Oc.First | Oc.Last | Oc.Frequency | Oc.Steps" = field() - @attrs.define + @dataclass(config=CFG) class Print(Item): _keyword: ClassVar[str] = "print" rtype: Union[float, str] = field() @@ -146,12 +147,12 @@ class Print(Item): optional=True, longname="number of particle tracking times", ) - tracktimes: Optional[list[Tracktimes]] = field( + tracktimes: Optional[SkipValidation[list[Tracktimes]]] = field( default=None, block="tracktimes", auto_from="tracktimes", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/prt/prp.py b/flopy4/mf6/prt/prp.py index 4400cb22..b9f01b02 100644 --- a/flopy4/mf6/prt/prp.py +++ b/flopy4/mf6/prt/prp.py @@ -2,21 +2,22 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Prp(Package): dfn_name: ClassVar[str] = "prt-prp" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Packagedata(Item): irptno: int = field(index=True, pk=True) cellid: tuple = field(cellid=True) @@ -26,33 +27,33 @@ class Packagedata(Item): aux: tuple = () boundname: Optional[str] = field(default=None, optional=True) - @attrs.define + @dataclass(config=CFG) class Releasetimes(Item): time: float = field() - @attrs.define + @dataclass(config=CFG) class All(Item): _keyword: ClassVar[str] = "all" - @attrs.define + @dataclass(config=CFG) class First(Item): _keyword: ClassVar[str] = "first" - @attrs.define + @dataclass(config=CFG) class Last(Item): _keyword: ClassVar[str] = "last" - @attrs.define + @dataclass(config=CFG) class Frequency(Item): _keyword: ClassVar[str] = "frequency" frequency: int = field() - @attrs.define + @dataclass(config=CFG) class Steps(Item): _keyword: ClassVar[str] = "steps" steps: tuple = field(default=(), array=True) - @attrs.define + @dataclass(config=CFG) class Fraction(Item): _keyword: ClassVar[str] = "fraction" fraction: tuple = field(default=(), array=True) @@ -168,17 +169,17 @@ class Fraction(Item): block="dimensions", longname="number of particle release times", ) - packagedata: Optional[list[Packagedata]] = field( + packagedata: Optional[SkipValidation[list[Packagedata]]] = field( default=None, block="packagedata", auto_from="packagedata", ) - releasetimes: Optional[list[Releasetimes]] = field( + releasetimes: Optional[SkipValidation[list[Releasetimes]]] = field( default=None, block="releasetimes", auto_from="releasetimes", ) - _stress_period_data: Optional[dict[int, list[_StressPeriodDataItem]]] = field( + _stress_period_data: Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]] = field( alias="stress_period_data", default=None, repr=False, diff --git a/flopy4/mf6/pts.py b/flopy4/mf6/pts.py index dc3a0b29..d4808b87 100644 --- a/flopy4/mf6/pts.py +++ b/flopy4/mf6/pts.py @@ -2,24 +2,25 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import _optional_path from flopy4.mf6.record import Record -from flopy4.mf6.solution import Solution +from flopy4.mf6.solution import CFG, Solution from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Pts(Solution): dfn_name: ClassVar[str] = "sln-pts" slntype: ClassVar[str] = "pts" - @attrs.define + @dataclass(config=CFG) class NoPtc(Record): _keyword: ClassVar[str] = "no_ptc" - no_ptc_option: Optional[str] = attrs.field(default=None) + no_ptc_option: Optional[str] = Field(default=None) print_option: Optional[str] = field( default=None, diff --git a/flopy4/mf6/record.py b/flopy4/mf6/record.py index bc41bc86..bcb8a2a8 100644 --- a/flopy4/mf6/record.py +++ b/flopy4/mf6/record.py @@ -9,90 +9,92 @@ another) rather than flattening the nested one's fields into itself -- see _nested_class, inferred from the field's own type annotation rather than a declared flag, and make.py's _build_record_class_specs. + +Generated Record/Item subclasses are `pydantic.dataclasses.dataclass` -- +`Record.fields()`/`_nested_class()` below read +`pydantic_fields()`/`FieldInfo.json_schema_extra` accordingly. A nested/ +composed field's annotation (e.g. `Headprint.formatrecord: "Oc.Format"`) is +a forward-reference string naming a SIBLING class inside the same enclosing +package class -- unresolvable via any module-global lookup at class-body- +execution time (Python class bodies can't see sibling names in an enclosing +class's scope). Pydantic resolves it lazily on first construction -- `Record.fields()`'s guarded +`rebuild_dataclass()` call handles the one case that doesn't self-heal on +its own: something (like `from_tokens()`) inspecting a class's fields +before any instance of it has ever been built. """ -import sys +from __future__ import annotations + import types -from functools import lru_cache -from pathlib import Path -from typing import Any, Union, cast, get_args, get_origin +from pathlib import Path, PurePath +from typing import Any, Union, get_args, get_origin -import attrs +from pydantic import ConfigDict +from pydantic.dataclasses import is_pydantic_dataclass, rebuild_dataclass +from flopy4.spec import field_meta, pydantic_fields -def _resolve_sibling_class(cls: type, name: str) -> Any | None: - """Resolve a bare sibling class name one level up from `cls` in - `__qualname__`, where every generated flat-sibling class (composed - Records, keystring-union arms) lives regardless of DFN nesting depth. - No type check here -- callers (`_nested_class`, item.py's - `_nested_union_classes`) apply their own, against different base - classes. +CFG = ConfigDict(arbitrary_types_allowed=True, validate_assignment=True, extra="forbid") - Must resolve at runtime: a class body can't see sibling names from an - enclosing scope, which is why the qualified string form - (``"Oc.Format"``) exists at all -- purely for mypy. - """ - obj = sys.modules[cls.__module__] - for part in cls.__qualname__.split(".")[:-1]: - obj = getattr(obj, part) - return getattr(obj, name, None) +def _nested_class(cls: type, annotation: Any) -> "type[Record] | None": + """If a field's (already-resolved) annotation is -- or wraps, via + `Optional[...]` -- a Record subclass, return it; else None. + Resolvability against a real Record subclass is itself the signal, no + declared "is this nested" flag needed. -@lru_cache(maxsize=None) -def _nested_class(cls: type, type_str: str) -> "type[Record] | None": - """If a field's raw type annotation (e.g. ``"Format"`` or - ``"Optional[Oc.Format]"``) names a Record subclass, return it; else - None. Resolvability against a real Record subclass is itself the - signal -- no declared "is this nested" flag needed. - - Cached since to_tokens/from_tokens call this per field, often - repeatedly while parsing many rows. + By the time `Record.fields()` has run, `annotation` (a pydantic + `FieldInfo.annotation`) is already the real class object, not a + string -- no `sys.modules`/qualname lookup needed. """ - name = type_str - if name.startswith("Optional[") and name.endswith("]"): - name = name[len("Optional[") : -1] - name = name.rsplit(".", 1)[-1] - resolved = _resolve_sibling_class(cls, name) - return resolved if isinstance(resolved, type) and issubclass(resolved, Record) else None + args = get_args(annotation) + candidate = next((a for a in args if a is not type(None)), annotation) + return candidate if isinstance(candidate, type) and issubclass(candidate, Record) else None -def _is_bool_field(f: attrs.Attribute) -> bool: - t = f.type +def _is_bool_field(finfo: Any) -> bool: + t = finfo.annotation origin = get_origin(t) if origin is types.UnionType or origin is Union: t = next((a for a in get_args(t) if a is not type(None)), t) - return t in (bool, "bool") + return t is bool -def _coerce(token: Any, f: attrs.Attribute) -> Any: +def _coerce(token: Any, finfo: Any) -> Any: """Cast a raw token to a field's declared type (time_series falls back to the raw string if it isn't a float). Only Optional[X] (a single non-None union arm) is unwrapped -- a genuine multi-type union like Union[float, str] is deliberately ambiguous and left as the raw token.""" - if f.metadata.get("time_series"): + meta = field_meta(finfo) + if isinstance(meta, dict) and meta.get("time_series"): try: return float(token) except (ValueError, TypeError): return str(token) - t = f.type + t = finfo.annotation origin = get_origin(t) if origin is types.UnionType or origin is Union: args = [a for a in get_args(t) if a is not type(None)] if len(args) != 1: return token t = args[0] - if t in (int, "int"): + if t is int: return int(float(str(token))) - if t in (float, "float"): + if t is float: return float(token) - if t in (Path, "Path"): + if t is Path: return Path(token) return token -def _is_list_field(f: attrs.Attribute) -> bool: +def _is_tagged(finfo: Any) -> bool: + meta = field_meta(finfo) + return bool(isinstance(meta, dict) and meta.get("tagged")) + + +def _is_list_field(finfo: Any) -> bool: """True for ``list[X]`` or ``Optional[list[X]]``""" - t = f.type + t = finfo.annotation origin = get_origin(t) if origin is types.UnionType or origin is Union: t = next((a for a in get_args(t) if a is not type(None)), t) @@ -100,9 +102,9 @@ def _is_list_field(f: attrs.Attribute) -> bool: return origin is list -def _list_elem_coerce(token: Any, f: attrs.Attribute) -> Any: +def _list_elem_coerce(token: Any, finfo: Any) -> Any: """Coerce one token to a list field's declared element type.""" - t = f.type + t = finfo.annotation origin = get_origin(t) if origin is types.UnionType or origin is Union: t = next((a for a in get_args(t) if a is not type(None)), t) @@ -122,29 +124,39 @@ def _tokens(name: str, value: Any) -> list: return [name.upper(), value] -def _consume_tagged(tokens: list, i: int, f: attrs.Attribute) -> "tuple[Any, int] | None": - """Match field f's tagged keyword at tokens[i]; a bool field needs no - value token, anything else does. None if unmatched or value missing.""" - if str(tokens[i]).upper() != f.name.upper(): +def _consume_tagged(tokens: list, i: int, name: str, finfo: Any) -> "tuple[Any, int] | None": + """Match field `name`'s tagged keyword at tokens[i]; a bool field needs + no value token, anything else does. None if unmatched or value + missing.""" + if str(tokens[i]).upper() != name.upper(): return None - if _is_bool_field(f): + if _is_bool_field(finfo): return True, 1 if i + 1 >= len(tokens): return None - return _coerce(tokens[i + 1], f), 2 + return _coerce(tokens[i + 1], finfo), 2 class Record: """Mixin for record types.""" @classmethod - def fields(cls: type["Record"]) -> list[attrs.Attribute]: - """Record (or Item) class' fields, in declaration order.""" - fields = attrs.fields(cast(type[attrs.AttrsInstance], cls)) - return [f for f in fields if not f.name.startswith("_")] + def fields(cls) -> dict[str, Any]: + """Record (or Item) class' non-private fields, in declaration + order, keyed by name -- each value a pydantic `FieldInfo`. + + Guards with a `rebuild_dataclass()` call so a nested/composed + field's annotation is the real class, not a stale `ForwardRef`, + even when called before any instance of `cls` has ever been + constructed (exactly what `from_tokens()` does). + """ + assert is_pydantic_dataclass(cls) + if not cls.__pydantic_complete__: + rebuild_dataclass(cls, force=True, _parent_namespace_depth=4) + return {n: f for n, f in pydantic_fields(cls).items() if not n.startswith("_")} @classmethod - def keyword(cls: type["Record"]) -> str: + def keyword(cls) -> str: return vars(cls).get("_keyword", "") def to_tokens(self) -> tuple: @@ -154,27 +166,29 @@ def to_tokens(self) -> tuple: for tok in vars(cls).get("_extra_tokens", ()): tokens.append(tok) fields = cls.fields() - tagged = [a for a in fields if a.metadata.get("tagged")] - untagged = [a for a in fields if not a.metadata.get("tagged")] - for a in tagged + untagged: - v = getattr(self, a.name) + tagged = [(n, f) for n, f in fields.items() if _is_tagged(f)] + untagged = [(n, f) for n, f in fields.items() if not _is_tagged(f)] + for name, finfo in tagged + untagged: + v = getattr(self, name) if v is None: continue if isinstance(v, Record): tokens.extend(v.to_tokens()) - elif a.metadata.get("tagged"): - tokens.extend(_tokens(a.name, v)) + elif _is_tagged(finfo): + tokens.extend(_tokens(name, v)) elif isinstance(v, bool): if v: - tokens.append(a.name.upper()) + tokens.append(name.upper()) elif isinstance(v, (list, tuple)): tokens.extend(v) + elif isinstance(v, PurePath): + tokens.append(v.as_posix()) else: tokens.append(v) return tuple(tokens) @classmethod - def from_tokens(cls, tokens: str | list[str]) -> "Record": + def from_tokens(cls, tokens: "str | list[str]") -> "Record": """Parse a token string/list back into an instance. Tagged fields are matched by keyword wherever it appears; whatever's @@ -196,10 +210,9 @@ def from_tokens(cls, tokens: str | list[str]) -> "Record": fields = cls.fields() - def _nested(f: attrs.Attribute) -> "type[Record] | None": - return _nested_class(cast(type, cls), f.type) if isinstance(f.type, str) else None - - nested_fields = [f for f in fields if _nested(f) is not None] + nested_fields = [ + (n, f) for n, f in fields.items() if _nested_class(cls, f.annotation) is not None + ] if nested_fields: # A record composed of nested record(s) has, in the current # corpus, no other fields of its own once _keyword/_extra_tokens @@ -208,53 +221,53 @@ def _nested(f: attrs.Attribute) -> "type[Record] | None": f"{cls.__name__}: exactly one nested record field, with no plain " "fields of its own, is the only shape supported so far" ) - nf = nested_fields[0] - nested_cls = _nested(nf) + nf_name, nf_finfo = nested_fields[0] + nested_cls = _nested_class(cls, nf_finfo.annotation) assert nested_cls is not None - return cls(**{nf.name: nested_cls.from_tokens(tokens)}) + return cls(**{nf_name: nested_cls.from_tokens(tokens)}) - tagged = {f.name.upper(): f for f in fields if f.metadata.get("tagged")} - untagged = [f for f in fields if not f.metadata.get("tagged")] + tagged = {n.upper(): (n, f) for n, f in fields.items() if _is_tagged(f)} + untagged = [(n, f) for n, f in fields.items() if not _is_tagged(f)] kwargs: dict = {} consumed: set[int] = set() i = 0 while i < len(tokens): - f = tagged.get(str(tokens[i]).upper()) - result = None if f is None else _consume_tagged(tokens, i, f) - if f is None or result is None: + entry = tagged.get(str(tokens[i]).upper()) + result = None if entry is None else _consume_tagged(tokens, i, entry[0], entry[1]) + if entry is None or result is None: i += 1 continue val, width = result - kwargs[f.name] = val + kwargs[entry[0]] = val for j in range(width): consumed.add(i + j) i += width required_tagged = [ - f - for f in fields - if f.metadata.get("tagged") and f.default is attrs.NOTHING and f.name not in kwargs + (n, f) + for n, f in fields.items() + if _is_tagged(f) and f.is_required() and n not in kwargs ] positional_queue = required_tagged + untagged remaining = [t for j, t in enumerate(tokens) if j not in consumed] - list_field = next((f for f in positional_queue if _is_list_field(f)), None) + list_field = next((nf for nf in positional_queue if _is_list_field(nf[1])), None) if list_field is not None: idx = positional_queue.index(list_field) assert idx == len(positional_queue) - 1, ( - f"{cls.__name__}.{list_field.name}: a list-typed record field must be " + f"{cls.__name__}.{list_field[0]}: a list-typed record field must be " "the last positional field -- it consumes all remaining tokens" ) scalar_fields = positional_queue[:idx] - for f, tok in zip(scalar_fields, remaining): - kwargs[f.name] = _coerce(tok, f) - kwargs[list_field.name] = [ - _list_elem_coerce(tok, list_field) for tok in remaining[len(scalar_fields) :] + for (name, finfo), tok in zip(scalar_fields, remaining): + kwargs[name] = _coerce(tok, finfo) + kwargs[list_field[0]] = [ + _list_elem_coerce(tok, list_field[1]) for tok in remaining[len(scalar_fields) :] ] else: - for f, tok in zip(positional_queue, remaining): - kwargs[f.name] = _coerce(tok, f) + for (name, finfo), tok in zip(positional_queue, remaining): + kwargs[name] = _coerce(tok, finfo) return cls(**kwargs) diff --git a/flopy4/mf6/simulation.py b/flopy4/mf6/simulation.py index 3e5da181..8cb49596 100644 --- a/flopy4/mf6/simulation.py +++ b/flopy4/mf6/simulation.py @@ -1,11 +1,12 @@ from os import PathLike -from typing import ClassVar +from pathlib import Path +from typing import ClassVar, Optional from warnings import warn -import attrs from modflow_devtools.misc import cd, run_cmd +from pydantic.dataclasses import dataclass -from flopy4.mf6.context import Context, update_child_attr +from flopy4.mf6.context import CFG, Context from flopy4.mf6.exchange import Exchange from flopy4.mf6.model import Model from flopy4.mf6.solution import Solution @@ -22,32 +23,34 @@ def convert_time(value): raise TypeError(f"Expected Time or Tdis, got {type(value)}") -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Simulation(Context): dfn_name: ClassVar[str] = "sim-nam" - tdis: Tdis = field(block="timing", converter=convert_time, default=attrs.Factory(Tdis)) - models: dict[str, Model] = field(block="models", default=attrs.Factory(dict)) - exchanges: dict[str, Exchange] = field(block="exchanges", default=attrs.Factory(dict)) - solutions: dict[str, Solution] = field(block="solutiongroup", default=attrs.Factory(dict)) + tdis: Tdis = field(block="timing", converter=convert_time, default_factory=Tdis) + models: dict[str, Model] = field(block="models", default_factory=dict) + exchanges: dict[str, Exchange] = field(block="exchanges", default_factory=dict) + solutions: dict[str, Solution] = field(block="solutiongroup", default_factory=dict) def default_filename(self) -> str: return "mfsim.nam" - def __attrs_post_init__(self): - from attrs import fields_dict - - super().__attrs_post_init__() - if self.filename != "mfsim.nam": + def __post_init__(self, dims: Optional[dict] = None): + super().__post_init__(dims) + if self.filename != Path("mfsim.nam"): if self.filename is not None: warn( "Simulation filename must be 'mfsim.nam'.", UserWarning, ) - self.filename = "mfsim.nam" - fields = fields_dict(type(self)) - field = fields["workspace"] - update_child_attr(self, field, self.workspace) + self.filename = Path("mfsim.nam") + # Re-propagate workspace to Simulation's own children (models/ + # exchanges/solutions/tdis) -- Context.__post_init__ (already run, + # via super() above) only saw whatever was attached at ITS point + # in the post-init chain; re-assigning through the property setter + # (see Context.workspace) re-runs the propagation now that every + # field on this concrete Simulation instance is attached. + self.workspace = self.workspace @property def time(self) -> Time: diff --git a/flopy4/mf6/solution.py b/flopy4/mf6/solution.py index 34bee4fe..b68b868c 100644 --- a/flopy4/mf6/solution.py +++ b/flopy4/mf6/solution.py @@ -1,15 +1,16 @@ from abc import ABC from typing import ClassVar -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Solution(Package, ABC): slntype: ClassVar[str] = "sln" - models: list[str] = attrs.field(default=attrs.Factory(list)) + models: list[str] = Field(default_factory=list) def default_filename(self) -> str: return f"solution.{self.slntype.lower()}" diff --git a/flopy4/mf6/spec.py b/flopy4/mf6/spec.py index 60910ec0..c4c75656 100644 --- a/flopy4/mf6/spec.py +++ b/flopy4/mf6/spec.py @@ -1,5 +1,5 @@ """ -Wrap `attrs` specification utilities for MF6. +Wrap `pydantic` specification utilities for MF6. These include field decorators and introspection functions. """ @@ -7,27 +7,32 @@ import types from datetime import datetime from pathlib import Path -from typing import Literal, Union, get_args, get_origin +from typing import Any, Literal, Union, get_args, get_origin -import attrs import numpy as np -from attrs import NOTHING, Attribute +from pydantic import Field +from pydantic.fields import FieldInfo from flopy4.mf6._types import FloatArrayLike, IntArrayLike +from flopy4.spec import field_meta from flopy4.spec import fields_dict as flopy_fields_dict FieldType = Literal["keyword", "integer", "double", "string", "list", "record"] +# Sentinel distinguishing "no default given" (a required field) from a +# real default of `None` -- pydantic.Field()'s own "no default" state is +# just not passing `default=` at all, so this marks that case through +# field()/path()'s own default=... parameter. +_UNSET = object() + def field( - default=NOTHING, - validator=None, + default=_UNSET, + default_factory=None, converter=None, repr=True, - eq=True, init=True, metadata=None, - on_setattr=None, alias: str | None = None, block: str | None = None, longname: str | None = None, @@ -47,8 +52,18 @@ def field( tagged: bool = False, array: bool = False, ): - """Define a field: always a plain ``attrs.field()``.""" + """Define a field: always a plain ``pydantic.Field()``. + + ``converter``: stashed into the field's own `json_schema_extra` + (`metadata["converter"]`) rather than passed to `Field()` directly -- + pydantic has no per-field `converter=` hook. + Applied uniformly by one shared `field_validator("*", mode="before")` + on `Component` (see `flopy4.mf6.component.Component._apply_converter`) + instead of one per field. + """ metadata = metadata or {} + if converter is not None: + metadata["converter"] = converter if block: metadata["block"] = block if longname: @@ -85,38 +100,32 @@ def field( metadata["tagged"] = True if array: metadata["array"] = True - return attrs.field( - default=default, - validator=validator, - converter=converter, - repr=repr, - eq=eq, - init=init, - on_setattr=on_setattr, - metadata=metadata, - alias=alias, - ) + kwargs: dict[str, Any] = dict(repr=repr, init=init, json_schema_extra=metadata) + if alias is not None: + kwargs["alias"] = alias + if default_factory is not None: + kwargs["default_factory"] = default_factory + elif default is not _UNSET: + kwargs["default"] = default + return Field(**kwargs) FileDirection = Literal[None, "in", "out"] def path( - default=NOTHING, - validator=None, + default=_UNSET, converter=None, repr=True, - eq=True, init=True, metadata=None, - on_setattr=None, block: str | None = None, direction: FileDirection | None = None, longname: str | None = None, optional: bool = False, keyword: str | None = None, ): - """Define a path field: always a plain ``attrs.field()``. + """Define a path field: always a plain ``pydantic.Field()``. ``keyword``: the file record's trigger keyword, stored as ``_keyword`` metadata (lowercase, same convention as a Record class's own @@ -124,8 +133,14 @@ def path( in ``TS6 FILEIN `` (block level, where it's also the row's key on ingress -- not derivable from the py name, ts_filerecord → ts_file) or ``tab6`` in a LAK tables row (read by Item.to_tokens()/from_tokens()). + + ``converter``: see `field()`'s own docstring -- stashed in metadata, + applied by the shared `Component`-level validator (Item/Record-level + path fields go through `Record`'s own equivalent validator instead). """ metadata = metadata or {} + if converter is not None: + metadata["converter"] = converter if keyword: metadata["_keyword"] = keyword.lower() if block: @@ -136,22 +151,16 @@ def path( metadata["longname"] = longname if optional: metadata["optional"] = True - return attrs.field( - default=default, - validator=validator, - converter=converter, - repr=repr, - eq=eq, - init=init, - on_setattr=on_setattr, - metadata=metadata, - ) + kwargs: dict[str, Any] = dict(repr=repr, init=init, json_schema_extra=metadata) + if default is not _UNSET: + kwargs["default"] = default + return Field(**kwargs) -Block = dict[str, Attribute] +Block = dict[str, FieldInfo] -def blocks(cls) -> list[list[Attribute]]: +def blocks(cls) -> list[list[FieldInfo]]: """Return an ordered list of blocks for a component class.""" return [list(v.values()) for v in blocks_dict(cls).values()] @@ -160,30 +169,30 @@ def blocks_dict(cls) -> dict[str, Block]: """ Return an ordered dictionary of blocks for a component class, whose keys are block names. Each block is a map from variable - (field) name to `attrs.Attribute`. + (field) name to `FieldInfo`. """ fields = fields_dict(cls) blocks: dict[str, Block] = {} for k, v in fields.items(): - block = v.metadata["block"] + block = field_meta(v)["block"] if block not in blocks: blocks[block] = {} blocks[block][k] = v return dict(sorted(blocks.items(), key=block_sort_key)) -def fields(cls) -> list[Attribute]: +def fields(cls) -> list[FieldInfo]: """Return an ordered list of fields for a component class.""" return list(fields_dict(cls).values()) -def fields_dict(cls) -> dict[str, Attribute]: +def fields_dict(cls) -> dict[str, FieldInfo]: """ Return an ordered dictionary of fields for a component class, - whose keys are field names. Each field is an `attrs.Attribute`. + whose keys are field names. Each field is a `FieldInfo`. """ fields = flopy_fields_dict(cls) - return {k: v for k, v in fields.items() if "block" in v.metadata} + return {k: v for k, v in fields.items() if "block" in field_meta(v)} def _ndarray_field_type(t) -> FieldType | None: @@ -234,7 +243,7 @@ def repeating_array_key_type(field_type) -> type | None: return key if val in (IntArrayLike, FloatArrayLike) else None -def to_field_type(t: type) -> FieldType: +def to_field_type(t: Any) -> FieldType: if (result := _ndarray_field_type(t)) is not None: return result match t: diff --git a/flopy4/mf6/tdis.py b/flopy4/mf6/tdis.py index 2107c79f..4d3b2c53 100644 --- a/flopy4/mf6/tdis.py +++ b/flopy4/mf6/tdis.py @@ -1,21 +1,22 @@ from datetime import datetime -from typing import ClassVar, Optional +from typing import Any, ClassVar, Optional -import attrs import numpy as np from numpy.typing import ArrayLike, NDArray +from pydantic import field_validator +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field from flopy4.mf6.utils.time import Time -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Tdis(Package): dfn_name: ClassVar[str] = "sim-tdis" - @attrs.define + @dataclass(config=CFG) class PeriodData(Item): perlen: float nstp: int @@ -24,17 +25,46 @@ class PeriodData(Item): time_units: Optional[str] = field(default=None, block="options", optional=True) start_date_time: Optional[str] = field( default=None, - converter=lambda v: v.isoformat() if isinstance(v, datetime) else v, + # A raw ISO datetime like "1970-01-01T00:00:00" tokenizes into TWO + # file tokens (e.g. [1970, "-01-01T00:00:00"]) -- the reader's + # numeric-token grammar splits at the leading digits, before the + # first embedded dash -- and a bare year like "1997" (no rest of an + # ISO string on the row at all) tokenizes as a plain int, not a + # str. Neither passes Optional[str] validation, so the converter + # (already needed for the datetime -> isoformat direction) also + # normalizes both back into one string here. + converter=lambda v: ( + v.isoformat() + if isinstance(v, datetime) + else "".join(str(t) for t in v) + if isinstance(v, list) + else str(v) + if isinstance(v, (int, float)) + else v + ), block="options", optional=True, ) nper: int = field(default=1, block="dimensions") - perlen: NDArray[np.float64] = attrs.field(default=1.0) - nstp: NDArray[np.int64] = attrs.field(default=1) - tsmult: NDArray[np.float64] = attrs.field(default=1.0) + perlen: NDArray[np.float64] = field(default_factory=lambda: np.array(1.0)) + nstp: NDArray[np.int64] = field(default_factory=lambda: np.array(1, dtype=np.int64)) + tsmult: NDArray[np.float64] = field(default_factory=lambda: np.array(1.0)) perioddata: Optional[list[PeriodData]] = field(default=None, block="perioddata") - def __attrs_post_init__(self): + # perlen/nstp/tsmult aren't Package "griddata" block fields (no + # block="griddata" metadata), so Package._coerce_arrays' shape-driven + # check doesn't touch them -- same underlying problem though (a + # declared array type with a bare-scalar default/override), so this + # needs its own mode="before" coercion, scoped to just these 3 fields + # by name rather than metadata. + @field_validator("perlen", "nstp", "tsmult", mode="before") + @classmethod + def _coerce_to_array(cls, v: Any) -> Any: + if isinstance(v, np.ndarray): + return v + return np.asarray(v) + + def __post_init__(self, dims: Optional[dict] = None): if self.perioddata: rows = [ row @@ -48,27 +78,30 @@ def __attrs_post_init__(self): object.__setattr__(self, "perlen", np.array([r.perlen for r in rows], dtype=np.float64)) object.__setattr__(self, "nstp", np.array([r.nstp for r in rows], dtype=np.int64)) object.__setattr__(self, "tsmult", np.array([r.tsmult for r in rows], dtype=np.float64)) - super().__attrs_post_init__() + super().__post_init__(dims) return nper = self.nper - if isinstance(self.perlen, (int, float)): - object.__setattr__(self, "perlen", np.full(nper, self.perlen, dtype=np.float64)) - elif not isinstance(self.perlen, np.ndarray): - object.__setattr__(self, "perlen", np.asarray(self.perlen, dtype=np.float64)) - if isinstance(self.nstp, (int, float)): - object.__setattr__(self, "nstp", np.full(nper, int(self.nstp), dtype=np.int64)) - elif not isinstance(self.nstp, np.ndarray): - object.__setattr__(self, "nstp", np.asarray(self.nstp, dtype=np.int64)) - if isinstance(self.tsmult, (int, float)): - object.__setattr__(self, "tsmult", np.full(nper, self.tsmult, dtype=np.float64)) - elif not isinstance(self.tsmult, np.ndarray): - object.__setattr__(self, "tsmult", np.asarray(self.tsmult, dtype=np.float64)) + # Defaults are 0-d arrays, and _coerce_to_array turns an explicit + # scalar/list into an array, so each is an array here; broadcast a + # single value to nper. + if self.perlen.size == 1: + object.__setattr__(self, "perlen", np.full(nper, self.perlen.item(), dtype=np.float64)) + elif self.perlen.dtype != np.float64: + object.__setattr__(self, "perlen", self.perlen.astype(np.float64)) + if self.nstp.size == 1: + object.__setattr__(self, "nstp", np.full(nper, int(self.nstp.item()), dtype=np.int64)) + elif self.nstp.dtype != np.int64: + object.__setattr__(self, "nstp", self.nstp.astype(np.int64)) + if self.tsmult.size == 1: + object.__setattr__(self, "tsmult", np.full(nper, self.tsmult.item(), dtype=np.float64)) + elif self.tsmult.dtype != np.float64: + object.__setattr__(self, "tsmult", self.tsmult.astype(np.float64)) rows = [ Tdis.PeriodData(perlen=float(p), nstp=int(n), tsmult=float(t)) for p, n, t in zip(self.perlen, self.nstp, self.tsmult) ] object.__setattr__(self, "perioddata", rows) - super().__attrs_post_init__() + super().__post_init__(dims) def get_dims(self) -> dict[str, int]: """Get all dimensions.""" diff --git a/flopy4/mf6/utils/cbc_reader.py b/flopy4/mf6/utils/cbc_reader.py index 18a054cf..68666865 100644 --- a/flopy4/mf6/utils/cbc_reader.py +++ b/flopy4/mf6/utils/cbc_reader.py @@ -9,15 +9,18 @@ import numpy as np import xarray as xr import xugrid as xu -from attrs import define from flopy.discretization import StructuredGrid +from pydantic import ConfigDict +from pydantic.dataclasses import dataclass from flopy4.adapters import read_binary_grid_file from flopy4.mf6.utils.grid import get_coords from flopy4.mf6.utils.time import assign_datetime_coords +_CFG = ConfigDict(arbitrary_types_allowed=True, extra="forbid") -@define + +@dataclass(config=_CFG) class Imeth1Header: kstp: int kper: int @@ -32,7 +35,7 @@ class Imeth1Header: pos: int -@define +@dataclass(config=_CFG) class Imeth6Header: kstp: int kper: int diff --git a/flopy4/mf6/utils/codegen/dfn2py.py b/flopy4/mf6/utils/codegen/dfn2py.py index bf2b3726..2e381d74 100644 --- a/flopy4/mf6/utils/codegen/dfn2py.py +++ b/flopy4/mf6/utils/codegen/dfn2py.py @@ -19,7 +19,7 @@ # *g / *a variants: gridded/array package variants, deferred. # # TODO (subpackage tier): detect `# flopy subpackage` DFN annotations and emit -# a typed child attrs field (e.g. ncf: Optional[Ncf]) alongside the existing path +# a typed child field (e.g. ncf: Optional[Ncf]) alongside the existing path # field; DisBase.write() already establishes the write pattern for NCF. # utl-ts also needs period values referencing timeseries by name written as strings. _SKIP = { diff --git a/flopy4/mf6/utils/codegen/filters.py b/flopy4/mf6/utils/codegen/filters.py index d4df64d2..0b2efe2b 100644 --- a/flopy4/mf6/utils/codegen/filters.py +++ b/flopy4/mf6/utils/codegen/filters.py @@ -273,11 +273,11 @@ def _record_child_supported(c: FieldV3) -> bool: def can_generate_record_class(f: FieldV3) -> bool: - """True when a compound record should be rendered as an inner attrs class. + """True when a compound record should be rendered as an inner dataclass. All non-file records whose children are entirely scalars, keywords, nested inline arrays, and/or nested records (see - _record_child_supported) become inner attrs classes. The first keyword + _record_child_supported) become inner dataclasses. The first keyword child (if any) is the trigger token (``_keyword``); remaining keyword children become ``Optional[bool]`` fields. A child that is itself a Record becomes its own composed class rather than being flattened in. @@ -619,12 +619,12 @@ def item_class( ``role="nested_union"`` column -- it qualifies that field's forward-reference union annotation (``"Oc.All | Oc.First | ..."``). - Produces a 4-space-indented ``@attrs.define`` class whose fields carry + Produces a 4-space-indented ``@dataclass(config=CFG)`` class whose fields carry real metadata (``index=``/``pk=``/``fk=``/``cellid=``/``time_series=``/ ``prefix=``/``tagged=``, via ``field()``) -- the class itself is the schema. Required fields (no default) are declared before optional fields to - satisfy attrs ordering constraints. + satisfy dataclass field-ordering constraints. ``is_period=True`` injects ``aux: tuple = ()`` between required value columns and optional columns, for packages that accept positional AUXILIARY @@ -765,8 +765,8 @@ def _field_line(col: dict, *, optional: bool) -> str: if meta: return f" {col['name']}: {py_type} = field({margs})" # A bare annotation here is equivalent to field() at runtime (both - # mean "no default") -- but mypy's attrs plugin doesn't recognize - # field() (a flopy4.mf6.spec wrapper, not attrs.field itself) as a + # mean "no default") -- but mypy's dataclass_transform support doesn't recognize + # field() (a flopy4.mf6.spec wrapper, not pydantic.Field itself) as a # field specifier, so it can't tell field()-declared columns above # (e.g. pk=/cellid=) don't actually have a default either. Left # bare, that misreading makes mypy treat *this* column as a @@ -796,7 +796,7 @@ def _field_line(col: dict, *, optional: bool) -> str: optional_non_boundname = [col for col in optional if col["role"] != "boundname"] boundname_cols = [col for col in optional if col["role"] == "boundname"] - lines = [" @attrs.define"] + lines = [" @dataclass(config=CFG)"] lines.append(f" class {class_name}(Item):") if keyword: lines.append(f' _keyword: ClassVar[str] = "{keyword}"') @@ -969,7 +969,7 @@ def collision_names( A name is a collision when it appears in more than one static list block, OR when it appears in any block AND is reserved by a period field. The - latter ensures that static block attrs never shadow bare period field names. + latter ensures that static block fields never shadow bare period field names. Prefix columns are excluded since they produce no attr. """ names = [col.name for cols in block_schemas.values() for col in cols if not col.is_prefix] diff --git a/flopy4/mf6/utils/codegen/make.py b/flopy4/mf6/utils/codegen/make.py index c42ab5ea..6f911bdd 100644 --- a/flopy4/mf6/utils/codegen/make.py +++ b/flopy4/mf6/utils/codegen/make.py @@ -53,7 +53,7 @@ class FieldSpec: @dataclass class InnerClassFieldSpec: - """Pre-computed context for one field of an inner attrs class.""" + """Pre-computed context for one field of an inner dataclass.""" py_name: str type_annotation: str @@ -64,7 +64,7 @@ class InnerClassFieldSpec: @dataclass class InnerClassSpec: - """Pre-computed context for a generated inner attrs class.""" + """Pre-computed context for a generated inner dataclass.""" class_name: str keyword: str @@ -109,7 +109,7 @@ class PeriodArmSpec: @dataclass class ComputedFieldSpec: """Pre-computed context for a read-only computed property, replacing a - stored attrs field entirely -- e.g. ``maxbound``, derived live from + stored field entirely -- e.g. ``maxbound``, derived live from ``stress_period_data``'s row counts rather than stored and kept in sync by hand (see ``build_component_spec``'s ``_maxbound_is_computed`` for when this applies).""" @@ -159,7 +159,7 @@ def _schema_dict_from_columns( accumulated and attached as a 'prefix' key on the next value column so the codec can emit the fixed token(s) before the value. is_row_keyword columns (optional keywords, e.g. MIXED) get role 'inline_keyword'. aux columns are - excluded -- appended dynamically in __attrs_post_init__. + excluded -- appended dynamically in __post_init__. ``nested_arm_classes``, when given, maps a column name to sibling arm class names already built for it (see ``_build_arm_specs_from_union``) @@ -325,8 +325,8 @@ def _ml_field( Produces continuation lines pre-indented at 8 spaces (args) and 4 spaces (closing paren) so the Jinja template can render it verbatim after `` {name}: {type} = ``. ``metadata`` here is the set of ``field()``/ - ``path()`` kwargs (block, schema, fill_forward, ...), not a raw attrs - metadata dict -- generated fields are plain attrs fields, so they go + ``path()`` kwargs (block, schema, fill_forward, ...), not a raw + metadata dict -- generated fields are plain dataclass fields, so they go through the same passive-metadata constructors hand-written classes use for their scalar fields. """ @@ -580,7 +580,7 @@ def _process_child(child: FieldV3) -> None: for child in data_children: _process_child(child) - # attrs requires fields with defaults to follow fields without defaults. + # Dataclass fields with defaults must follow fields without defaults. inner_fields.sort(key=lambda field: str(field.optional)) words = _strip_record_words(f.name) @@ -783,15 +783,35 @@ def _generated_imports( if typing_parts: stdlib.append(f"from typing import {', '.join(sorted(typing_parts))}") - third_party: list[str] = ["import attrs"] + # Bare pydantic Field() (as opposed to the flopy4.mf6.spec field()/ + # path() wrappers, imported separately below via _spec_parts) is only + # ever emitted by the template for spec.inner_classes -- composed + # Record fields (e.g. Oc.Headprint.formatrecord). Every top-level + # package field and every item_class()-rendered Item/Record field + # routes through field()/path() instead. Confirmed empirically: a + # generated file with no inner_classes but an unconditional Field + # import left 49 F401 (unused import) errors across the regenerated + # corpus before this was scoped to has_inner_classes. + _pydantic_parts = ["Field"] if has_inner_classes else [] + if has_period_schema: + # has_period_schema is already the OR of period_schema/block_schemas/ + # period_arms (see the call site) -- SkipValidation is needed + # whenever any Item-list field is generated (see item.py's + # item_list_type()/Package._init_item_lists() for why: pydantic + # would otherwise validate a raw tuple/dict input eagerly). + _pydantic_parts.append("SkipValidation") + third_party: list[str] = [] + if _pydantic_parts: + third_party.append(f"from pydantic import {', '.join(sorted(_pydantic_parts))}") + third_party.append("from pydantic.dataclasses import dataclass") if has_array: third_party.append("import numpy as np") third_party.append("from numpy.typing import NDArray") _base_imports = { - "Package": "from flopy4.mf6.package import Package", - "Solution": "from flopy4.mf6.solution import Solution", - "Context": "from flopy4.mf6.context import Context", + "Package": "from flopy4.mf6.package import CFG, Package", + "Solution": "from flopy4.mf6.solution import CFG, Solution", + "Context": "from flopy4.mf6.context import CFG, Context", } flopy4: list[str] = [_base_imports.get(base_class, _base_imports["Package"])] if has_inner_classes: @@ -910,7 +930,7 @@ def build_component_spec( for block_name, f in all_fields: if filters.is_list_field(f) and block_name in _bp_block_names: - continue # covered by BlockPropertySpec; column attrs generated below + continue # covered by BlockPropertySpec; column fields generated below if block_name in _fill_forward_blocks and filters.is_list_field(f): union = filters.find_keystring_union(f) @@ -1053,7 +1073,7 @@ def build_component_spec( FieldSpec( dfn_name=bp.block_name, py_name=bp.block_name, - type_annotation=f"Optional[list[{_item_cls_name}]]", + type_annotation=f"Optional[SkipValidation[list[{_item_cls_name}]]]", spec_call=_ml_field(metadata=_meta), generatable=True, ) @@ -1073,7 +1093,9 @@ def build_component_spec( FieldSpec( dfn_name="_stress_period_data", py_name="_stress_period_data", - type_annotation="Optional[dict[int, list[_StressPeriodDataItem]]]", + type_annotation=( + "Optional[SkipValidation[dict[int, list[_StressPeriodDataItem]]]]" + ), spec_call=_ml_field(alias="stress_period_data", repr_=False, metadata=_spd_meta), generatable=True, ) @@ -1084,13 +1106,13 @@ def build_component_spec( FieldSpec( dfn_name="_stress_period_data", py_name="_stress_period_data", - type_annotation="Optional[dict[int, list[StressPeriodData]]]", + type_annotation="Optional[SkipValidation[dict[int, list[StressPeriodData]]]]", spec_call=_ml_field(alias="stress_period_data", repr_=False, metadata=_spd_meta), generatable=True, ) ) - # READARRAY period fields → individual Optional[Int|FloatArrayLike] attrs + # READARRAY period fields → individual Optional[Int|FloatArrayLike] # fields. G-variant packages (CHDG, DRNG, WELG, RCHA …) declare each period # array separately. Each field is a full-grid array passed directly by # the user; the egress side (unstructure.py's _unstructure_package) diff --git a/flopy4/mf6/utils/codegen/templates/package.py.jinja b/flopy4/mf6/utils/codegen/templates/package.py.jinja index 6941294e..1dc41cb9 100644 --- a/flopy4/mf6/utils/codegen/templates/package.py.jinja +++ b/flopy4/mf6/utils/codegen/templates/package.py.jinja @@ -14,7 +14,7 @@ {% endfor %} -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class {{ spec.class_name }}({{ spec.base_class }}): dfn_name: ClassVar[str] = "{{ spec.dfn_name }}" @@ -27,7 +27,7 @@ class {{ spec.class_name }}({{ spec.base_class }}): {% endif %} {% for rec in spec.inner_classes %} - @attrs.define + @dataclass(config=CFG) class {{ rec.class_name }}(Record): _keyword: ClassVar[str] = "{{ rec.keyword }}" {% if rec.extra_tokens_repr %} @@ -35,17 +35,17 @@ class {{ spec.class_name }}({{ spec.base_class }}): {% endif %} {% for f in rec.fields %} {% if f.nested and f.optional %} - {{ f.py_name }}: "Optional[{{ spec.class_name }}.{{ f.type_annotation }}]" = attrs.field(default=None) + {{ f.py_name }}: "Optional[{{ spec.class_name }}.{{ f.type_annotation }}]" = Field(default=None) {% elif f.nested %} - {{ f.py_name }}: "{{ spec.class_name }}.{{ f.type_annotation }}" = attrs.field() + {{ f.py_name }}: "{{ spec.class_name }}.{{ f.type_annotation }}" = Field() {% elif f.tagged and f.optional %} - {{ f.py_name }}: {{ f.type_annotation }} = attrs.field(default=None, metadata={"tagged": True}) + {{ f.py_name }}: {{ f.type_annotation }} = Field(default=None, json_schema_extra={"tagged": True}) {% elif f.tagged %} - {{ f.py_name }}: {{ f.type_annotation }} = attrs.field(metadata={"tagged": True}) + {{ f.py_name }}: {{ f.type_annotation }} = Field(json_schema_extra={"tagged": True}) {% elif f.optional %} - {{ f.py_name }}: {{ f.type_annotation }} = attrs.field(default=None) + {{ f.py_name }}: {{ f.type_annotation }} = Field(default=None) {% else %} - {{ f.py_name }}: {{ f.type_annotation }} = attrs.field() + {{ f.py_name }}: {{ f.type_annotation }} = Field() {% endif %} {% endfor %} diff --git a/flopy4/mf6/utl/ats.py b/flopy4/mf6/utl/ats.py index 6c660332..eb368691 100644 --- a/flopy4/mf6/utl/ats.py +++ b/flopy4/mf6/utl/ats.py @@ -1,18 +1,19 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ats(Package): dfn_name: ClassVar[str] = "utl-ats" - @attrs.define + @dataclass(config=CFG) class Perioddata(Item): iperats: int = field(index=True, pk=True) dt0: float = field() @@ -26,7 +27,7 @@ class Perioddata(Item): block="dimensions", longname="number of ATS periods", ) - perioddata: Optional[list[Perioddata]] = field( + perioddata: Optional[SkipValidation[list[Perioddata]]] = field( default=None, block="perioddata", auto_from="perioddata", diff --git a/flopy4/mf6/utl/hpc.py b/flopy4/mf6/utl/hpc.py index 0b1d0912..5eb2ac35 100644 --- a/flopy4/mf6/utl/hpc.py +++ b/flopy4/mf6/utl/hpc.py @@ -1,18 +1,19 @@ # autogenerated file, do not modify from typing import ClassVar, Optional, Union -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Hpc(Package): dfn_name: ClassVar[str] = "utl-hpc" - @attrs.define + @dataclass(config=CFG) class Partitions(Item): mname: Union[float, str] = field() mrank: int = field() @@ -23,7 +24,7 @@ class Partitions(Item): optional=True, longname="model print table to listing file", ) - partitions: Optional[list[Partitions]] = field( + partitions: Optional[SkipValidation[list[Partitions]]] = field( default=None, block="partitions", ) diff --git a/flopy4/mf6/utl/laktab.py b/flopy4/mf6/utl/laktab.py index 2b57235b..f9e3a3ee 100644 --- a/flopy4/mf6/utl/laktab.py +++ b/flopy4/mf6/utl/laktab.py @@ -1,20 +1,21 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Laktab(Package): dfn_name: ClassVar[str] = "utl-laktab" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Table(Item): stage: float = field() volume: float = field() @@ -31,7 +32,7 @@ class Table(Item): block="dimensions", longname="number of table columns", ) - table: Optional[list[Table]] = field( + table: Optional[SkipValidation[list[Table]]] = field( default=None, block="table", auto_from="table", diff --git a/flopy4/mf6/utl/ncf.py b/flopy4/mf6/utl/ncf.py index 60f8fa91..2d8f9196 100644 --- a/flopy4/mf6/utl/ncf.py +++ b/flopy4/mf6/utl/ncf.py @@ -1,16 +1,16 @@ from typing import ClassVar, Optional from warnings import warn -import attrs import numpy as np from numpy.typing import NDArray +from pydantic.dataclasses import dataclass from flopy4.mf6.enums import NetCDFFormat from flopy4.mf6.spec import field -from flopy4.mf6.utl.ncf_base import NcfBase +from flopy4.mf6.utl.ncf_base import CFG, NcfBase -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Ncf(NcfBase): dfn_name: ClassVar[str] = "utl-ncf" diff --git a/flopy4/mf6/utl/ncf_base.py b/flopy4/mf6/utl/ncf_base.py index 7e32d1c4..b4c1e02d 100644 --- a/flopy4/mf6/utl/ncf_base.py +++ b/flopy4/mf6/utl/ncf_base.py @@ -1,12 +1,12 @@ from typing import Optional -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class NcfBase(Package): """Field-type overrides for the utl-ncf CRS options; see Ncf for the rest.""" diff --git a/flopy4/mf6/utl/sfrtab.py b/flopy4/mf6/utl/sfrtab.py index 716e2613..c45332e0 100644 --- a/flopy4/mf6/utl/sfrtab.py +++ b/flopy4/mf6/utl/sfrtab.py @@ -1,20 +1,21 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic import SkipValidation +from pydantic.dataclasses import dataclass from flopy4.mf6.item import Item -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Sfrtab(Package): dfn_name: ClassVar[str] = "utl-sfrtab" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class Table(Item): xfraction: float = field() height: float = field() @@ -30,7 +31,7 @@ class Table(Item): block="dimensions", longname="number of table columns", ) - table: Optional[list[Table]] = field( + table: Optional[SkipValidation[list[Table]]] = field( default=None, block="table", auto_from="table", diff --git a/flopy4/mf6/utl/spca.py b/flopy4/mf6/utl/spca.py index ba4dd999..3b1673a2 100644 --- a/flopy4/mf6/utl/spca.py +++ b/flopy4/mf6/utl/spca.py @@ -2,14 +2,14 @@ from pathlib import Path from typing import ClassVar, Optional -import attrs +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike, _optional_path -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.spec import field, path -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Spca(Package): dfn_name: ClassVar[str] = "utl-spca" diff --git a/flopy4/mf6/utl/tas.py b/flopy4/mf6/utl/tas.py index 5d1b36e6..5cf290fd 100644 --- a/flopy4/mf6/utl/tas.py +++ b/flopy4/mf6/utl/tas.py @@ -1,34 +1,35 @@ # autogenerated file, do not modify from typing import ClassVar, Optional -import attrs +from pydantic import Field +from pydantic.dataclasses import dataclass from flopy4.mf6._types import FloatArrayLike -from flopy4.mf6.package import Package +from flopy4.mf6.package import CFG, Package from flopy4.mf6.record import Record from flopy4.mf6.spec import field -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class Tas(Package): dfn_name: ClassVar[str] = "utl-tas" multi_package: ClassVar[bool] = True - @attrs.define + @dataclass(config=CFG) class TimeSeriesName(Record): _keyword: ClassVar[str] = "name" - time_series_name: list[str] = attrs.field() + time_series_name: list[str] = Field() - @attrs.define + @dataclass(config=CFG) class InterpolationMethod(Record): _keyword: ClassVar[str] = "method" - interpolation_method: str = attrs.field() + interpolation_method: str = Field() - @attrs.define + @dataclass(config=CFG) class Sfac(Record): _keyword: ClassVar[str] = "sfac" - sfacval: list[float] = attrs.field() + sfacval: list[float] = Field() time_series_name: Optional[TimeSeriesName] = field( default=None, diff --git a/flopy4/mf6/write_context.py b/flopy4/mf6/write_context.py index eda6138a..de04faf5 100644 --- a/flopy4/mf6/write_context.py +++ b/flopy4/mf6/write_context.py @@ -3,15 +3,18 @@ import threading from typing import TYPE_CHECKING, ClassVar, Literal, Optional -from attrs import define, field +from pydantic import ConfigDict, Field +from pydantic.dataclasses import dataclass if TYPE_CHECKING: from threading import local ArrayFormat = Literal["internal", "constant", "open/close"] +_CFG = ConfigDict(arbitrary_types_allowed=True, validate_assignment=True, extra="forbid") -@define + +@dataclass(config=_CFG) class WriteContext: """ Configuration context for writing MODFLOW 6 input files. @@ -51,12 +54,12 @@ class WriteContext: ... sim.write() """ - use_binary: bool = field(default=False) - use_netcdf: bool = field(default=False) - binary_threshold: Optional[int] = field(default=None) - float_precision: int = field(default=8) - use_relative_paths: bool = field(default=True) - array_format: Optional[ArrayFormat] = field(default=None) + use_binary: bool = Field(default=False) + use_netcdf: bool = Field(default=False) + binary_threshold: Optional[int] = Field(default=None) + float_precision: int = Field(default=8) + use_relative_paths: bool = Field(default=True) + array_format: Optional[ArrayFormat] = Field(default=None) # Class-level thread-local storage for context stack _global_context_stack: ClassVar["local"] diff --git a/flopy4/mixins.py b/flopy4/mixins.py index 4414f05c..512a9b31 100644 --- a/flopy4/mixins.py +++ b/flopy4/mixins.py @@ -1,37 +1,37 @@ """ Mixins implementing the `DatasetConvertible`/`DataTreeConvertible` -protocols (`flopy4/protocols.py`) on top of the generic attrs<->xarray -conversion functions (`flopy4/attrs_xarray.py`). +protocols (`flopy4/protocols.py`) on top of the generic dataclass<->xarray +conversion functions (`flopy4/dataclass_xarray.py`). """ import xarray as xr -from flopy4.attrs_xarray import ( - attrs_to_dataset, - attrs_to_datatree, - dataset_to_attrs, - datatree_to_attrs, +from flopy4.dataclass_xarray import ( + dataclass_to_dataset, + dataclass_to_datatree, + dataset_to_dataclass, + datatree_to_dataclass, ) class DatasetConvertibleMixin: - """Mixin for leaf attrs classes with no attrs-typed child fields.""" + """Mixin for leaf dataclasses with no dataclass-typed child fields.""" def to_xarray(self) -> xr.Dataset: - return attrs_to_dataset(self) + return dataclass_to_dataset(self) @classmethod def from_dataset(cls, dataset: xr.Dataset): - return dataset_to_attrs(cls, dataset) + return dataset_to_dataclass(cls, dataset) class DataTreeConvertibleMixin: - """Mixin for internal-node attrs classes with one or more attrs-typed - child fields.""" + """Mixin for internal-node dataclasses with one or more + dataclass-typed child fields.""" def to_xarray(self) -> xr.DataTree: - return attrs_to_datatree(self) + return dataclass_to_datatree(self) @classmethod def from_datatree(cls, tree: xr.DataTree): - return datatree_to_attrs(cls, tree) + return datatree_to_dataclass(cls, tree) diff --git a/flopy4/protocols.py b/flopy4/protocols.py index 88d31978..231185ca 100644 --- a/flopy4/protocols.py +++ b/flopy4/protocols.py @@ -1,6 +1,6 @@ """ -Runtime-checkable protocols for the explicit attrs<->xarray conversion -layer (`flopy4/attrs_xarray.py`, `flopy4/mixins.py`). +Runtime-checkable protocols for the explicit dataclass<->xarray conversion +layer (`flopy4/dataclass_xarray.py`, `flopy4/mixins.py`). """ from typing import Protocol, runtime_checkable @@ -12,7 +12,7 @@ class DatasetConvertible(Protocol): """A leaf object: convertible to/from a flat `xr.Dataset`. - For attrs classes with no attrs-typed child fields (e.g. a DFN leaf + For dataclasses with no dataclass-typed child fields (e.g. a DFN leaf package with only scalar/array fields). """ @@ -27,7 +27,7 @@ class DataTreeConvertible(Protocol): """An internal-node object: convertible to/from a hierarchical `xr.DataTree`. - For attrs classes with one or more attrs-typed child fields (single, + For dataclasses with one or more dataclass-typed child fields (single, list, or dict), whose own scalar/array fields form the tree's root dataset and whose children form named child nodes. """ diff --git a/flopy4/spec.py b/flopy4/spec.py index 74963cc7..3f61ddad 100644 --- a/flopy4/spec.py +++ b/flopy4/spec.py @@ -1,14 +1,47 @@ """ -Wrap `attrs` specification utilities. +Generic pydantic-dataclass specification utilities. """ -from attrs import Attribute -from attrs import fields_dict as attrs_fields_dict +from typing import Any +from pydantic.dataclasses import is_pydantic_dataclass +from pydantic.fields import FieldInfo -def fields_dict(cls) -> dict[str, Attribute]: + +def is_dataclass_instance(value: Any) -> bool: + """True if `value` is an instance of a pydantic dataclass. Used + wherever generic tree-walking code needs to tell a nested + dataclass-typed value apart from a plain scalar/array leaf value.""" + return is_pydantic_dataclass(type(value)) + + +def pydantic_fields(cls: type) -> dict[str, FieldInfo]: + """A pydantic dataclass's fields, by name. + + Use this rather than `cls.__pydantic_fields__` directly: pydantic sets + that attribute at runtime but doesn't declare it on the decorated + class's type, so mypy only sees it after `is_pydantic_dataclass()` + narrows `cls`. + """ + if not is_pydantic_dataclass(cls): + raise TypeError(f"{cls!r} is not a pydantic dataclass") + return cls.__pydantic_fields__ + + +def field_meta(finfo: FieldInfo) -> dict[str, Any]: + """A field's flopy4 metadata (`block`, `shape`, ...), or `{}`. + + Stored in `json_schema_extra`, which pydantic types as a JSON dict or + a callable; flopy4 only ever stores a dict there. + """ + extra = finfo.json_schema_extra + return extra if isinstance(extra, dict) else {} + + +def fields_dict(cls) -> dict[str, Any]: """ Return an ordered dictionary of fields for a component class, - whose keys are field names. Each field is an `attrs.Attribute`. + whose keys are field names. Each field is a pydantic `FieldInfo`. + Empty for anything that isn't a pydantic dataclass. """ - return dict(attrs_fields_dict(cls)) + return dict(pydantic_fields(cls)) if is_pydantic_dataclass(cls) else {} diff --git a/pixi.lock b/pixi.lock index 1878f69f..e6eeb4e9 100644 --- a/pixi.lock +++ b/pixi.lock @@ -10817,8 +10817,6 @@ packages: name: flopy4 requires_dist: - modflow-devtools[ecosystem] @ git+https://github.com/MODFLOW-ORG/modflow-devtools.git - - attrs>=25.4.0,<26 - - cattrs>=25.3.0,<26 - flopy>=3.9.5,<4 - jinja2>=3.1.6,<4 - lark>=1.3.1,<2 diff --git a/pyproject.toml b/pyproject.toml index 046e59dc..b34d2e5e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,8 +35,6 @@ classifiers = [ requires-python = ">=3.11,<3.14" dependencies = [ "modflow-devtools[ecosystem] @ git+https://github.com/MODFLOW-ORG/modflow-devtools.git", - "attrs>=25.4.0,<26", - "cattrs>=25.3.0,<26", "flopy>=3.9.5,<4", "jinja2>=3.1.6,<4", "lark>=1.3.1,<2", diff --git a/test/mf6/test_mf6_adapters.py b/test/mf6/test_mf6_adapters.py index 39ad671d..eaf97aa0 100644 --- a/test/mf6/test_mf6_adapters.py +++ b/test/mf6/test_mf6_adapters.py @@ -71,8 +71,13 @@ def test_flopy3_model(tmp_path): inner_dvclose=1e-6, linear_acceleration="cg", ) - ims.inner_hclose = 1e-6 - ims.inner_rclose = 0.1000000 + # Legacy (pre-current-DFN) attribute names, not real Ims fields. + # extra="forbid" (needed for init=False derived fields like DisBase's + # nlay/nrow/...) rejects an ordinary `ims.inner_hclose = ...`, so this + # uses the same object.__setattr__ escape hatch the source uses + # internally (e.g. Package._broadcast_griddata) to bypass validation. + object.__setattr__(ims, "inner_hclose", 1e-6) + object.__setattr__(ims, "inner_rclose", 0.1000000) ic = Ic(dims=dims) oc = Oc(dims=dims, stress_period_data={0: [("SAVE", "HEAD", "ALL"), ("SAVE", "BUDGET", "ALL")]}) npf = Npf(dims=dims) @@ -222,13 +227,21 @@ def test_flopy3_package(tmp_path): assert not dis3.has_stress_period_data # package data - # Flopy3Package.data_list is built from attrs_to_dataset(dis): every + # Flopy3Package.data_list is built from dataclass_to_dataset(dis): every # non-None scalar field first (dataset-level .attrs, in declaration # order), then every array field (.data_vars, in declaration order) -- - # not a curated flopy3-only subset. See attrs_to_dataset()'s own - # docstring (flopy4/attrs_xarray.py) for the scalar/array split rule. + # not a curated flopy3-only subset. See dataclass_to_dataset()'s own + # docstring (flopy4/dataclass_xarray.py) for the scalar/array split rule. + # + # nlay/nrow/ncol/ncpl/nvert/nodes' relative order: DisBase declares + # them all together (as init=False derived fields); Dis then redeclares + # nlay/ncol/nrow as its own required fields. A redeclared dataclass + # field keeps the base class's original position. data_list = [ "name", + "nlay", + "nrow", + "ncol", "ncpl", "nvert", "nodes", @@ -236,9 +249,6 @@ def test_flopy3_package(tmp_path): "xorigin", "yorigin", "export_array_netcdf", - "nlay", - "ncol", - "nrow", "delr", "delc", "top", diff --git a/test/mf6/test_mf6_codec.py b/test/mf6/test_mf6_codec.py index ceae13fc..906c8fb2 100644 --- a/test/mf6/test_mf6_codec.py +++ b/test/mf6/test_mf6_codec.py @@ -170,7 +170,10 @@ def test_dumps_dis_with_constant_arrays(dis_with_constant_arrays): pprint(loaded) assert ["LENGTH_UNITS", "feet"] in loaded["OPTIONS"] - assert loaded["DIMENSIONS"] == [["NLAY", 2], ["NCOL", 10], ["NROW", 10]] + # NROW/NCOL order (not NCOL/NROW): DisBase declares nlay/nrow/ncol/... + # together; Dis redeclares nlay/ncol/nrow as its own fields, and a + # redeclared dataclass field keeps the base class's original position. + assert loaded["DIMENSIONS"] == [["NLAY", 2], ["NROW", 10], ["NCOL", 10]] assert ["DELR"] in loaded["GRIDDATA"] assert ["DELC"] in loaded["GRIDDATA"] @@ -1460,7 +1463,14 @@ def test_ssm_fileinput_row_format(): fileinput={ "pname": np.array(["rch-1", "wel-1"]), "spc6_filename": np.array(["rch.spc6", "wel.spc6"]), - "mixed": np.array([True, False]), + # `mixed` is a bare-presence-flag tagged field (see item.py's + # module docstring) but declared Optional[str] (item.py's + # codegen types every "inline_keyword"-role tagged field str, + # regardless of how it's actually used) -- to_tokens() only + # ever checks its truthiness (`if val:`), so a real string, + # not a bool, is the field's actual contract (a raw + # np.True_/np.False_ fails Optional[str] validation). + "mixed": np.array(["MIXED", ""], dtype=object), }, ) text = dumps(unstructure_component(ssm)) @@ -1620,10 +1630,12 @@ def test_rclose_from_tokens_with_option(): def test_from_tokens_missing_required_raises(): - """attrs raises TypeError when a required field has no value.""" + """A required field with no value raises pydantic's ValidationError.""" + from pydantic import ValidationError + from flopy4.mf6.gwf.oc import Oc - with pytest.raises(TypeError): + with pytest.raises(ValidationError): Oc.Headprint.from_tokens("") diff --git a/test/mf6/test_mf6_codegen.py b/test/mf6/test_mf6_codegen.py index 61e690ac..cbaa91d3 100644 --- a/test/mf6/test_mf6_codegen.py +++ b/test/mf6/test_mf6_codegen.py @@ -254,7 +254,7 @@ def test_item_class_static_block_no_aux(self): {"name": "boundname", "role": "boundname", "dfn_type": "string"}, ] result = item_class(schema, "Packagedata") - assert "@attrs.define" in result + assert "@dataclass(config=CFG)" in result assert "class Packagedata(Item):" in result assert "ifno: int = field(index=True, pk=True)" in result assert "strt: float" in result @@ -370,7 +370,7 @@ def test_imports_include_base(self, dfn_name, expected_class, expected_base, all + spec.imports.get("third_party", []) + spec.imports.get("flopy4", []) ) - assert "attrs" in all_imports + assert "pydantic" in all_imports assert "Package" in all_imports def test_outpath(self, dfn_name, expected_class, expected_base, all_dfns): @@ -424,7 +424,7 @@ def test_imports_include_solution( + spec.imports.get("third_party", []) + spec.imports.get("flopy4", []) ) - assert "attrs" in all_imports + assert "pydantic" in all_imports assert "Solution" in all_imports assert "ClassVar" in all_imports @@ -458,7 +458,8 @@ def test_mvr_list_fields_expanded_and_optional(all_dfns): assert spec.period_schema, "MVR should have a period_schema" assert "_stress_period_data" in field_map spd_field = field_map["_stress_period_data"] - assert spd_field.type_annotation == "Optional[dict[int, list[StressPeriodData]]]" + expected_type = "Optional[SkipValidation[dict[int, list[StressPeriodData]]]]" + assert spd_field.type_annotation == expected_type # Packages block → single recarray field assert "packages" in field_map or "packages" in spec.block_schemas @@ -501,7 +502,7 @@ def test_lak_outlets_dim_declared(self, lak_spec): assert bp.dim_is_dfn_declared is True def test_lak_ifno_collision_prefixed(self, lak_spec): - # ifno appears in packagedata, connectiondata, and tables — all get block-prefixed attrs + # ifno appears in packagedata, connectiondata, and tables — all get block-prefixed fields for block in ("packagedata", "connectiondata", "tables"): bp = next(b for b in lak_spec.block_properties if b.block_name == block) assert bp.attr_name_map.get("ifno") == f"{block}_ifno", ( diff --git a/test/mf6/test_mf6_component.py b/test/mf6/test_mf6_component.py index 6b655f6a..08f32050 100644 --- a/test/mf6/test_mf6_component.py +++ b/test/mf6/test_mf6_component.py @@ -294,7 +294,7 @@ def test_init_sim_explicit_dims(): assert gwf.oc is oc assert gwf.npf is npf assert gwf.chd[0] is chd - # k is stored as a plain attrs field; use to_xarray() for Dataset access + # k is stored as a plain field; use to_xarray() for Dataset access assert np.array_equal(sim.models["gwf"].npf.k, np.ones(100)) assert np.array_equal(sim.models["gwf"].npf.to_xarray()["k"].values, np.ones((1, 10, 10))) @@ -536,7 +536,7 @@ def test_to_xarray_on_context(function_tmpdir): assert isinstance(dt, xr.DataTree) assert isinstance(dt.kper, xr.DataArray) assert np.array_equal(dt.kper, [0]) - assert dt.attrs["filename"] == "mfsim.nam" + assert dt.attrs["filename"] == Path("mfsim.nam") assert dt.attrs["workspace"] == Path(function_tmpdir) @@ -1598,15 +1598,13 @@ def test_disv_class_identity(): def test_prt_dis_no_ncf(): """prt.Dis and prt.Disv must not expose NCF fields.""" - import attrs - from flopy4.mf6 import prt - dis_field_names = {f.name for f in attrs.fields(prt.Dis)} + dis_field_names = set(prt.Dis.__pydantic_fields__) assert "ncf6_filerecord" not in dis_field_names assert "ncf" not in dis_field_names - disv_field_names = {f.name for f in attrs.fields(prt.Disv)} + disv_field_names = set(prt.Disv.__pydantic_fields__) assert "ncf6_filerecord" not in disv_field_names assert "ncf" not in disv_field_names @@ -1721,7 +1719,7 @@ def test_explicit_parent_top_down(): def test_explicit_parent_bottom_up_non_package(): """The bottom-up (`parent=`) half of _parent tracking works for a non-Package component (Model/Context/Component's own chain), whose - own __attrs_post_init__ chains to Component's.""" + own __post_init__ chains to Component's.""" sim = Simulation() gwf = Gwf(parent=sim, name="gwf") assert gwf._parent is sim @@ -1729,8 +1727,8 @@ def test_explicit_parent_bottom_up_non_package(): def test_explicit_parent_bottom_up_package(): """The bottom-up (`parent=`) half of `_parent` tracking also works - for a Package subclass, which requires Package.__attrs_post_init__ - to chain to Component.__attrs_post_init__ via super() -- see that + for a Package subclass, which requires Package.__post_init__ + to chain to Component.__post_init__ via super() -- see that method's own docstring for why the chaining order matters. """ gwf = Gwf() @@ -1768,3 +1766,47 @@ def test_parent_setter_detach(): assert ic.parent is None assert ic._parent is None assert gwf.ic is None + + +def test_eq_ignores_dims_and_parent(): + """Equality ignores `dims` and `_parent`, on subclasses too (each + generated dataclass gets its own `__eq__`).""" + assert Ims(dims={"nodes": 3}) == Ims(dims={"nodes": 4}) + assert Ims(parent=Simulation()) == Ims() + assert Ims(inner_maximum=10) != Ims(inner_maximum=20) + + +def test_filename_is_path(): + """`filename` accepts a str or Path, is stored as a Path, and goes into + the name file with POSIX separators.""" + from pathlib import PureWindowsPath + + from flopy4.mf6.codec.writer.filters import quote_if_needed + from flopy4.mf6.converter.binding import Binding + + ic = Ic(filename="gwf.ic") + assert ic.filename == Path("gwf.ic") + ic.filename = Path("sub") / "gwf.ic" + assert Binding.from_component(ic).fname == "sub/gwf.ic" + assert quote_if_needed(PureWindowsPath("sub\\gwf.ic")) == "sub/gwf.ic" + + +def test_external_array_path_is_posix(): + """An external array's OPEN/CLOSE path is written with POSIX separators.""" + from pathlib import PureWindowsPath + + from flopy4.mf6.codec.writer import _JINJA_ENV + + macros = _JINJA_ENV.get_template("macros.jinja").module + out = str(macros.array("top", PureWindowsPath("data\\top.dat"), how="external")) # type: ignore[attr-defined] + assert "OPEN/CLOSE data/top.dat" in out + + +def test_layered_int_griddata_keeps_int_dtype(): + # A scalar layered integer array (idomain=1) is repeated per layer in + # _coerce_griddata; it must stay integer, or it's written as a float + # CONSTANT, which some MF6 builds reject for an integer array. + from flopy4.mf6.gwf import Dis + + dis = Dis(nlay=2, nrow=1, ncol=3, top=1.0, botm=[0.0, -1.0], idomain=1) + assert dis.idomain.dtype == np.int64 diff --git a/test/mf6/test_mf6_integration.py b/test/mf6/test_mf6_integration.py index 0af8bb0d..2f367cef 100644 --- a/test/mf6/test_mf6_integration.py +++ b/test/mf6/test_mf6_integration.py @@ -764,7 +764,7 @@ def test_quickstart_netcdf(function_tmpdir): with open(function_tmpdir / f"{gwf_name}.nam", "r") as fh: lines = fh.readlines() nc_fpth = function_tmpdir / f"{gwf_name}.input.nc" - assert f" NETCDF FILEIN {nc_fpth}\n" in lines + assert f" NETCDF FILEIN {nc_fpth.as_posix()}\n" in lines with open(function_tmpdir / f"{gwf_name}.dis", "r") as fh: lines = fh.readlines() assert " DELR NETCDF\n" in lines @@ -876,7 +876,7 @@ def test_quickstart_netcdf_mesh(function_tmpdir): with open(function_tmpdir / f"{gwf_name}.nam", "r") as fh: lines = fh.readlines() nc_fpth = function_tmpdir / f"{gwf_name}.input.nc" - assert f" NETCDF FILEIN {nc_fpth}\n" in lines + assert f" NETCDF FILEIN {nc_fpth.as_posix()}\n" in lines with open(function_tmpdir / f"{gwf_name}.dis", "r") as fh: lines = fh.readlines() assert " DELR NETCDF\n" in lines diff --git a/test/mf6/test_mf6_spec.py b/test/mf6/test_mf6_spec.py index 21ec9bfa..ac9f7aaf 100644 --- a/test/mf6/test_mf6_spec.py +++ b/test/mf6/test_mf6_spec.py @@ -5,7 +5,9 @@ def test_blocks(): block_spec = blocks(Gwf) options = block_spec[0] - assert options[-1].name == "netcdf_input_file" + options_by_name = blocks_dict(Gwf)["options"] + assert next(reversed(options_by_name)) == "netcdf_input_file" + assert len(options) == len(options_by_name) def test_blocks_dict(): diff --git a/test/test_attrs_xarray.py b/test/test_dataclass_xarray.py similarity index 74% rename from test/test_attrs_xarray.py rename to test/test_dataclass_xarray.py index 24bc25d0..4587eedc 100644 --- a/test/test_attrs_xarray.py +++ b/test/test_dataclass_xarray.py @@ -1,7 +1,7 @@ -"""Tests for the generic attrs<->xarray conversion layer -(flopy4/attrs_xarray.py, flopy4/protocols.py, flopy4/mixins.py). +"""Tests for the generic dataclass<->xarray conversion layer +(flopy4/dataclass_xarray.py, flopy4/protocols.py, flopy4/mixins.py). -Uses synthetic attrs classes to validate the conversion functions in +Uses synthetic dataclasses to validate the conversion functions in isolation, plus one round-trip against a real leaf MF6 package (Dis) to confirm it also works unmodified on an actual generated class. """ @@ -10,67 +10,70 @@ import numpy as np import xarray as xr -from attrs import define, field - -from flopy4.attrs_xarray import ( - attrs_to_dataset, - attrs_to_datatree, - dataset_to_attrs, - datatree_to_attrs, +from pydantic import ConfigDict, Field +from pydantic.dataclasses import dataclass + +from flopy4.dataclass_xarray import ( + dataclass_to_dataset, + dataclass_to_datatree, + dataset_to_dataclass, + datatree_to_dataclass, ) from flopy4.mf6.gwf.dis import Dis from flopy4.mf6.spec import field as mf6_field from flopy4.mixins import DatasetConvertibleMixin, DataTreeConvertibleMixin from flopy4.protocols import DatasetConvertible, DataTreeConvertible +_CFG = ConfigDict(arbitrary_types_allowed=True, extra="forbid") + -@define +@dataclass(config=_CFG) class Leaf(DatasetConvertibleMixin): """A leaf class with only scalar/array fields, no children.""" name: str = "leaf" value: float = 1.0 - # Deliberately not named "dims"/"parent"/"_parent" -- flopy4.attrs_xarray + # Deliberately not named "dims"/"parent"/"_parent" -- flopy4.dataclass_xarray # always excludes those (see its module docstring); no real DFN field # uses them either. values: Optional[np.ndarray] = mf6_field(default=None, shape=("nlay",)) -@define +@dataclass(config=_CFG) class OnlyChild(DatasetConvertibleMixin): label: str = "child" -@define +@dataclass(config=_CFG) class ListChild(DatasetConvertibleMixin): idx: int = 0 -@define +@dataclass(config=_CFG) class DictChild(DatasetConvertibleMixin): key: str = "k" -@define +@dataclass(config=_CFG) class Node(DataTreeConvertibleMixin): """An internal-node class with all three child-field kinds.""" title: str = "node" only: Optional[OnlyChild] = None - items: list[ListChild] = field(factory=list) - mapping: dict[str, DictChild] = field(factory=dict) + items: list[ListChild] = Field(default_factory=list) + mapping: dict[str, DictChild] = Field(default_factory=dict) def test_leaf_scalar_and_array_round_trip(): leaf = Leaf(name="a", value=2.5, values=np.array([1.0, 2.0, 3.0])) - ds = attrs_to_dataset(leaf) + ds = dataclass_to_dataset(leaf) assert isinstance(ds, xr.Dataset) assert ds.attrs["name"] == "a" assert ds.attrs["value"] == 2.5 assert list(ds["values"].dims) == ["nlay"] np.testing.assert_array_equal(ds["values"].values, [1.0, 2.0, 3.0]) - rebuilt = dataset_to_attrs(Leaf, ds) + rebuilt = dataset_to_dataclass(Leaf, ds) assert rebuilt.name == "a" assert rebuilt.value == 2.5 np.testing.assert_array_equal(rebuilt.values, [1.0, 2.0, 3.0]) @@ -87,12 +90,12 @@ def test_leaf_mixin_to_xarray_and_from_dataset(): def test_node_single_child_round_trip(): node = Node(title="root", only=OnlyChild(label="x")) - tree = attrs_to_datatree(node) + tree = dataclass_to_datatree(node) assert isinstance(tree, xr.DataTree) assert "only" in tree.children assert tree.dataset.attrs["title"] == "root" - rebuilt = datatree_to_attrs(Node, tree) + rebuilt = datatree_to_dataclass(Node, tree) assert rebuilt.title == "root" assert isinstance(rebuilt.only, OnlyChild) assert rebuilt.only.label == "x" @@ -100,26 +103,26 @@ def test_node_single_child_round_trip(): def test_node_list_children_round_trip(): node = Node(items=[ListChild(idx=0), ListChild(idx=1), ListChild(idx=2)]) - tree = attrs_to_datatree(node) + tree = dataclass_to_datatree(node) assert "items0" in tree.children assert "items1" in tree.children assert "items2" in tree.children - rebuilt = datatree_to_attrs(Node, tree) + rebuilt = datatree_to_dataclass(Node, tree) assert [c.idx for c in rebuilt.items] == [0, 1, 2] def test_node_dict_children_not_reconstructed(): """Documented limitation: dict-kind children round-trip into the tree (by their real dict key) but can't be recovered back into the dict - field from the tree alone -- see flopy4/attrs_xarray.py's module docstring. + field from the tree alone -- see flopy4/dataclass_xarray.py's module docstring. """ node = Node(mapping={"foo": DictChild(key="foo"), "bar": DictChild(key="bar")}) - tree = attrs_to_datatree(node) + tree = dataclass_to_datatree(node) assert "foo" in tree.children assert "bar" in tree.children - rebuilt = datatree_to_attrs(Node, tree) + rebuilt = datatree_to_dataclass(Node, tree) assert rebuilt.mapping == {} @@ -134,22 +137,22 @@ def test_node_mixin_to_xarray_and_from_datatree(): def test_empty_node_has_no_children(): node = Node() - tree = attrs_to_datatree(node) + tree = dataclass_to_datatree(node) assert dict(tree.children) == {} def test_dis_leaf_round_trip_real_mf6_class(): - """Dis has no attrs-typed children when ncf is unset, so it's a real + """Dis has no dataclass-typed children when ncf is unset, so it's a real (not synthetic) example of a DatasetConvertible-shaped leaf class.""" dis = Dis(nlay=2, nrow=3, ncol=4) - ds = attrs_to_dataset(dis) + ds = dataclass_to_dataset(dis) assert ds.attrs["nlay"] == 2 assert ds.attrs["nrow"] == 3 assert ds.attrs["ncol"] == 4 assert dict(ds["delr"].sizes) == {"ncol": 4} assert dict(ds["botm"].sizes) == {"nodes": 24} - rebuilt = dataset_to_attrs(Dis, ds) + rebuilt = dataset_to_dataclass(Dis, ds) assert rebuilt.nlay == 2 assert rebuilt.nrow == 3 assert rebuilt.ncol == 4 diff --git a/test/test_dimensions.py b/test/test_dimensions.py index d039fd92..7af18b30 100644 --- a/test/test_dimensions.py +++ b/test/test_dimensions.py @@ -2,12 +2,15 @@ from typing import Optional -from attrs import define, field +from pydantic import ConfigDict, Field +from pydantic.dataclasses import dataclass from flopy4.dimensions import DimensionResolverMixin +_CFG = ConfigDict(arbitrary_types_allowed=True, extra="forbid") -@define + +@dataclass(config=_CFG) class MockDimensionProvider: """Mock component that provides dimensions.""" @@ -26,25 +29,25 @@ def get_dims(self) -> dict[str, int]: } -@define +@dataclass(config=_CFG) class MockContainer(DimensionResolverMixin): """Mock container that uses the dimension registry mixin.""" provider: Optional[MockDimensionProvider] = None -@define +@dataclass(config=_CFG) class MockContainerWithDict(DimensionResolverMixin): """Mock container with dict of providers.""" - providers: dict[str, MockDimensionProvider] = field(factory=dict) + providers: dict[str, MockDimensionProvider] = Field(default_factory=dict) -@define +@dataclass(config=_CFG) class MockContainerWithList(DimensionResolverMixin): """Mock container with list of providers.""" - providers: list[MockDimensionProvider] = field(factory=list) + providers: list[MockDimensionProvider] = Field(default_factory=list) def test_resolve_dimension_from_direct_child(): diff --git a/test/test_uio.py b/test/test_uio.py index 54783aa5..6feac819 100644 --- a/test/test_uio.py +++ b/test/test_uio.py @@ -2,14 +2,14 @@ from pathlib import Path -import attrs +from pydantic.dataclasses import dataclass -from flopy4.mf6.component import Component +from flopy4.mf6.component import CFG, Component from flopy4.mf6.constants import MF6 from flopy4.uio import DEFAULT_REGISTRY, IO, Loader, Registry, Writer -@attrs.define(kw_only=True, slots=False) +@dataclass(config=CFG, kw_only=True) class MockComponent(Component): """Minimal test component for IO testing.""" @@ -100,7 +100,7 @@ def test_loader_registry_subclass_lookup(): """Test that registry correctly finds loaders for subclasses.""" test_registry = Registry() - @attrs.define(kw_only=True, slots=False) + @dataclass(config=CFG, kw_only=True) class SubComponent(MockComponent): """Subclass of MockComponent."""