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
53 changes: 53 additions & 0 deletions autotest/dfns/test_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -379,6 +379,59 @@ def test_get_fields_and_get_block_include_header(dev3_spec):
assert wel.get_block("iper") is wel.blocks["period"]


def test_block_get_fields_no_recurse():
record = Record(
name="afrcsv_filerecord",
fields={
"auto_flow_reduce_csv": Keyword(name="auto_flow_reduce_csv"),
"afrcsvfile": File(name="afrcsvfile", direction="out"),
},
)
block = Block(name="options", fields={"afrcsv_filerecord": record})
fields = block.get_fields()
assert fields.keys() == ["afrcsv_filerecord"]
assert fields["afrcsv_filerecord"] is record


def test_block_get_fields_recurse_descends_record():
inner = File(name="afrcsvfile", direction="out")
record = Record(
name="afrcsv_filerecord",
fields={"auto_flow_reduce_csv": Keyword(name="auto_flow_reduce_csv"), "afrcsvfile": inner},
)
block = Block(name="options", fields={"afrcsv_filerecord": record})
fields = block.get_fields(recurse=True)
assert set(fields.keys()) == {"afrcsv_filerecord", "auto_flow_reduce_csv", "afrcsvfile"}
assert fields["afrcsvfile"] is inner


def test_block_get_fields_recurse_descends_list_item_union():
keyword_arm = Keyword(name="save")
record_arm = Record(name="print", fields={"rtype": String(name="rtype")})
union = Union(name="ocsetting", arms={"save": keyword_arm, "print": record_arm})
lst = List(name="steps", item=union)
block = Block(name="period", fields={"steps": lst})
fields = block.get_fields(recurse=True)
assert set(fields.keys()) == {"steps", "save", "print", "rtype"}


def test_block_get_fields_recurse_includes_header():
header = Integer(name="iper", tagged=False)
block = Block(name="period", fields={}, header=header)
fields = block.get_fields(recurse=True)
assert fields["iper"] is header


def test_component_get_fields_matches_block_get_fields(dev3_spec):
"""ComponentBase.get_fields() is a flatten of each block's get_fields()."""
wel = dev3_spec.components["gwf-wel"]
expected: list[tuple] = []
for block in wel.blocks.values():
expected.extend(block.get_fields(recurse=True).items(multi=True))
actual = wel.get_fields(recurse=True)
assert list(actual.items(multi=True)) == expected


def test_render_block_header_scalar(dev3_spec):
"""render() attaches a scalar header to the BEGIN line, matching mf6io.pdf."""
render = dev3_spec.components["gwf-wel"].blocks["period"].render()
Expand Down
35 changes: 21 additions & 14 deletions modflow_devtools/dfns/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -169,7 +169,8 @@ def _check_shape_length(self) -> "List":

@property
def children(self) -> "dict[str, Field]":
return {"item": self.item} # type: ignore[return-value]
# item.name duplicates the List's own name, so it's not a useful key here.
return self.item.children


Field = Annotated[
Expand All @@ -183,6 +184,16 @@ def children(self) -> "dict[str, Field]":
List.model_rebuild()


def _collect_fields(
fields: "dict[str, Field]", items: "list[tuple[str, Field]]", *, recurse: bool
) -> None:
"""Append `fields` to `items`, descending into Record/Union/List children if `recurse`."""
for name, field in fields.items():
items.append((name, field))
if recurse and isinstance(field, (Record, Union, List)):
_collect_fields(field.children, items, recurse=True)


def _render_shape(field: "Array") -> str:
return f"({', '.join(field.shape)})" if field.shape else ""

Expand Down Expand Up @@ -586,6 +597,14 @@ def repeats(self) -> bool:
def render(self, *, developmode: bool = False) -> str:
return _render_block(self, developmode=developmode)

def get_fields(self, recurse: bool = False) -> OMD:
"""Fields keyed by name, including `header`; `recurse` descends into children."""
items: list[tuple[str, Field]] = []
_collect_fields(self.fields, items, recurse=recurse)
if self.header is not None:
_collect_fields({self.header.name: self.header}, items, recurse=recurse)
return OMD(items)


Blocks = Mapping[str, Block]

Expand Down Expand Up @@ -651,20 +670,8 @@ def _serialize(self, handler: Any) -> dict[str, Any]:

def get_fields(self, recurse: bool = False) -> OMD:
items: list[tuple[str, Field]] = []

def _collect(fields: dict) -> None:
for name, field in fields.items():
items.append((name, field))
if recurse:
if isinstance(field, (Record, Union)):
_collect(field.children)
elif isinstance(field, List):
_collect(field.item.children)

for block in (self.blocks or {}).values():
_collect(block.fields)
if block.header is not None:
_collect({block.header.name: block.header})
items.extend(block.get_fields(recurse=recurse).items(multi=True))
return OMD(items)

def get_block(self, field_name: str) -> Block | None:
Expand Down
Loading