mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 14:29:21 -04:00
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:
co-authored by
Claude Opus 5.5
parent
5d651c4ad5
commit
c842818fb5
@@ -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
|
||||
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user