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
363 changes: 362 additions & 1 deletion autotest/dfns/test_schema_dimensions.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
"""Tests for DFN schema array shape expressions and dimension resolution"""

import ast
from collections import ChainMap

import pytest

from autotest.dfns.test_schema import _pkg
from autotest.dfns.test_schema import _DEV3_SNAPSHOT_DIR, _pkg
from modflow_devtools.dfns.schema import (
Array,
Block,
Dfns,
Double,
InputDim,
Integer,
List,
Expand All @@ -18,11 +20,15 @@
RuntimeDim,
String,
_names_in_expr,
_parse_shape_element,
_resolve_derived_dims,
_ShapeRef,
_validate_len_call,
_validate_list_shape_element,
_validate_shape_element,
_validate_sum_call,
dim_input,
dim_value,
split_bound,
)

Expand Down Expand Up @@ -531,6 +537,116 @@ def test_validate_shape_element_fk_block_mismatch():
_validate_shape_element("packagedata.nlakeconn(lakeno)", arr, lak, enc, known)


@pytest.mark.parametrize(
"element,expected",
[
("nvert", _ShapeRef("dim", "nvert")),
("nseg-1", _ShapeRef("dim", "nseg", offset=-1)),
("ncol + 1", _ShapeRef("dim", "ncol", offset=1)),
("<=maxats", _ShapeRef("dim", "maxats", bound="<=")),
(">= nper", _ShapeRef("dim", "nper", bound=">=")),
("<ncol - 2", _ShapeRef("dim", "ncol", offset=-2, bound="<")),
(
"packagedata.ncon(ifno)",
_ShapeRef("lookup", "ncon", block="packagedata", fk_field="ifno"),
),
(
"gwf-x.block.col(fk)",
_ShapeRef("lookup", "col", component="gwf-x", block="block", fk_field="fk"),
),
(
"<=packagedata.ncon(ifno)",
_ShapeRef("lookup", "ncon", bound="<=", block="packagedata", fk_field="ifno"),
),
],
)
def test_parse_shape_element(element, expected):
assert _parse_shape_element(element) == expected


@pytest.mark.parametrize(
"element,match",
[
("", "invalid shape element"),
("123bad", "invalid shape element"),
("nrow * 2", "invalid shape element"),
("<=nrow * 2", "invalid shape element"),
("len(auxiliary)", "invalid shape element"),
("packagedata.ncon", "invalid shape element"),
("<=<nrow", "at most one bound operator"),
],
)
def test_parse_shape_element_invalid(element, match):
with pytest.raises(ValueError, match=match):
_parse_shape_element(element)


def test_parse_shape_element_sibling():
ncvert = Integer(name="ncvert")
arr = Array(name="icvert", dtype="integer", shape=["ncvert"])
enc = Record(name="item", fields={"ncvert": ncvert, "icvert": arr})
assert _parse_shape_element("ncvert", enc) == _ShapeRef("sibling", "ncvert")
assert _parse_shape_element("ncvert") == _ShapeRef("dim", "ncvert")
assert _parse_shape_element("nvert", enc) == _ShapeRef("dim", "nvert")


def test_parse_shape_element_non_integer_not_sibling():
arr = Array(name="vals", dtype="double", shape=["n"])
enc = Record(name="item", fields={"n": String(name="n"), "vals": arr})
assert _parse_shape_element("n", enc).kind == "dim"


def test_validate_shape_element_sibling():
arr = Array(name="icvert", dtype="integer", shape=[])
enc = Record(name="item", fields={"ncvert": Integer(name="ncvert"), "icvert": arr})
_validate_shape_element("ncvert", arr, _pkg("test"), enc, set())


def test_validate_shape_element_sibling_shadows_dim():
arr = Array(name="icvert", dtype="integer", shape=[])
enc = Record(name="item", fields={"ncvert": Integer(name="ncvert"), "icvert": arr})
with pytest.raises(ValueError, match="both a dim and a sibling field"):
_validate_shape_element("ncvert", arr, _pkg("test"), enc, {"ncvert"})


def _cross_component_lookup_ctx():
"""An array in another component sized by a column of gwf-lak's packagedata."""
_arr, _enc, lak, _known = _lookup_ctx()
gwf = Model(name="gwf-nam", blocks=None)
spec = Dfns(components={"gwf-nam": gwf, "gwf-lak": lak})
fk_lakeno = Integer(name="lakeno", fk="packagedata.lakeno")
arr = Array(name="outflow", dtype="double", shape=[])
enc = Record(name="item", fields={"lakeno": fk_lakeno, "outflow": arr})
other = Package(name="gwf-other", parent="gwf-nam", blocks=None)
return arr, enc, other, spec


def test_validate_shape_element_cross_component_lookup():
arr, enc, other, spec = _cross_component_lookup_ctx()
elem = "gwf-lak.packagedata.nlakeconn(lakeno)"
_validate_shape_element(elem, arr, other, enc, set(), spec)


def test_validate_shape_element_cross_component_lookup_requires_spec():
arr, enc, other, _spec = _cross_component_lookup_ctx()
with pytest.raises(ValueError, match="requires a Dfns spec"):
_validate_shape_element("gwf-lak.packagedata.nlakeconn(lakeno)", arr, other, enc, set())


