Skip to content
Merged
Show file tree
Hide file tree
Changes from 5 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
102 changes: 55 additions & 47 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] = {}
Comment thread
puddly marked this conversation as resolved.
Outdated
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,74 @@ 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 != 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",
Comment thread
puddly marked this conversation as resolved.
Outdated
idx,
entry.multicastId,
entry.endpoint,
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
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
Comment thread
puddly marked this conversation as resolved.
Outdated
) -> 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
Comment thread
puddly marked this conversation as resolved.

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
Comment thread
puddly marked this conversation as resolved.

# 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)
Comment thread
puddly marked this conversation as resolved.

# The multicast table was never touched
assert len(app._ezsp._protocol.setMulticastTableEntry.mock_calls) == 0
12 changes: 6 additions & 6 deletions tests/test_multicast.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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


Expand Down Expand Up @@ -167,15 +167,15 @@ 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()
ret = await _unsubscribe(multicast, grp_id, success=True)
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


Expand All @@ -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
Loading