Skip to content
Merged
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
109 changes: 61 additions & 48 deletions bellows/multicast.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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)

Expand All @@ -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
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: 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()
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
29 changes: 29 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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)
44 changes: 44 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,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
Comment thread
puddly marked this conversation as resolved.
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
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