def test_validate_shape_element_cross_component_lookup_unknown_component():
arr, enc, other, spec = _cross_component_lookup_ctx()
with pytest.raises(ValueError, match="not found in spec"):
_validate_shape_element(
"gwf-nope.packagedata.nlakeconn(lakeno)", arr, other, enc, set(), spec
)


def test_validate_list_shape_element_lookup():
lst = List(name="stress_period_data", item=Record(name="item", fields={}))
with pytest.raises(ValueError, match="not inside a record"):
_validate_list_shape_element("packagedata.ncon(ifno)", lst, {"maxbound"})


def test_local_dims():
block = _dim_block("nlay", "nrow", "ncol")
pkg = Package(
Expand Down Expand Up @@ -981,3 +1097,248 @@ def test_input_dims_visible_to_model_component_itself():
gwf = Model(name="gwf-nam", parent="sim-nam", blocks=None)
spec = Dfns(components={"gwf-nam": gwf, "gwf-dis": dis})
assert "nodesuser" in spec.input_dims("gwf-nam")


def test_input_dim_must_name_its_own_field():
pkg = Package(
name="test",
blocks={"dimensions": _dim_block("nlay")},
dims={"nlayers": InputDim(value="nlay", scope="component")},
)
with pytest.raises(ValueError, match="must be the dim's own name"):
Dfns(components={"test": pkg})


_DIS_DIMS = {
"nlay": "nlay",
"nrow": "nrow",
"ncol": "ncol",
"ncpl": "nrow * ncol",
"nodes": "nlay * ncpl",
"ncelldim": "3",
"naux": "len(auxiliary)",
"nconn": "sum(packagedata.nlakeconn)",
"njas": "(nja - nodes) / 2",
}
_DIS_INPUTS = {
"nlay": 2,
"nrow": 3,
"ncol": 4,
"nja": 30,
"auxiliary": ["temp", "conc"],
"packagedata.nlakeconn": [1, 2, 3],
}


@pytest.mark.parametrize(
"name,expected",
[
("nlay", 2),
("ncpl", 12),
("nodes", 24),
("ncelldim", 3),
("naux", 2),
("nconn", 6),
("njas", 3),
("nja", 30), # not a dim: an input field
],
)
def test_dim_value(name, expected):
assert dim_value(name, _DIS_DIMS, _DIS_INPUTS.get) == expected


@pytest.mark.parametrize("name", ["nlay", "nodes", "naux", "nconn"])
def test_dim_value_unset_input(name):
assert dim_value(name, _DIS_DIMS, {}.get) is None


def test_dim_value_inexact_division():
with pytest.raises(ValueError, match="not an integer"):
dim_value("njas", _DIS_DIMS, {**_DIS_INPUTS, "nja": 31}.get)


def test_dim_value_cycle():
with pytest.raises(ValueError, match="cycle"):
dim_value("a", {"a": "b + 1", "b": "a * 2"}, {}.get)


@pytest.mark.parametrize("expr", ["nlay ** 2", "max(nlay, 1)", "nlay if nlay else 1", "1.5"])
def test_dim_value_unsupported(expr):
with pytest.raises(ValueError, match="unsupported"):
dim_value("d", {"d": expr}, {"nlay": 2}.get)


class _One(int):
"""1, usable as an Integer field, a self-sizing array, or a list column."""

def __new__(cls):
return super().__new__(cls, 1)

def __len__(self):
return 1

def __iter__(self):
return iter([1])


def test_dim_value_snapshot_dims():
"""Every input dim in the current DFNs evaluates, given its inputs."""
spec = Dfns.load(_DEV3_SNAPSHOT_DIR)
dis = spec.components["gwf-dis"]
dims = {n: d.value for n, d in dis.dims.items()}
assert dim_value("nodes", dims, {"nlay": 2, "nrow": 3, "ncol": 4}.get) == 24
for component in spec.components.values():
dims = {n: d.value for n, d in (component.dims or {}).items()}
for name in dims:
assert dim_value(name, dims, lambda _: _One()) is not None


def test_row_varying_array_must_be_rightmost():
ncvert = Integer(name="ncvert")
icvert = Array(name="icvert", dtype="integer", shape=["ncvert"])
xc = Double(name="xc")
enc = Record(name="item", fields={"ncvert": ncvert, "icvert": icvert, "xc": xc})
with pytest.raises(ValueError, match="must be the rightmost field"):
_validate_shape_element("ncvert", icvert, _pkg("test"), enc, set())


def test_row_varying_lookup_must_be_rightmost():
arr, enc, pkg, known = _lookup_ctx()
enc = Record(name="item", fields={**enc.fields, "after": Double(name="after")})
with pytest.raises(ValueError, match="must be the rightmost field"):
_validate_shape_element("packagedata.nlakeconn(lakeno)", arr, pkg, enc, known)


_WEL_DIMS = {"maxbound": "maxbound", "auxiliary": "len(auxiliary)", "nseg": "nseg"}


@pytest.mark.parametrize(
"expr,expected",
[
("nseg-1", 3), # a dim, offset
("auxiliary", 2), # a derived dim
("ncvert", 5), # a field in the row
],
)
def test_dim_value_shape_expression(expr, expected):
package = {"nseg": 4, "auxiliary": ["temp", "conc"]}
row = {"ncvert": 5}
assert dim_value(expr, _WEL_DIMS, ChainMap(row, package).get) == expected


def test_dim_value_row_lookup():
packagedata = {1: {"ncon": 2}, 2: {"ncon": 3}}
calls = []

def select(path, key):
calls.append((path, key))
return packagedata[key]["ncon"]

assert dim_value("packagedata.ncon(ifno)", {}, {"ifno": 2}.get, select) == 3
assert dim_value("gwf-x.packagedata.ncon(ifno)", {}, {"ifno": 1}.get, select) == 2
assert calls == [("packagedata.ncon", 2), ("gwf-x.packagedata.ncon", 1)]
assert dim_value("packagedata.ncon(ifno)", {}, {}.get, select) is None


def test_dim_value_row_lookup_needs_select():
with pytest.raises(ValueError, match="needs select"):
dim_value("packagedata.ncon(ifno)", {}, {"ifno": 1}.get)


def test_dim_value_bounded():
with pytest.raises(ValueError, match="split_bound"):
dim_value("<=maxbound", _WEL_DIMS, {"maxbound": 3}.get)


def test_dim_value_snapshot_shapes():
"""Every array and list shape in the current DFNs evaluates, given its inputs."""
spec = Dfns.load(_DEV3_SNAPSHOT_DIR)

def shapes(field):
if isinstance(field, (Array, List)):
yield from field.shape
if isinstance(field, List):
yield from shapes(field.item)
children = getattr(field, "fields", None) or getattr(field, "arms", None) or {}
for child in children.values():
yield from shapes(child)

for component in spec.components.values():
dims = {n: d.value for n, d in (component.dims or {}).items()}
for block in (component.blocks or {}).values():
for field in block.fields.values():
for element in shapes(field):
_bound, expr = split_bound(element)
value = dim_value(expr, dims, lambda _: _One(), lambda _p, _k: 1)
assert value is not None, (component.name, element)


@pytest.mark.parametrize(
"expr,inputs,length,expected",
[
("ncvert", {}, 5, ("ncvert", 5)), # a field in the row
("numalphaj", {}, 3, ("numalphaj", 3)), # an input dim
("nseg-1", {}, 3, ("nseg", 4)),
("1 + nseg", {}, 3, ("nseg", 2)),
("10 - nseg", {}, 3, ("nseg", 7)),
("-nseg", {}, -3, ("nseg", 3)),
("nlay", {}, 2, ("nlay", 2)),
("nlayp", {}, 3, ("nlay", 2)), # through a derived dim
("nseg-1", {"nseg": 4}, 3, None), # nothing unset
("auxiliary", {}, 2, None), # len() can't be undone
("nconn", {}, 6, None), # nor sum()
("ncpl", {"nrow": 3}, 12, None), # nor *
("nlay + nseg", {}, 5, None), # two unset inputs
("packagedata.ncon(ifno)", {"ifno": 1}, 2, None), # nor a row-level lookup
],
)
def test_dim_input(expr, inputs, length, expected):
dims = {
"numalphaj": "numalphaj",
"nseg": "nseg",
"nlay": "nlay",
"nlayp": "nlay + 1",
"ncpl": "nrow * ncol",
"auxiliary": "len(auxiliary)",
"nconn": "sum(packagedata.nlakeconn)",
}
assert dim_input(expr, dims, inputs.get, length=length) == expected


def test_dim_input_snapshot_shapes():
"""Solving any shape in the current DFNs for its sole unset input, then
evaluating it with that input set, gives back the length."""
spec = Dfns.load(_DEV3_SNAPSHOT_DIR)
solved = set()
for component in spec.components.values():
dims = {n: d.value for n, d in (component.dims or {}).items()}

def shapes(field):
if isinstance(field, (Array, List)):
yield from field.shape
if isinstance(field, List):
yield from shapes(field.item)
children = getattr(field, "fields", None) or getattr(field, "arms", None) or {}
for child in children.values():
yield from shapes(child)

for block in (component.blocks or {}).values():
for field in block.fields.values():
for element in shapes(field):
_bound, expr = split_bound(element)
if (solution := dim_input(expr, dims, length=7)) is None:
continue
name, n = solution
assert dim_value(expr, dims, {name: n}.get) == 7, (component.name, expr)
solved.add(expr)
assert {"ncvert", "numalphaj", "nseg-1", "maxbound"} <= solved
assert "auxiliary" not in solved


def test_dim_defaults():
"""Without dims there are none; without a lookup no input is set."""
assert dim_value("2 * 3") == 6
assert dim_value("nlay") is None
assert dim_value("nlayp", {"nlayp": "nlay + 1"}) is None
assert dim_input("ncvert", length=5) == ("ncvert", 5)
assert dim_input("nlayp", {"nlayp": "nlay + 1"}, length=3) == ("nlay", 2)
Loading
Loading