Avoid cancelling shared loader futures when a waiter is cancelled (#183836)

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
epenet
2026-10-01 14:50:11 +02:00
committed by GitHub
co-authored by Claude Opus 5.5
parent 5d651c4ad5
commit c842818fb5
4 changed files with 183 additions and 4 deletions
+4 -4
View File
@@ -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
+12
View File
@@ -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]:
+124
View File
@@ -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
+43
View File
@@ -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)