Skip to content
Merged
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
21 changes: 20 additions & 1 deletion flopy4/mf6/dis_methods.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, ClassVar

import numpy as np
from flopy.discretization import StructuredGrid as LegacyStructuredGrid
Expand All @@ -15,6 +15,25 @@ class DisMethods:
"""Methods for the generated structured discretizations (`Dis` in GWF,
GWT, GWE and PRT); fields come from the DFN."""

# DFN fields these methods use, checked against the DFNs at sync time.
_requires: ClassVar[frozenset[str]] = frozenset(
{
"angrot",
"botm",
"crs",
"delc",
"delr",
"idomain",
"length_units",
"ncol",
"nlay",
"nrow",
"top",
"xorigin",
"yorigin",
}
)

def to_grid(self: "Dis") -> StructuredGrid: # type: ignore[misc]
"""Convert the discretization to a `StructuredGrid`."""
nlay, nrow, ncol = self.nlay, self.nrow, self.ncol
Expand Down
22 changes: 21 additions & 1 deletion flopy4/mf6/disu_methods.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, ClassVar

import numpy as np
from flopy.discretization.unstructuredgrid import UnstructuredGrid
Expand All @@ -11,6 +11,26 @@ class DisuMethods:
"""Methods for the generated unstructured discretizations (`Disu` in GWF,
GWT and GWE); fields come from the DFN."""

# DFN fields these methods use, checked against the DFNs at sync time.
_requires: ClassVar[frozenset[str]] = frozenset(
{
"angrot",
"bot",
"cell2d",
"crs",
"iac",
"idomain",
"ihc",
"ja",
"length_units",
"nodes",
"top",
"vertices",
"xorigin",
"yorigin",
}
)

def to_grid(self: "Disu") -> UnstructuredGrid: # type: ignore[misc]
"""Convert the discretization to an `UnstructuredGrid`."""
if self.nodes is None:
Expand Down
20 changes: 19 additions & 1 deletion flopy4/mf6/disv_methods.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from abc import ABCMeta
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, ClassVar

import numpy as np
from flopy.discretization import VertexGrid as LegacyVertexGrid
Expand Down Expand Up @@ -45,6 +45,24 @@ class DisvMethods(metaclass=_LegacyInputs):
"""Methods for the generated vertex discretizations (`Disv` in GWF, GWT,
GWE and PRT); fields come from the DFN."""

# DFN fields these methods use, checked against the DFNs at sync time.
_requires: ClassVar[frozenset[str]] = frozenset(
{
"angrot",
"botm",
"cell2d",
"crs",
"idomain",
"length_units",
"ncpl",
"nlay",
"top",
"vertices",
"xorigin",
"yorigin",
}
)

Cell2dRecord = _Cell2dRecord()

@property
Expand Down
12 changes: 11 additions & 1 deletion flopy4/mf6/tdis_methods.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, ClassVar, Optional

import numpy as np
from numpy.typing import ArrayLike
Expand All @@ -12,6 +12,16 @@
class TdisMethods:
"""Methods for the generated `Tdis`; fields come from the DFN."""

# DFN fields these methods use, checked against the DFNs at sync time.
_requires: ClassVar[frozenset[str]] = frozenset(
{
"nper",
"perioddata",
"start_date_time",
"time_units",
}
)

def to_time(self: "Tdis") -> Time: # type: ignore[misc]
"""Convert to a `Time` object."""
perlen, nstp, tsmult = zip(*((r.perlen, r.nstp, r.tsmult) for r in self.perioddata or []))
Expand Down
40 changes: 38 additions & 2 deletions flopy4/mf6/utils/codegen/make.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
schema this replaces.
"""

import ast
from collections.abc import Mapping
from dataclasses import dataclass
from dataclasses import field as dc_field
Expand Down Expand Up @@ -1013,7 +1014,8 @@ def _generated_imports(

# flopy API methods (factories, conversions to and from flopy types) for
# generated classes, as "module:Class" method-only mixins. Mixins declare no
# fields; those come from the DFN only.
# fields; those come from the DFN only. A mixin lists the DFN fields it uses
# in `_requires`, which check_mixins checks against the DFNs at sync time.
_DIS = ["flopy4.mf6.dis_methods:DisMethods"]
_DISV = ["flopy4.mf6.disv_methods:DisvMethods"]
_DISU = ["flopy4.mf6.disu_methods:DisuMethods"]
Expand Down Expand Up @@ -1047,9 +1049,43 @@ def _generated_imports(


def check_mixins(dfns: Mapping[str, Component]) -> None:
"""Raise if a `MIXINS` key names no DFN component."""
"""Raise if a `MIXINS` key names no DFN component, or a mixin's
`_requires` names a field its component's DFN lacks."""
if unknown := sorted(set(MIXINS) - set(dfns)):
raise ValueError(f"MIXINS keys match no DFN component: {unknown}")
lacking = []
for name, mixins in MIXINS.items():
fields = dfn_field_names(dfns[name])
for mixin in mixins:
if missing := sorted(mixin_requires(mixin) - fields):
lacking.append(f"{mixin.partition(':')[2]} on {name}: {', '.join(missing)}")
if lacking:
raise ValueError("Mixins need DFN fields that are missing:\n " + "\n ".join(lacking))


