diff --git a/bellows/multicast.py b/bellows/multicast.py index 1e7955d0..a1fe3634 100644 --- a/bellows/multicast.py +++ b/bellows/multicast.py @@ -10,7 +10,9 @@ class Multicast: def __init__(self, ezsp): self._ezsp = ezsp - self._multicast = {} + self._multicast: dict[ + tuple[int, int], tuple[t.EmberMulticastTableEntry, int] + ] = {} self._available = set() async def _initialize(self) -> None: @@ -30,7 +32,7 @@ async def _initialize(self) -> None: continue LOGGER.debug("MulticastTableEntry[%s] = %s", i, entry) if entry.endpoint != 0: - self._multicast[entry.multicastId] = (entry, i) + self._multicast[entry.multicastId, entry.endpoint] = (entry, i) else: self._available.add(i) @@ -42,66 +44,77 @@ async def startup(self, coordinator) -> None: for group_id in ep.member_of: await self.subscribe(group_id) - async def subscribe(self, group_id) -> t.sl_Status: - if group_id in self._multicast: - LOGGER.debug("%s is already subscribed", t.EmberMulticastId(group_id)) - return t.sl_Status.OK - - try: - idx = self._available.pop() - except KeyError: - LOGGER.error("No more available slots MulticastId subscription") - return t.sl_Status.INVALID_INDEX + async def _set_multicast_entry( + self, idx: int, group_id: int, endpoint_id: int + ) -> tuple[t.sl_Status, t.EmberMulticastTableEntry]: entry = t.EmberMulticastTableEntry() - entry.endpoint = t.uint8_t(1) + entry.endpoint = t.uint8_t(endpoint_id) entry.multicastId = t.EmberMulticastId(group_id) entry.networkIndex = t.uint8_t(0) - status = await self._ezsp.setMulticastTableEntry(idx, entry) - if t.sl_Status.from_ember_status(status[0]) != t.sl_Status.OK: + + (status,) = await self._ezsp.setMulticastTableEntry(idx, entry) + status = t.sl_Status.from_ember_status(status) + + if status is t.sl_Status.OK: + LOGGER.debug( + "Set MulticastTableEntry #%s for %s multicast id %s for endpoint %d: %s", + idx, + group_id, + entry.multicastId, + entry.endpoint, + status, + ) + else: LOGGER.warning( - "Set MulticastTableEntry #%s for %s multicast id: %s", + "Failed to set MulticastTableEntry #%s for %s multicast id %s for endpoint %d: %s", idx, + group_id, entry.multicastId, + entry.endpoint, status, ) - self._available.add(idx) - return status[0] - - self._multicast[entry.multicastId] = (entry, idx) - LOGGER.debug( - "Set MulticastTableEntry #%s for %s multicast id: %s", - idx, - entry.multicastId, - status, + + return status, entry + + async def subscribe(self, group_id: int, endpoint_id: int = 1) -> t.sl_Status: + if (group_id, endpoint_id) in self._multicast: + LOGGER.debug("%s is already subscribed", t.EmberMulticastId(group_id)) + return t.sl_Status.OK + + try: + idx = self._available.pop() + except KeyError: + LOGGER.error("No more available slots MulticastId subscription") + return t.sl_Status.INVALID_INDEX + + status, entry = await self._set_multicast_entry( + idx=idx, group_id=group_id, endpoint_id=endpoint_id ) - return status[0] - async def unsubscribe(self, group_id) -> t.sl_Status: + if status is t.sl_Status.OK: + self._multicast[entry.multicastId, entry.endpoint] = (entry, idx) + else: + self._available.add(idx) + + return status + + async def unsubscribe(self, group_id: int, endpoint_id: int = 1) -> t.sl_Status: try: - entry, idx = self._multicast[group_id] + _entry, idx = self._multicast[group_id, endpoint_id] except KeyError: - LOGGER.error( + LOGGER.debug( "Couldn't find MulticastTableEntry for %s multicast_id", group_id ) return t.sl_Status.INVALID_INDEX - entry.endpoint = t.uint8_t(0) - status = await self._ezsp.setMulticastTableEntry(idx, entry) - if t.sl_Status.from_ember_status(status[0]) != t.sl_Status.OK: - LOGGER.warning( - "Set MulticastTableEntry #%s for %s multicast id: %s", - idx, - entry.multicastId, - status, - ) - return status[0] - - self._multicast.pop(group_id) - self._available.add(idx) - LOGGER.debug( - "Set MulticastTableEntry #%s for %s multicast id: %s", - idx, - entry.multicastId, - status, + status, _entry = await self._set_multicast_entry( + idx=idx, + group_id=group_id, + endpoint_id=0, ) - return status[0] + + if status is t.sl_Status.OK: + self._multicast.pop((group_id, endpoint_id)) + self._available.add(idx) + + return status diff --git a/bellows/zigbee/application.py b/bellows/zigbee/application.py index f64254ac..c8870b86 100644 --- a/bellows/zigbee/application.py +++ b/bellows/zigbee/application.py @@ -1131,6 +1131,24 @@ async def permit_with_link_key( return await super().permit(time_s) + async def _subscribe_to_multicast_group( + self, group_id: zigpy.types.Group, endpoint_id: int + ) -> None: + """Ask the coordinator firmware to subscribe to a group, if needed.""" + if self._multicast is None: + return None + + await self._multicast.subscribe(group_id=group_id, endpoint_id=endpoint_id) + + async def _unsubscribe_from_multicast_group( + self, group_id: zigpy.types.Group, endpoint_id: int + ) -> None: + """Ask the coordinator firmware to unsubscribe from a group, if needed.""" + if self._multicast is None: + return None + + await self._multicast.unsubscribe(group_id=group_id, endpoint_id=endpoint_id) + def _handle_id_conflict(self, nwk: t.EmberNodeId) -> None: LOGGER.warning("NWK conflict is reported for 0x%04x", nwk) self.state.counters[COUNTERS_CTRL][COUNTER_NWK_CONFLICTS].increment() diff --git a/pyproject.toml b/pyproject.toml index 98db341c..84074acd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,7 +17,7 @@ dependencies = [ "click", "click-log>=0.2.1", "voluptuous", - "zigpy>=0.87.0", + "zigpy>=2.1.0", ] [tool.setuptools.packages.find] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..7f08caf6 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,29 @@ +"""Common pytest fixtures for all tests.""" + +import logging + +import pytest + + +class FailOnBadFormattingHandler(logging.Handler): + def emit(self, record): + try: + record.msg % record.args + except Exception as e: # noqa: BLE001 + pytest.fail( + f"Failed to format log message {record.msg!r} with {record.args!r}: {e}" + ) + + +@pytest.fixture(autouse=True) +def raise_on_bad_log_formatting(): + handler = FailOnBadFormattingHandler() + + root = logging.getLogger() + root.addHandler(handler) + root.setLevel(logging.DEBUG) + + try: + yield + finally: + root.removeHandler(handler) diff --git a/tests/test_application.py b/tests/test_application.py index 9170747b..890ab2ab 100644 --- a/tests/test_application.py +++ b/tests/test_application.py @@ -24,6 +24,7 @@ GetRouteTableEntryRsp, GetTxPowerInfoRsp, ) +from bellows.multicast import Multicast import bellows.types import bellows.types as t import bellows.types.struct @@ -2683,3 +2684,46 @@ async def test_set_tx_power(app: ControllerApplication) -> None: assert result == 12.0 assert app._ezsp.setRadioPower.mock_calls == [call(power=12)] assert mock_update.mock_calls == [call(app._ezsp, tx_power=12)] + + +async def test_multicast_group_subscription(app: ControllerApplication) -> None: + """Test multicast group subscription APIs when there are no XNCP extensions.""" + app._ezsp._xncp_features = FirmwareFeatures.NONE + + app._multicast = Multicast(app._ezsp) + await app._multicast._initialize() + + # Subscribe to a group + await app.subscribe_to_multicast_group(0x1234) + assert app._ezsp._protocol.setMulticastTableEntry.mock_calls == [ + call( + 0, + t.EmberMulticastTableEntry(multicastId=0x1234, endpoint=1, networkIndex=0), + ) + ] + + app._ezsp._protocol.setMulticastTableEntry.reset_mock() + + # Unsubscribe from a group + await app.unsubscribe_from_multicast_group(0x1234) + assert app._ezsp._protocol.setMulticastTableEntry.mock_calls == [ + call( + 0, + t.EmberMulticastTableEntry(multicastId=0x1234, endpoint=0, networkIndex=0), + ) + ] + + +async def test_multicast_group_subscription_xncp(app: ControllerApplication) -> None: + """Test multicast group subscription APIs when XNCP extensions are available.""" + app._ezsp._xncp_features |= FirmwareFeatures.MEMBER_OF_ALL_GROUPS + assert app._multicast is None + + # Subscribe to a group (no-op) + await app.subscribe_to_multicast_group(0x1234) + + # Unsubscribe from a group (no-op) + await app.unsubscribe_from_multicast_group(0x1234) + + # The multicast table was never touched + assert len(app._ezsp._protocol.setMulticastTableEntry.mock_calls) == 0 diff --git a/tests/test_multicast.py b/tests/test_multicast.py index c08e0685..ef44bf6d 100644 --- a/tests/test_multicast.py +++ b/tests/test_multicast.py @@ -115,14 +115,14 @@ async def test_subscribe(multicast): set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 1 assert set_entry.call_args[0][1].multicastId == grp_id - assert grp_id in multicast._multicast + assert (grp_id, 1) in multicast._multicast set_entry.reset_mock() ret = await _subscribe(multicast, grp_id, success=True) assert ret == t.EmberStatus.SUCCESS set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 0 - assert grp_id in multicast._multicast + assert (grp_id, 1) in multicast._multicast async def test_subscribe_fail(multicast): @@ -134,7 +134,7 @@ async def test_subscribe_fail(multicast): set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 1 assert set_entry.call_args[0][1].multicastId == grp_id - assert grp_id not in multicast._multicast + assert (grp_id, 1) not in multicast._multicast assert len(multicast._available) == 1 @@ -167,7 +167,7 @@ async def test_unsubscribe(multicast): assert ret == t.EmberStatus.SUCCESS set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 1 - assert grp_id not in multicast._multicast + assert (grp_id, 1) not in multicast._multicast assert len(multicast._available) == 1 multicast._ezsp.setMulticastTableEntry.reset_mock() @@ -175,7 +175,7 @@ async def test_unsubscribe(multicast): assert ret != t.EmberStatus.SUCCESS set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 0 - assert grp_id not in multicast._multicast + assert (grp_id, 1) not in multicast._multicast assert len(multicast._available) == 1 @@ -190,5 +190,5 @@ async def test_unsubscribe_fail(multicast): assert ret != t.EmberStatus.SUCCESS set_entry = multicast._ezsp.setMulticastTableEntry assert set_entry.call_count == 1 - assert grp_id in multicast._multicast + assert (grp_id, 1) in multicast._multicast assert len(multicast._available) == 0