diff --git a/homeassistant/components/alexa_devices/coordinator.py b/homeassistant/components/alexa_devices/coordinator.py index 3e7b23c19a04..7874736bdbd0 100644 --- a/homeassistant/components/alexa_devices/coordinator.py +++ b/homeassistant/components/alexa_devices/coordinator.py @@ -1,7 +1,7 @@ """Support for Alexa Devices.""" from asyncio import Lock -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager from datetime import timedelta from pathlib import Path @@ -184,6 +184,10 @@ class AmazonDevicesCoordinator(DataUpdateCoordinator[dict[str, AmazonDevice]]): self.api.on_media_state_event.append(self.media_state_event_handler) self.api.on_media_state_event.freeze() + self._dnd_states: dict[str, bool] = {} + self.api.on_dnd_event.append(self.dnd_event_handler) + self.api.on_dnd_event.freeze() + @override async def _async_update_data(self) -> dict[str, AmazonDevice]: """Update device data.""" @@ -246,7 +250,7 @@ class AmazonDevicesCoordinator(DataUpdateCoordinator[dict[str, AmazonDevice]]): async def _async_sync_on_device_list_change(self) -> None: """Sync per-device state on first refresh and after the device list changes.""" - for sync_call in (self.sync_media_state,): + for sync_call in (self.sync_dnd_state, self.sync_media_state): try: await sync_call() except ConfigEntryNotReady as err: @@ -420,3 +424,22 @@ class AmazonDevicesCoordinator(DataUpdateCoordinator[dict[str, AmazonDevice]]): def volume_states(self) -> dict[str, AmazonVolumeState]: """Volumes of devices.""" return self._volume_states + + async def sync_dnd_state(self) -> None: + """Sync dnd state.""" + async with alexa_config_entry_errors(): + await self.api.sync_dnd_state() + + async def dnd_event_handler(self, dnd_states: dict[str, bool]) -> None: + """Handle pushed dnd events.""" + self._dnd_states = dict(dnd_states) + self.async_update_listeners() + + def set_dnd_state(self, serial_num: str, state: bool) -> None: + """Set the local DND state; caller writes its own state, so listeners aren't notified.""" + self._dnd_states[serial_num] = state + + @property + def dnd_states(self) -> Mapping[str, bool]: + """DND states of devices.""" + return self._dnd_states diff --git a/homeassistant/components/alexa_devices/switch.py b/homeassistant/components/alexa_devices/switch.py index 60193a8d41cf..b608dfe2633a 100644 --- a/homeassistant/components/alexa_devices/switch.py +++ b/homeassistant/components/alexa_devices/switch.py @@ -5,7 +5,6 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Final, override from aioamazondevices.const.devices import SPEAKER_GROUP_FAMILY -from aioamazondevices.structures import AmazonDevice from homeassistant.components.switch import ( DOMAIN as SWITCH_DOMAIN, @@ -16,37 +15,59 @@ from homeassistant.const import EntityCategory from homeassistant.core import HomeAssistant from homeassistant.helpers.entity_platform import AddConfigEntryEntitiesCallback -from .coordinator import AmazonConfigEntry, alexa_api_call +from .coordinator import AmazonConfigEntry, AmazonDevicesCoordinator, alexa_api_call from .entity import AmazonEntity from .utils import async_remove_entities, async_update_unique_id PARALLEL_UPDATES = 1 -TYPE_SENSOR = "sensor" -TYPE_COMMUNICATION = "communication" + +def _communication_is_on( + coordinator: AmazonDevicesCoordinator, + serial_num: str, + entity_description_key: str, +) -> bool: + """Return the local communication settings state.""" + return ( + coordinator.data[serial_num].communication_settings[entity_description_key] + == "ON" + ) + + +def _update_communication_state( + coordinator: AmazonDevicesCoordinator, + serial_num: str, + entity_description_key: str, + state: bool, +) -> None: + """Update the local communication settings state.""" + coordinator.data[serial_num].communication_settings[entity_description_key] = ( + "ON" if state else "OFF" + ) @dataclass(frozen=True, kw_only=True) class AmazonSwitchEntityDescription(SwitchEntityDescription): """Alexa Devices switch entity description.""" - is_on_fn: Callable[[AmazonDevice], bool] - is_available_fn: Callable[[AmazonDevice, str], bool] = lambda device, key: ( - device.online - and (sensor := device.sensors.get(key)) is not None - and sensor.error is False - ) + is_on_fn: Callable[[AmazonDevicesCoordinator, str, str], bool] + is_available_fn: Callable[[AmazonDevicesCoordinator, str, str], bool] method: str - switch_type: str + update_state_fn: Callable[[AmazonDevicesCoordinator, str, str, bool], None] -SENSOR_SWITCHES: Final = ( - AmazonSwitchEntityDescription( - key="dnd", - translation_key="do_not_disturb", - is_on_fn=lambda device: bool(device.sensors["dnd"].value), - method="set_do_not_disturb", - switch_type=TYPE_SENSOR, +DND_SWITCH: Final = AmazonSwitchEntityDescription( + key="dnd", + translation_key="do_not_disturb", + is_on_fn=lambda coordinator, serial_num, _: coordinator.dnd_states.get( + serial_num, False + ), + is_available_fn=lambda coordinator, serial_num, _: ( + serial_num in coordinator.dnd_states + ), + method="set_do_not_disturb", + update_state_fn=lambda coordinator, serial_num, _, state: coordinator.set_dnd_state( + serial_num, state ), ) COMMUNICATION_SWITCHES: Final = ( @@ -54,25 +75,25 @@ COMMUNICATION_SWITCHES: Final = ( key="announcements", translation_key="announcements", entity_category=EntityCategory.CONFIG, - is_on_fn=lambda device: device.communication_settings["announcements"] == "ON", - is_available_fn=lambda device, key: ( - device.online - and device.communication_settings.get(key) is not None - and device.communication_settings.get("communications") != "OFF" + is_on_fn=_communication_is_on, + is_available_fn=lambda coordinator, serial_num, key: ( + (settings := coordinator.data[serial_num].communication_settings).get(key) + is not None + and settings.get("communications") != "OFF" ), method="set_announcement_status", - switch_type=TYPE_COMMUNICATION, + update_state_fn=_update_communication_state, ), AmazonSwitchEntityDescription( key="communications", translation_key="communications", entity_category=EntityCategory.CONFIG, - is_on_fn=lambda device: device.communication_settings["communications"] == "ON", - is_available_fn=lambda device, key: ( - device.online and device.communication_settings.get(key) is not None + is_on_fn=_communication_is_on, + is_available_fn=lambda coordinator, serial_num, key: ( + coordinator.data[serial_num].communication_settings.get(key) is not None ), method="set_communication_status", - switch_type=TYPE_COMMUNICATION, + update_state_fn=_update_communication_state, ), ) @@ -103,27 +124,35 @@ async def async_setup_entry( await async_update_unique_id(hass, coordinator, SWITCH_DOMAIN, old_key, new_key) known_devices: set[str] = set() + known_dnd_devices: set[str] = set() def _check_device() -> None: current_devices = set(coordinator.data) known_devices.intersection_update(current_devices) new_devices = current_devices - known_devices - if new_devices: - known_devices.update(new_devices) - sensor_switches = [ - AmazonSwitchEntity(coordinator, serial_num, switch_desc) - for switch_desc in SENSOR_SWITCHES - for serial_num in new_devices - if switch_desc.key in coordinator.data[serial_num].sensors + + # DND state may arrive after device discovery (initial sync failure, + # or a later push), so track it separately from `known_devices`. + known_dnd_devices.intersection_update(current_devices) + new_dnd_devices = ( + current_devices & coordinator.dnd_states.keys() + ) - known_dnd_devices + + async_add_entities( + [ + AmazonSwitchEntity(coordinator, serial_num, DND_SWITCH) + for serial_num in new_dnd_devices ] - communication_switches = [ + + [ AmazonSwitchEntity(coordinator, serial_num, switch_desc) for switch_desc in COMMUNICATION_SWITCHES for serial_num in new_devices if switch_desc.key in coordinator.data[serial_num].communication_settings ] - async_add_entities(sensor_switches + communication_switches) + ) + known_dnd_devices.update(new_dnd_devices) + known_devices.update(new_devices) _check_device() entry.async_on_unload(coordinator.async_add_listener(_check_device)) @@ -143,14 +172,12 @@ class AmazonSwitchEntity(AmazonEntity, SwitchEntity): async with alexa_api_call(self.coordinator): await method(self.device, state) - if self.entity_description.switch_type == TYPE_SENSOR: - self.coordinator.data[self.device.serial_number].sensors[ - self.entity_description.key - ].value = state - elif self.entity_description.switch_type == TYPE_COMMUNICATION: - self.coordinator.data[self.device.serial_number].communication_settings[ - self.entity_description.key - ] = "ON" if state else "OFF" + self.entity_description.update_state_fn( + self.coordinator, + self.device.serial_number, + self.entity_description.key, + state, + ) self.async_write_ha_state() @override @@ -167,15 +194,17 @@ class AmazonSwitchEntity(AmazonEntity, SwitchEntity): @override def is_on(self) -> bool: """Return True if switch is on.""" - return self.entity_description.is_on_fn(self.device) + + return self.entity_description.is_on_fn( + self.coordinator, + self.device.serial_number, + self.entity_description.key, + ) @property @override def available(self) -> bool: """Return if entity is available.""" - return ( - self.entity_description.is_available_fn( - self.device, self.entity_description.key - ) - and super().available + return super().available and self.entity_description.is_available_fn( + self.coordinator, self.device.serial_number, self.entity_description.key ) diff --git a/tests/components/alexa_devices/conftest.py b/tests/components/alexa_devices/conftest.py index 51f1e78ef7da..22e03d89c09a 100644 --- a/tests/components/alexa_devices/conftest.py +++ b/tests/components/alexa_devices/conftest.py @@ -1,7 +1,7 @@ """Alexa Devices tests configuration.""" import asyncio -from collections.abc import Generator +from collections.abc import Awaitable, Callable, Generator from copy import deepcopy from unittest.mock import AsyncMock, MagicMock, patch @@ -67,6 +67,17 @@ def mock_amazon_devices_client() -> Generator[AsyncMock]: client.on_volume_state_event = MagicMock() client.on_media_state_event = MagicMock() client.on_todo_event = MagicMock() + client.on_dnd_event = MagicMock() + dnd_event_handler: list[Callable[[dict[str, bool]], Awaitable[None]]] = [] + client.on_dnd_event.append.side_effect = dnd_event_handler.append + + async def _sync_dnd_state() -> None: + assert dnd_event_handler, "on_dnd_event handler was not registered" + await dnd_event_handler[0]( + dict.fromkeys(client.get_devices_data.return_value, False) + ) + + client.sync_dnd_state = AsyncMock(side_effect=_sync_dnd_state) async def _start_http2_processing(*_args, **_kwargs) -> asyncio.Task[None]: async def _completed_task() -> None: diff --git a/tests/components/alexa_devices/test_coordinator.py b/tests/components/alexa_devices/test_coordinator.py index b2a720d520f1..6e6cd3e0d9a7 100644 --- a/tests/components/alexa_devices/test_coordinator.py +++ b/tests/components/alexa_devices/test_coordinator.py @@ -61,6 +61,34 @@ async def test_coordinator_stale_device( assert not hass.states.get(entity_id_1) +async def test_coordinator_stale_device_clears_dnd_state( + hass: HomeAssistant, + freezer: FrozenDateTimeFactory, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test a removed device's DND state does not linger after resync.""" + mock_amazon_devices_client.get_devices_data.return_value = { + TEST_DEVICE_1_SN: TEST_DEVICE_1, + TEST_DEVICE_2_SN: TEST_DEVICE_2, + } + + await setup_integration(hass, mock_config_entry) + + coordinator = mock_config_entry.runtime_data + assert TEST_DEVICE_2_SN in coordinator.dnd_states + + mock_amazon_devices_client.get_devices_data.return_value = { + TEST_DEVICE_1_SN: TEST_DEVICE_1, + } + + freezer.tick(SCAN_INTERVAL) + async_fire_time_changed(hass) + await hass.async_block_till_done() + + assert TEST_DEVICE_2_SN not in coordinator.dnd_states + + async def test_coordinator_load_previous_devices_from_registry( hass: HomeAssistant, mock_amazon_devices_client: AsyncMock, @@ -170,6 +198,7 @@ async def test_sync_history_state_error( assert mock_config_entry.state is expected_state +@pytest.mark.parametrize("method_name", ["sync_dnd_state", "sync_media_state"]) @pytest.mark.parametrize( ("side_effect", "expected_state"), [ @@ -200,21 +229,26 @@ async def test_sync_history_state_error( ), ], ) -async def test_sync_media_state_auth_failed( +async def test_sync_state_error( hass: HomeAssistant, mock_amazon_devices_client: AsyncMock, mock_config_entry: MockConfigEntry, + method_name: str, side_effect: type[Exception], expected_state: ConfigEntryState, ) -> None: - """Test setup fails with ConfigEntryAuthFailed when sync_media_state raises CannotAuthenticate.""" - mock_amazon_devices_client.sync_media_state.side_effect = side_effect + """Test setup state when a sync method raises during setup.""" + getattr(mock_amazon_devices_client, method_name).side_effect = side_effect await setup_integration(hass, mock_config_entry) assert mock_config_entry.state is expected_state +@pytest.mark.parametrize( + "method_name", + ["sync_dnd_state", "sync_media_state"], +) @pytest.mark.parametrize( "error", [ @@ -228,19 +262,20 @@ async def test_sync_media_state_auth_failed( ), ], ) -async def test_media_state_sync_failure_logged_on_first_refresh( +async def test_device_state_sync_failure_logged_on_first_refresh( hass: HomeAssistant, caplog: pytest.LogCaptureFixture, mock_amazon_devices_client: AsyncMock, mock_config_entry: MockConfigEntry, + method_name: str, error: Exception, ) -> None: - """Test a failing media sync on first refresh is logged but does not block setup.""" - mock_amazon_devices_client.sync_media_state.side_effect = error + """Test a failing DND/media sync on first refresh is logged but does not block setup.""" + getattr(mock_amazon_devices_client, method_name).side_effect = error await setup_integration(hass, mock_config_entry) - assert "Sync failed for sync_media_state:" in caplog.text + assert f"Sync failed for {method_name}:" in caplog.text assert str(error) in caplog.text assert ( "Data may be missing or incomplete until updates are pushed by Amazon" @@ -249,19 +284,20 @@ async def test_media_state_sync_failure_logged_on_first_refresh( assert mock_config_entry.state is ConfigEntryState.LOADED -async def test_media_state_sync_on_device_list_change( +async def test_device_state_sync_on_device_list_change( hass: HomeAssistant, freezer: FrozenDateTimeFactory, mock_amazon_devices_client: AsyncMock, mock_config_entry: MockConfigEntry, ) -> None: - """Test media state is resynced when a device is added.""" + """Test DND and media state are resynced when a device is added.""" mock_amazon_devices_client.get_devices_data.return_value = { TEST_DEVICE_1_SN: TEST_DEVICE_1, } await setup_integration(hass, mock_config_entry) + mock_amazon_devices_client.sync_dnd_state.assert_awaited_once() mock_amazon_devices_client.sync_media_state.assert_awaited_once() mock_amazon_devices_client.get_devices_data.return_value = { @@ -273,6 +309,7 @@ async def test_media_state_sync_on_device_list_change( async_fire_time_changed(hass) await hass.async_block_till_done() + assert mock_amazon_devices_client.sync_dnd_state.call_count == 2 assert mock_amazon_devices_client.sync_media_state.call_count == 2 freezer.tick(SCAN_INTERVAL) @@ -280,4 +317,5 @@ async def test_media_state_sync_on_device_list_change( await hass.async_block_till_done() # Device list unchanged: no additional resync + assert mock_amazon_devices_client.sync_dnd_state.call_count == 2 assert mock_amazon_devices_client.sync_media_state.call_count == 2 diff --git a/tests/components/alexa_devices/test_init.py b/tests/components/alexa_devices/test_init.py index 2ce63291cc9e..6ea0ae0bb987 100644 --- a/tests/components/alexa_devices/test_init.py +++ b/tests/components/alexa_devices/test_init.py @@ -305,3 +305,4 @@ async def test_initial_sync_failure_does_not_prevent_other_syncs( assert mock_config_entry.state is ConfigEntryState.LOADED mock_amazon_devices_client.get_todo_list_items.assert_awaited_once() mock_amazon_devices_client.sync_media_state.assert_awaited_once() + mock_amazon_devices_client.sync_dnd_state.assert_awaited_once() diff --git a/tests/components/alexa_devices/test_switch.py b/tests/components/alexa_devices/test_switch.py index 4d45f52a992b..66af0fcf2465 100644 --- a/tests/components/alexa_devices/test_switch.py +++ b/tests/components/alexa_devices/test_switch.py @@ -3,7 +3,6 @@ from copy import deepcopy from unittest.mock import AsyncMock, patch -from aioamazondevices.structures import AmazonDeviceSensor from freezegun.api import FrozenDateTimeFactory import pytest from syrupy.assertion import SnapshotAssertion @@ -25,7 +24,7 @@ from homeassistant.core import HomeAssistant from homeassistant.helpers import entity_registry as er from . import assert_device_removed_and_readded, setup_integration -from .const import TEST_DEVICE_1, TEST_DEVICE_1_SN +from .const import TEST_DEVICE_1, TEST_DEVICE_1_SN, TEST_DEVICE_2, TEST_DEVICE_2_SN from tests.common import MockConfigEntry, async_fire_time_changed, snapshot_platform @@ -49,11 +48,10 @@ async def test_all_entities( async def test_switch_dnd( hass: HomeAssistant, - freezer: FrozenDateTimeFactory, mock_amazon_devices_client: AsyncMock, mock_config_entry: MockConfigEntry, ) -> None: - """Test switching DND.""" + """Test switching DND updates state optimistically.""" await setup_integration(hass, mock_config_entry) assert (state := hass.states.get(ENTITY_ID)) @@ -67,34 +65,6 @@ async def test_switch_dnd( ) assert mock_amazon_devices_client.set_do_not_disturb.call_count == 1 - - device_data = deepcopy(TEST_DEVICE_1) - device_data.sensors = { - "dnd": AmazonDeviceSensor( - name="dnd", - value=True, - error=False, - error_msg=None, - error_type=None, - scale=None, - ), - "temperature": AmazonDeviceSensor( - name="temperature", - value="22.5", - error=False, - error_msg=None, - error_type=None, - scale="CELSIUS", - ), - } - mock_amazon_devices_client.get_devices_data.return_value = { - TEST_DEVICE_1_SN: device_data - } - - freezer.tick(SCAN_INTERVAL) - async_fire_time_changed(hass) - await hass.async_block_till_done() - assert (state := hass.states.get(ENTITY_ID)) assert state.state == STATE_ON @@ -105,34 +75,121 @@ async def test_switch_dnd( blocking=True, ) - device_data.sensors = { - "dnd": AmazonDeviceSensor( - name="dnd", - value=False, - error=False, - error_msg=None, - error_type=None, - scale=None, - ), - "temperature": AmazonDeviceSensor( - name="temperature", - value="22.5", - error=False, - error_msg=None, - error_type=None, - scale="CELSIUS", - ), - } + assert mock_amazon_devices_client.set_do_not_disturb.call_count == 2 + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_OFF + + +async def test_switch_dnd_pushed_event( + hass: HomeAssistant, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test the DND switch reflects state pushed by the Amazon API.""" + await setup_integration(hass, mock_config_entry) + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_OFF + + mock_amazon_devices_client.on_dnd_event.append.assert_called_once() + event_handler = mock_amazon_devices_client.on_dnd_event.append.call_args.args[0] + + await event_handler({TEST_DEVICE_1_SN: True}) + await hass.async_block_till_done() + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_ON + + await event_handler({TEST_DEVICE_1_SN: False}) + await hass.async_block_till_done() + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_OFF + + +async def test_switch_dnd_unavailable_when_missing_from_push( + hass: HomeAssistant, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test the DND switch is unavailable while its state is missing from a push.""" + await setup_integration(hass, mock_config_entry) + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_OFF + + event_handler = mock_amazon_devices_client.on_dnd_event.append.call_args.args[0] + + await event_handler({}) + await hass.async_block_till_done() + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_UNAVAILABLE + + await event_handler({TEST_DEVICE_1_SN: True}) + await hass.async_block_till_done() + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_ON + + +async def test_switch_dnd_not_created_without_synced_state( + hass: HomeAssistant, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test the DND switch is not created until a DND state has been synced.""" + mock_amazon_devices_client.sync_dnd_state = AsyncMock() + + await setup_integration(hass, mock_config_entry) + + assert hass.states.get(ENTITY_ID) is None + + +async def test_switch_dnd_created_by_late_pushed_event( + hass: HomeAssistant, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test the DND switch is created by a push arriving after initial sync failed.""" + mock_amazon_devices_client.sync_dnd_state = AsyncMock() + + await setup_integration(hass, mock_config_entry) + + assert hass.states.get(ENTITY_ID) is None + + mock_amazon_devices_client.on_dnd_event.append.assert_called_once() + event_handler = mock_amazon_devices_client.on_dnd_event.append.call_args.args[0] + await event_handler({TEST_DEVICE_1_SN: True}) + await hass.async_block_till_done() + + assert (state := hass.states.get(ENTITY_ID)) + assert state.state == STATE_ON + + +async def test_switch_dnd_created_for_new_device( + hass: HomeAssistant, + freezer: FrozenDateTimeFactory, + mock_amazon_devices_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test a DND switch is created for a device added while HA is running.""" + new_entity_id = "switch.echo_test_2_do_not_disturb" + + await setup_integration(hass, mock_config_entry) + + assert hass.states.get(new_entity_id) is None + mock_amazon_devices_client.get_devices_data.return_value = { - TEST_DEVICE_1_SN: device_data + TEST_DEVICE_1_SN: TEST_DEVICE_1, + TEST_DEVICE_2_SN: TEST_DEVICE_2, } freezer.tick(SCAN_INTERVAL) async_fire_time_changed(hass) await hass.async_block_till_done() - assert mock_amazon_devices_client.set_do_not_disturb.call_count == 2 - assert (state := hass.states.get(ENTITY_ID)) + assert (state := hass.states.get(new_entity_id)) assert state.state == STATE_OFF