def dfn_field_names(component: Component) -> set[str]:
"""The names of a component's top-level fields, in every block."""
return {n for block in (component.blocks or {}).values() for n in (block.fields or {})}


def mixin_requires(mixin: str) -> frozenset[str]:
"""A ``module:Class`` mixin's ``_requires``. Read from its source, not
imported, since importing a mixin can import generated modules, and
sync has to work when those are broken."""
module, _, cls = mixin.partition(":")
path = _MF6_ROOT.joinpath(*module.split(".")[2:]).with_suffix(".py")
for node in ast.parse(path.read_text()).body:
if isinstance(node, ast.ClassDef) and node.name == cls:
for stmt in node.body:
if (
isinstance(stmt, ast.AnnAssign)
and isinstance(stmt.target, ast.Name)
and stmt.target.id == "_requires"
and isinstance(stmt.value, ast.Call)
):
return frozenset(ast.literal_eval(stmt.value.args[0]))
return frozenset()
raise ValueError(f"No class {cls} in {path}")


def _find_link(f: FieldV3) -> tuple[str, File, bool] | None:
Expand Down
9 changes: 8 additions & 1 deletion flopy4/mf6/utl/ts_methods.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from collections.abc import Mapping
from typing import TYPE_CHECKING, Optional, Union
from typing import TYPE_CHECKING, ClassVar, Optional, Union

import numpy as np
import pandas as pd
Expand All @@ -11,6 +11,13 @@
class TsMethods:
"""Methods for the generated `Ts`; fields come from the DFN."""

# DFN fields these methods use, checked against the DFNs at sync time.
_requires: ClassVar[frozenset[str]] = frozenset(
{
"timeseries",
}
)

@classmethod
def from_series( # type: ignore[misc]
cls: type["Ts"],
Expand Down
41 changes: 40 additions & 1 deletion test/mf6/test_mf6_codegen.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
- Add assertions specific to that tier's expected base class / extras.
"""

import ast
import importlib
import importlib.util
import sys
Expand Down Expand Up @@ -41,7 +42,14 @@
shared_dims,
value_column,
)
from flopy4.mf6.utils.codegen.make import build_component_spec, check_mixins, make_modules
from flopy4.mf6.utils.codegen.make import (
MIXINS,
build_component_spec,
check_mixins,
dfn_field_names,
make_modules,
mixin_requires,
)


# Shared fixtures
Expand Down Expand Up @@ -950,6 +958,37 @@ def test_check_mixins_rejects_unknown_component(all_dfns):
check_mixins({n: c for n, c in all_dfns.items() if n != "sim-tdis"})


def test_check_mixins_rejects_missing_field(all_dfns):
"""A DFN lacking a field a mixin needs fails the check, naming both."""
dis = all_dfns["gwf-dis"].model_copy(deep=True)
for block in dis.blocks.values():
(block.fields or {}).pop("crs", None)
with pytest.raises(ValueError, match="DisMethods on gwf-dis: crs"):
check_mixins({**all_dfns, "gwf-dis": dis})


def _self_attrs(mixin: str) -> set[str]:
module, _, cls = mixin.partition(":")
tree = ast.parse(Path(importlib.util.find_spec(module).origin).read_text())
node = next(n for n in tree.body if isinstance(n, ast.ClassDef) and n.name == cls)
return {
n.attr
for n in ast.walk(node)
if isinstance(n, ast.Attribute) and isinstance(n.value, ast.Name) and n.value.id == "self"
}


@pytest.mark.parametrize("name", sorted(MIXINS))
def test_mixin_requires_declared(all_dfns, name):
"""Every DFN field a mixin reads off ``self`` is in its ``_requires``,
and everything in ``_requires`` is a DFN field."""
fields = dfn_field_names(all_dfns[name])
for mixin in MIXINS[name]:
required = mixin_requires(mixin)
assert _self_attrs(mixin) & fields <= required, mixin
assert required <= fields, mixin


@pytest.mark.parametrize(
"selector,match",
[("bogus", "matches no component"), ("model", "no component_ftype")],
Expand Down
Loading