From 638ac58127c4a2df78bf07807805c749bbc7ebd5 Mon Sep 17 00:00:00 2001 From: CAOShurong <170531907+CAOShurong@users.noreply.github.com> Date: Thu, 17 Sep 2026 02:44:03 +0800 Subject: [PATCH] Fix NX class membership checks Signed-off-by: CAOShurong <170531907+CAOShurong@users.noreply.github.com> --- src/scippnexus/base.py | 7 +++++-- tests/nexus_test.py | 25 ++++++++++++++++++++++--- 2 files changed, 27 insertions(+), 5 deletions(-) diff --git a/src/scippnexus/base.py b/src/scippnexus/base.py index a6c16ca0..2e1c43fc 100644 --- a/src/scippnexus/base.py +++ b/src/scippnexus/base.py @@ -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 diff --git a/tests/nexus_test.py b/tests/nexus_test.py index 110dabec..efb0afa3 100644 --- a/tests/nexus_test.py +++ b/tests/nexus_test.py @@ -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'}