Skip to content
Open
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
7 changes: 5 additions & 2 deletions src/scippnexus/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -308,11 +308,14 @@ def _get_children_by_nx_class(
self, select: type | list[type]
) -> dict[str, NXobject | Field]:
children = {}
select = (select,) if isinstance(select, type) else select
requested = select
selectors = (select,) if isinstance(select, type) else select
for key, child in self._children.items():
nx_class = Field if isinstance(child, Field) else child.nx_class
if nx_class is not None and any(nx_class == sel for sel in select):
if nx_class is not None and any(nx_class == sel for sel in selectors):
children[key] = self[key]
if not children:
raise KeyError(requested)
return children

@overload
Expand Down
25 changes: 22 additions & 3 deletions tests/nexus_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,18 +133,37 @@ def test_nxobject_getitem_by_class(nxroot) -> None:
nxroot['entry'].create_class('events_1', NXevent_data)
assert list(nxroot[NXentry]) == ['entry']
assert list(nxroot[NXmonitor]) == ['monitor']
assert list(nxroot['entry'][NXmonitor]) == [] # not nested
assert list(nxroot[NXlog]) == [] # nested
with pytest.raises(KeyError, match='NXmonitor') as error:
nxroot['entry'][NXmonitor] # not nested
assert error.value.args == (NXmonitor,)
with pytest.raises(KeyError, match='NXlog'):
nxroot[NXlog] # nested
assert list(nxroot['entry'][NXlog]) == ['log']
assert set(nxroot['entry'][NXevent_data]) == {'events_0', 'events_1'}


def test_nxobject_contains_by_class(nxroot) -> None:
nxroot.create_class('monitor', NXmonitor)

assert NXentry in nxroot
assert NXmonitor in nxroot
assert NXlog not in nxroot


def test_nxobject_get_by_class_returns_default_when_absent(nxroot) -> None:
default = object()

assert nxroot.get(NXmonitor) is None
assert nxroot.get(NXmonitor, default) is default


def test_nxobject_getitem_by_class_get_fields(nxroot) -> None:
nxroot['entry'].create_class('log', NXlog)
nxroot['entry'].create_class('events_0', NXevent_data)
nxroot['entry']['field1'] = sc.arange('event', 4.0, unit='ns')
nxroot['entry']['field2'] = sc.arange('event', 2.0, unit='ns')
assert list(nxroot[snx.Field]) == []
with pytest.raises(KeyError, match='Field'):
nxroot[snx.Field]
assert set(nxroot['entry'][snx.Field]) == {'field1', 'field2'}


Expand Down