Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions docs/dev/sdd.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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 &mdash; blocks delimited by `BEGIN <name>` / `END <name>`, 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.
Expand Down
207 changes: 114 additions & 93 deletions flopy4/attrs_xarray.py → flopy4/dataclass_xarray.py

Large diffs are not rendered by default.

69 changes: 29 additions & 40 deletions flopy4/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]:
"""
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -271,51 +273,38 @@ 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"):
for dim in dims_needed:
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
2 changes: 1 addition & 1 deletion flopy4/mf6/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
14 changes: 12 additions & 2 deletions flopy4/mf6/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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".

Expand All @@ -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
Expand All @@ -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.
"""
Expand Down
26 changes: 13 additions & 13 deletions flopy4/mf6/adapters.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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.
"""
Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand All @@ -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(
Expand Down Expand Up @@ -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):
Expand Down
12 changes: 9 additions & 3 deletions flopy4/mf6/codec/writer/filters.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -150,15 +151,20 @@ 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.

WKT CRS strings (e.g. PROJCS["NAD83 / UTM zone 11N",...]) always
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)
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion flopy4/mf6/codec/writer/templates/macros.jinja
Original file line number Diff line number Diff line change
Expand Up @@ -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 %}
Loading
Loading