diff --git a/autotest/dfns/test_schema.py b/autotest/dfns/test_schema.py index fa389584..279c5681 100644 --- a/autotest/dfns/test_schema.py +++ b/autotest/dfns/test_schema.py @@ -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() diff --git a/modflow_devtools/dfns/schema.py b/modflow_devtools/dfns/schema.py index 561d3a0f..0d446cd4 100644 --- a/modflow_devtools/dfns/schema.py +++ b/modflow_devtools/dfns/schema.py @@ -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[ @@ -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 "" @@ -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] @@ -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: