diff --git a/homeassistant/components/device_tracker/entity.py b/homeassistant/components/device_tracker/entity.py index 4351ba175091..6461123c2f3e 100644 --- a/homeassistant/components/device_tracker/entity.py +++ b/homeassistant/components/device_tracker/entity.py @@ -1,6 +1,5 @@ """Provide functionality to keep track of devices.""" -import asyncio import logging from typing import TYPE_CHECKING, Any, final, override @@ -38,7 +37,6 @@ from homeassistant.helpers.device_registry import ( ) from homeassistant.helpers.dispatcher import async_dispatcher_send from homeassistant.helpers.entity import Entity, EntityDescription -from homeassistant.helpers.entity_platform import EntityPlatform from homeassistant.helpers.event import async_track_state_change_event from homeassistant.loader import async_suggest_report_issue from homeassistant.util.hass_dict import HassKey @@ -657,26 +655,24 @@ class ScannerEntity( or self._async_mac_address_registered() ) - @callback @override - def add_to_platform_start( - self, - hass: HomeAssistant, - platform: EntityPlatform, - parallel_updates: asyncio.Semaphore | None, - ) -> None: - """Start adding an entity to a platform.""" - super().add_to_platform_start(hass, platform, parallel_updates) + async def async_prepare_to_add_to_hass(self) -> None: + """Run before the entity is added to hass. + + Registers the MAC address before the entity is added so a tracker that is + created disabled can still be enabled later when its device becomes known. + """ + await super().async_prepare_to_add_to_hass() if self.mac_address and self.unique_id: _async_register_mac( - hass, - platform.platform_name, + self.hass, + self.platform.platform_name, self.mac_address, self.unique_id, ) if self.is_connected and self.ip_address: _async_connected_device_registered( - hass, + self.hass, self.mac_address, self.ip_address, self.hostname, diff --git a/homeassistant/helpers/entity.py b/homeassistant/helpers/entity.py index de7b9d506a8a..ec6685d16bc5 100644 --- a/homeassistant/helpers/entity.py +++ b/homeassistant/helpers/entity.py @@ -1509,14 +1509,42 @@ class Entity( else: self.hass.states.async_remove(self.entity_id, context=self._context) + async def async_prepare_to_add_to_hass(self) -> None: + """Run before the entity is added to hass. + + Called on every add attempt, before the platform processes the entity + registry and before its state is written, including for adds which + will be aborted, e.g. because the entity is disabled. Adding may not + complete; register cleanup with async_on_remove. + + To be extended by integrations. + """ + async def async_added_to_hass(self) -> None: - """Run when entity about to be added to hass. + """Run when the entity has been added to hass. + + Called as the last step of a successful add: after the entity has its + entity_id (and its registry entry, if it has a unique_id) and immediately + before its state is written for the first time. Use it to subscribe to + events, register update listeners and fetch initial data. + + Not called when adding the entity is aborted, e.g. because the entity is + disabled or its entity_id or unique_id collides with an existing entity. To be extended by integrations. """ async def async_will_remove_from_hass(self) -> None: - """Run when entity will be removed from hass. + """Run when the entity is about to be removed from hass. + + The counterpart to async_added_to_hass: called when the entity is removed + for an entity that was successfully added. Use it to undo work done in + async_added_to_hass, e.g. unsubscribe from events or release resources. + + Not called when adding the entity is aborted before it finished being + added; on that path only the callbacks registered with async_on_remove + run. Register cleanup for anything set up before the add completed with + async_on_remove so it runs on both an aborted add and a normal removal. To be extended by integrations. """ diff --git a/homeassistant/helpers/entity_platform.py b/homeassistant/helpers/entity_platform.py index 9a93678068ee..b97569d8956f 100644 --- a/homeassistant/helpers/entity_platform.py +++ b/homeassistant/helpers/entity_platform.py @@ -868,6 +868,7 @@ class EntityPlatform: self._get_parallel_updates_semaphore(hasattr(entity, "update")), ) try: + await entity.async_prepare_to_add_to_hass() restored = await self._async_add_entity_impl( entity, update_before_add, entity_registry, config_subentry_id ) diff --git a/tests/helpers/test_entity.py b/tests/helpers/test_entity.py index 78ac0ff62e3a..485d3d3e702c 100644 --- a/tests/helpers/test_entity.py +++ b/tests/helpers/test_entity.py @@ -3145,6 +3145,117 @@ async def test_platform_state_fail_to_add_rollback_raises( assert "Failed to add entity" in caplog.text +async def test_async_prepare_to_add_to_hass_runs_before_registration( + hass: HomeAssistant, entity_registry: er.EntityRegistry +) -> None: + """Test async_prepare_to_add_to_hass runs before registration and the state write. + + It is awaited during the add, before async_added_to_hass, before the entity is + registered in the entity registry and before it is written to the state machine. + """ + events: list[str] = [] + observed: dict[str, Any] = {} + + class MockEntity(entity.Entity): + _attr_unique_id = "5678" + + async def async_prepare_to_add_to_hass(self) -> None: + await super().async_prepare_to_add_to_hass() + events.append("before") + observed["registry_entry"] = self.registry_entry + observed["registered"] = entity_registry.async_get_entity_id( + "test", "test_platform", "5678" + ) + observed["states_before"] = len(hass.states.async_all()) + + async def async_added_to_hass(self) -> None: + await super().async_added_to_hass() + events.append("added") + + platform = MockEntityPlatform(hass, domain="test") + ent = MockEntity() + await platform.async_add_entities([ent]) + + assert events == ["before", "added"] + # During the hook the entity was not yet registered nor written to the state machine + assert observed["registry_entry"] is None + assert observed["registered"] is None + assert len(hass.states.async_all()) == observed["states_before"] + 1 + # After the add it is registered, added and has a state + assert entity_registry.async_get_entity_id("test", "test_platform", "5678") + assert ent._platform_state is entity.EntityPlatformState.ADDED + assert hass.states.get(ent.entity_id) is not None + + +async def test_async_prepare_to_add_to_hass_runs_for_disabled_entity( + hass: HomeAssistant, entity_registry: er.EntityRegistry +) -> None: + """Test async_prepare_to_add_to_hass runs even for a disabled, aborted entity.""" + events: list[str] = [] + + class MockEntity(entity.Entity): + _attr_unique_id = "5678" + _attr_entity_registry_enabled_default = False + + async def async_prepare_to_add_to_hass(self) -> None: + await super().async_prepare_to_add_to_hass() + events.append("before") + + async def async_added_to_hass(self) -> None: + await super().async_added_to_hass() + events.append("added") + + platform = MockEntityPlatform(hass, domain="test") + ent = MockEntity() + await platform.async_add_entities([ent]) + + # The entity is aborted for being disabled: async_added_to_hass never runs, + # but async_prepare_to_add_to_hass still does. + assert events == ["before"] + entity_id = entity_registry.async_get_entity_id("test", "test_platform", "5678") + assert entity_id is not None + assert ( + entity_registry.async_get(entity_id).disabled_by + is er.RegistryEntryDisabler.INTEGRATION + ) + assert ent._platform_state is entity.EntityPlatformState.REMOVED + assert ent.hass is None + assert hass.states.get(entity_id) is None + + +async def test_async_prepare_to_add_to_hass_raising_aborts_add( + hass: HomeAssistant, + entity_registry: er.EntityRegistry, + caplog: pytest.LogCaptureFixture, +) -> None: + """Test a raising async_prepare_to_add_to_hass aborts the add cleanly. + + The entity must be aborted instead of being left stuck in the ADDING state, + and it must never be registered or written to the state machine. + """ + + class MockEntity(entity.Entity): + _attr_unique_id = "5678" + + async def async_prepare_to_add_to_hass(self) -> None: + raise ValueError("Failed before add") + + async def async_added_to_hass(self) -> None: + raise AssertionError("async_added_to_hass must not run") + + platform = MockEntityPlatform(hass, domain="test") + ent = MockEntity() + assert ent._platform_state is entity.EntityPlatformState.NOT_ADDED + await platform.async_add_entities([ent]) + + assert ent._platform_state is entity.EntityPlatformState.REMOVED + assert ent.hass is None + assert ent.platform is None + assert entity_registry.async_get_entity_id("test", "test_platform", "5678") is None + assert hass.states.async_all() == [] + assert "Failed before add" in caplog.text + + async def test_platform_state_write_from_init( hass: HomeAssistant, caplog: pytest.LogCaptureFixture ) -> None: