diff --git a/homeassistant/loader.py b/homeassistant/loader.py index 9891109dd043..f4a15981e646 100644 --- a/homeassistant/loader.py +++ b/homeassistant/loader.py @@ -47,7 +47,7 @@ from .generated.usb import USB from .generated.zeroconf import HOMEKIT, ZEROCONF from .helpers.json import cached_json_fragment, json_fragment from .helpers.typing import UNDEFINED, UndefinedType -from .util.async_ import create_eager_task +from .util.async_ import create_eager_task, wait_shared_future from .util.hass_dict import HassKey from .util.json import JSON_DECODE_EXCEPTIONS, json_loads @@ -347,7 +347,7 @@ async def async_get_custom_components( return comps if isinstance(comps_or_future, asyncio.Future): - return await comps_or_future + return await wait_shared_future(comps_or_future) return comps_or_future @@ -1013,7 +1013,7 @@ class Integration: return cache[domain] if self._component_future: - return await self._component_future + return await wait_shared_future(self._component_future) if debug := _LOGGER.isEnabledFor(logging.DEBUG): start = time.perf_counter() @@ -1226,7 +1226,7 @@ class Integration: if in_progress_imports: for platform_name, future in in_progress_imports.items(): - platforms[platform_name] = await future + platforms[platform_name] = await wait_shared_future(future) return platforms diff --git a/homeassistant/util/async_.py b/homeassistant/util/async_.py index 1a910b3bbe8a..701d7b19e3d9 100644 --- a/homeassistant/util/async_.py +++ b/homeassistant/util/async_.py @@ -8,6 +8,7 @@ from asyncio import ( TimerHandle, gather, get_running_loop, + wait, ) from collections.abc import Awaitable, Callable, Coroutine import concurrent.futures @@ -47,6 +48,17 @@ def cancelling(task: Future[Any]) -> bool: return bool((cancelling_ := getattr(task, "cancelling", None)) and cancelling_()) +async def wait_shared_future[_T](future: Future[_T]) -> _T: + """Wait for a future shared with other callers without cancelling it. + + Awaiting a shared future directly cancels it when the waiter is cancelled, + which breaks the other waiters and makes the owner's set_result raise. + """ + if not future.done(): + await wait((future,)) + return future.result() + + def run_callback_threadsafe[_T, *_Ts]( loop: AbstractEventLoop, callback: Callable[[*_Ts], _T], *args: *_Ts ) -> concurrent.futures.Future[_T]: diff --git a/tests/test_loader.py b/tests/test_loader.py index f2b764c8b79a..574e7593d30b 100644 --- a/tests/test_loader.py +++ b/tests/test_loader.py @@ -839,6 +839,33 @@ async def test_get_custom_components(hass: HomeAssistant) -> None: mock_get.assert_called_once_with(hass) +@pytest.mark.usefixtures("enable_custom_integrations") +async def test_get_custom_components_concurrent_load_cancelled( + hass: HomeAssistant, +) -> None: + """Verify cancelling a waiting caller does not break the in-progress load.""" + custom_components = {"test_1": _get_test_integration(hass, "test_1", False)} + start_event = threading.Event() + load_event = asyncio.Event() + + def get_custom_components(hass: HomeAssistant) -> dict[str, loader.Integration]: + hass.loop.call_soon_threadsafe(load_event.set) + start_event.wait() + return custom_components + + with patch("homeassistant.loader._get_custom_components", get_custom_components): + load_task1 = asyncio.create_task(loader.async_get_custom_components(hass)) + load_task2 = asyncio.create_task(loader.async_get_custom_components(hass)) + await load_event.wait() + load_task2.cancel() + with pytest.raises(asyncio.CancelledError): + await load_task2 + start_event.set() + assert await load_task1 == custom_components + + assert await loader.async_get_custom_components(hass) == custom_components + + @pytest.mark.usefixtures("enable_custom_integrations") async def test_custom_component_overwriting_core(hass: HomeAssistant) -> None: """Test loading a custom component that overwrites a core component.""" @@ -1496,6 +1523,55 @@ async def test_async_get_component_concurrent_loads(hass: HomeAssistant) -> None assert config_flow_module_name in imports +@pytest.mark.usefixtures("enable_custom_integrations") +async def test_async_get_component_concurrent_load_cancelled( + hass: HomeAssistant, +) -> None: + """Verify cancelling a waiting caller does not break the in-progress load.""" + integration = await loader.async_get_integration( + hass, "test_package_loaded_executor" + ) + config_flow_module_name = f"{integration.pkg_path}.config_flow" + module_mock = MagicMock(__file__="__init__.py") + config_flow_module_mock = MagicMock(__file__="config_flow.py") + start_event = threading.Event() + import_event = asyncio.Event() + + def import_module(name: str) -> Any: + hass.loop.call_soon_threadsafe(import_event.set) + start_event.wait() + if name == integration.pkg_path: + return module_mock + if name == config_flow_module_name: + return config_flow_module_mock + raise ImportError + + modules_without_integration = { + k: v + for k, v in sys.modules.items() + if k not in (config_flow_module_name, integration.pkg_path) + } + with ( + patch.dict( + "sys.modules", + {**modules_without_integration}, + clear=True, + ), + patch("homeassistant.loader.importlib.import_module", import_module), + ): + load_task1 = asyncio.create_task(integration.async_get_component()) + load_task2 = asyncio.create_task(integration.async_get_component()) + await import_event.wait() + load_task2.cancel() + with pytest.raises(asyncio.CancelledError): + await load_task2 + start_event.set() + comp1 = await load_task1 + assert integration._component_future is None + + assert comp1 is module_mock + + async def test_async_get_component_deadlock_fallback( hass: HomeAssistant, caplog: pytest.LogCaptureFixture ) -> None: @@ -1973,6 +2049,54 @@ async def test_async_get_platforms_concurrent_loads(hass: HomeAssistant) -> None assert integration.get_platform_cached("button") is button_module_mock +@pytest.mark.usefixtures("enable_custom_integrations") +async def test_async_get_platforms_concurrent_load_cancelled( + hass: HomeAssistant, +) -> None: + """Verify cancelling a waiting caller does not break the in-progress load.""" + integration = await loader.async_get_integration( + hass, "test_package_loaded_executor" + ) + await integration.async_get_component() + + button_module_name = f"{integration.pkg_path}.button" + button_module_mock = MagicMock() + start_event = threading.Event() + import_event = asyncio.Event() + + def import_module(name: str) -> Any: + hass.loop.call_soon_threadsafe(import_event.set) + start_event.wait() + if name == button_module_name: + return button_module_mock + raise ImportError + + modules_without_button = { + k: v + for k, v in sys.modules.items() + if k not in (button_module_name, integration.pkg_path) + } + with ( + patch.dict( + "sys.modules", + modules_without_button, + clear=True, + ), + patch("homeassistant.loader.importlib.import_module", import_module), + ): + load_task1 = asyncio.create_task(integration.async_get_platforms(["button"])) + load_task2 = asyncio.create_task(integration.async_get_platforms(["button"])) + await import_event.wait() + load_task2.cancel() + with pytest.raises(asyncio.CancelledError): + await load_task2 + start_event.set() + load_result1 = await load_task1 + + assert load_result1 == {"button": button_module_mock} + assert integration.get_platform_cached("button") is button_module_mock + + @pytest.mark.usefixtures("enable_custom_integrations") async def test_integration_warnings( hass: HomeAssistant, caplog: pytest.LogCaptureFixture diff --git a/tests/util/test_async.py b/tests/util/test_async.py index 0ea50995d328..c04d3e2328ca 100644 --- a/tests/util/test_async.py +++ b/tests/util/test_async.py @@ -61,6 +61,49 @@ async def test_gather_with_limited_concurrency() -> None: assert results == [2, 2, -1, -1] +async def test_wait_shared_future_result() -> None: + """Test wait_shared_future returns the result of a pending or done future.""" + loop = asyncio.get_running_loop() + future: asyncio.Future[int] = loop.create_future() + loop.call_soon(future.set_result, 42) + + assert await hasync.wait_shared_future(future) == 42 + assert await hasync.wait_shared_future(future) == 42 + + +async def test_wait_shared_future_exception() -> None: + """Test wait_shared_future raises the exception of the future.""" + loop = asyncio.get_running_loop() + future: asyncio.Future[int] = loop.create_future() + loop.call_soon(future.set_exception, ValueError("boom")) + + with pytest.raises(ValueError, match="boom"): + await hasync.wait_shared_future(future) + + +async def test_wait_shared_future_waiter_cancelled() -> None: + """Test cancelling a waiter leaves the shared future intact and logs nothing.""" + loop = asyncio.get_running_loop() + future: asyncio.Future[int] = loop.create_future() + exception_handler = Mock() + loop.set_exception_handler(exception_handler) + + waiter = asyncio.create_task(hasync.wait_shared_future(future)) + await asyncio.sleep(0) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + + assert not future.cancelled() + future.set_exception(ValueError("boom")) + with pytest.raises(ValueError, match="boom"): + future.result() + await asyncio.sleep(0) + + # asyncio.shield would report "exception in shielded future" here + exception_handler.assert_not_called() + + async def test_shutdown_run_callback_threadsafe(hass: HomeAssistant) -> None: """Test we can shutdown run_callback_threadsafe.""" hasync.shutdown_run_callback_threadsafe(hass.loop)