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
90 changes: 46 additions & 44 deletions bellows/multicast.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ class Multicast:

def __init__(self, ezsp):
self._ezsp = ezsp
self._multicast = {}
self._multicast: dict[int, int] = {}
self._available = set()

async def _initialize(self) -> None:
Expand All @@ -30,7 +30,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)

Expand All @@ -42,66 +42,68 @@ 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:
LOGGER.warning(
"Set MulticastTableEntry #%s for %s multicast id: %s",
idx,
entry.multicastId,
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,
else:
LOGGER.debug(
"Set MulticastTableEntry #%s for %s multicast id %s for endpoint %d: %s",
idx,
entry.multicastId,
entry.endpoint,
status,
)

return status[0], 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:
self._multicast[entry.multicastId, entry.endpoint] = (entry, 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]
status, _entry = await self._set_multicast_entry(
idx=idx,
group_id=group_id,
endpoint_id=0,
)

self._multicast.pop(group_id)
self._multicast.pop((group_id, endpoint_id))
self._available.add(idx)
LOGGER.debug(
"Set MulticastTableEntry #%s for %s multicast id: %s",
idx,
entry.multicastId,
status,
)
return status[0]

return status
18 changes: 18 additions & 0 deletions bellows/zigbee/application.py
Original file line number Diff line number Diff line change
Expand Up @@ -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: t.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: t.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()
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ dependencies = [
"click",
"click-log>=0.2.1",
"voluptuous",
"zigpy>=0.87.0",
"zigpy>=2.1.0",
]

[tool.setuptools.packages.find]
Expand Down
43 changes: 43 additions & 0 deletions tests/test_application.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
GetRouteTableEntryRsp,
GetTxPowerInfoRsp,
)
from bellows.multicast import Multicast
import bellows.types
import bellows.types as t
import bellows.types.struct
Expand Down Expand Up @@ -2683,3 +2684,45 @@ 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

# 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
Loading