diff --git a/homeassistant/helpers/entity.py b/homeassistant/helpers/entity.py index 5fcf8a5c57a6..ce38b085bd45 100644 --- a/homeassistant/helpers/entity.py +++ b/homeassistant/helpers/entity.py @@ -54,6 +54,7 @@ from homeassistant.core_config import DATA_CUSTOMIZE from homeassistant.exceptions import HomeAssistantError, NoEntitySpecifiedError from homeassistant.loader import async_suggest_report_issue from homeassistant.util import ensure_unique_string, slugify +from homeassistant.util.async_ import wait_shared_future from homeassistant.util.frozen_dataclass_compat import FrozenOrThawed from . import device_registry as dr, entity_registry as er @@ -1463,7 +1464,7 @@ class Entity( or if force_remove=True, its state will be removed. """ if self.__remove_future is not None: - await self.__remove_future + await wait_shared_future(self.__remove_future) return self.__remove_future = self.hass.loop.create_future() diff --git a/tests/helpers/test_entity.py b/tests/helpers/test_entity.py index f645afb4275a..c81a443dd935 100644 --- a/tests/helpers/test_entity.py +++ b/tests/helpers/test_entity.py @@ -685,6 +685,38 @@ async def test_async_remove_twice(hass: HomeAssistant) -> None: assert ent._platform_state is entity.EntityPlatformState.REMOVED +async def test_async_remove_cancel_concurrent_waiter(hass: HomeAssistant) -> None: + """Test cancelling a concurrent remove does not break the in-progress remove.""" + release = asyncio.Event() + + class MockEntitySlowRemoval(entity.Entity): + """Entity that blocks while being removed.""" + + async def async_will_remove_from_hass(self) -> None: + """Block until released.""" + await release.wait() + + platform = MockEntityPlatform(hass, domain="test") + ent = MockEntitySlowRemoval() + ent.entity_id = "test.test" + await platform.async_add_entities([ent]) + + owner = hass.async_create_task(ent.async_remove()) + await asyncio.sleep(0) + waiter = hass.async_create_task(ent.async_remove()) + other_waiter = hass.async_create_task(ent.async_remove()) + await asyncio.sleep(0) + + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + + release.set() + await owner + await other_waiter + assert hass.states.get("test.test") is None + + async def test_set_context(hass: HomeAssistant) -> None: """Test setting context.""" context = Context()