diff --git a/homeassistant/components/bluetooth/manager.py b/homeassistant/components/bluetooth/manager.py index 4a6c701ad8c6..5cb49668dd2e 100644 --- a/homeassistant/components/bluetooth/manager.py +++ b/homeassistant/components/bluetooth/manager.py @@ -180,11 +180,10 @@ class HomeAssistantBluetoothManager(BluetoothManager): def _address_disappeared(self, address: str) -> None: """Dismiss all discoveries for the given address.""" self._integration_matcher.async_clear_address(address) - for flow in self.hass.config_entries.flow.async_progress_by_init_data_type( + self.hass.config_entries.flow.async_dismiss_discovery_flows( BluetoothServiceInfoBleak, lambda service_info: bool(service_info.address == address), - ): - self.hass.config_entries.flow.async_abort(flow["flow_id"]) + ) @override async def async_setup(self) -> None: diff --git a/homeassistant/components/ssdp/scanner.py b/homeassistant/components/ssdp/scanner.py index 6e011664aedb..fd1fd4e5d162 100644 --- a/homeassistant/components/ssdp/scanner.py +++ b/homeassistant/components/ssdp/scanner.py @@ -416,14 +416,13 @@ class Scanner: self, byebye_discovery_info: _SsdpServiceInfo ) -> None: """Dismiss all discoveries for the given address.""" - for flow in self.hass.config_entries.flow.async_progress_by_init_data_type( + self.hass.config_entries.flow.async_dismiss_discovery_flows( _SsdpServiceInfo, lambda service_info: bool( service_info.ssdp_st == byebye_discovery_info.ssdp_st and service_info.ssdp_location == byebye_discovery_info.ssdp_location ), - ): - self.hass.config_entries.flow.async_abort(flow["flow_id"]) + ) async def _async_get_description_dict( self, location: str | None diff --git a/homeassistant/components/zeroconf/discovery.py b/homeassistant/components/zeroconf/discovery.py index 203936667306..d289a2b7afb7 100644 --- a/homeassistant/components/zeroconf/discovery.py +++ b/homeassistant/components/zeroconf/discovery.py @@ -255,13 +255,10 @@ class ZeroconfDiscovery: def _async_dismiss_discoveries(self, name: str) -> None: """Dismiss all discoveries for the given name.""" - for flow in self.hass.config_entries.flow.async_progress_by_init_data_type( + self.hass.config_entries.flow.async_dismiss_discovery_flows( _ZeroconfServiceInfo, lambda service_info: bool(service_info.name == name), - ): - if flow.get("context", {}).get("dismiss_protected"): - continue - self.hass.config_entries.flow.async_abort(flow["flow_id"]) + ) @callback def async_service_update( diff --git a/homeassistant/config_entries.py b/homeassistant/config_entries.py index 40f3f60abd22..a1848f3dceab 100644 --- a/homeassistant/config_entries.py +++ b/homeassistant/config_entries.py @@ -1985,6 +1985,21 @@ class ConfigEntriesFlowManager( return True return False + @callback + def async_dismiss_discovery_flows( + self, init_data_type: type, matcher: Callable[[Any], bool] + ) -> None: + """Abort discovery flows for a thing that is no longer reachable. + + Flows the user has started interacting with are left alone, because a + device often stops answering discovery precisely because it is being + paired. + """ + for flow in self.async_progress_by_init_data_type(init_data_type, matcher): + if flow["context"].get("dismiss_protected"): + continue + self.async_abort(flow["flow_id"]) + @callback def async_has_matching_flow(self, flow: ConfigFlow) -> bool: """Check if an existing matching flow is in progress.""" diff --git a/tests/components/bluetooth/test_manager.py b/tests/components/bluetooth/test_manager.py index 01d648cd5c5d..497526f1f64c 100644 --- a/tests/components/bluetooth/test_manager.py +++ b/tests/components/bluetooth/test_manager.py @@ -3,7 +3,7 @@ from datetime import timedelta import time from typing import Any -from unittest.mock import patch +from unittest.mock import call, patch from bleak.backends.scanner import AdvertisementData, BLEDevice from bluetooth_adapters import AdvertisementHistory @@ -1104,7 +1104,7 @@ async def test_goes_unavailable_dismisses_discovery_and_makes_discoverable( patch.object( hass.config_entries.flow, "async_progress_by_init_data_type", - return_value=[{"flow_id": "mock_flow_id"}], + return_value=[{"flow_id": "mock_flow_id", "context": {}}], ) as mock_async_progress_by_init_data_type, patch.object(hass.config_entries.flow, "async_abort") as mock_async_abort, patch_bluetooth_time( @@ -2088,3 +2088,23 @@ async def test_async_register_advertisement_callback(hass: HomeAssistant) -> Non cancel() inject_advertisement_with_source(hass, device, adv, "hci0") assert len(seen) == 2 + + +@pytest.mark.usefixtures("enable_bluetooth") +async def test_address_disappeared_keeps_dismiss_protected_flows( + hass: HomeAssistant, +) -> None: + """Flows the user is interacting with survive the device disappearing.""" + with ( + patch.object( + hass.config_entries.flow, + "async_progress_by_init_data_type", + return_value=[ + {"flow_id": "pending", "context": {}}, + {"flow_id": "pairing", "context": {"dismiss_protected": True}}, + ], + ), + patch.object(hass.config_entries.flow, "async_abort") as mock_async_abort, + ): + _get_manager()._address_disappeared("44:44:33:11:23:45") + assert mock_async_abort.mock_calls == [call("pending")] diff --git a/tests/components/ssdp/test_init.py b/tests/components/ssdp/test_init.py index 25dac8d3b058..2776fe703d77 100644 --- a/tests/components/ssdp/test_init.py +++ b/tests/components/ssdp/test_init.py @@ -1,7 +1,7 @@ """Test the SSDP integration.""" from ipaddress import IPv4Address -from unittest.mock import ANY, AsyncMock, patch +from unittest.mock import ANY, AsyncMock, call, patch from async_upnp_client.server import UpnpServer from async_upnp_client.ssdp_listener import SsdpListener @@ -899,12 +899,15 @@ async def test_flow_dismiss_on_byebye( ) mock_ssdp_advertisement["nts"] = "ssdp:byebye" - # ssdp:byebye advertisement should dismiss existing flows + # ssdp:byebye dismisses existing flows, but not one the user is working through with ( patch.object( hass.config_entries.flow, "async_progress_by_init_data_type", - return_value=[{"flow_id": "mock_flow_id"}], + return_value=[ + {"flow_id": "mock_flow_id", "context": {}}, + {"flow_id": "pairing", "context": {"dismiss_protected": True}}, + ], ) as mock_async_progress_by_init_data_type, patch.object(hass.config_entries.flow, "async_abort") as mock_async_abort, ): @@ -912,7 +915,7 @@ async def test_flow_dismiss_on_byebye( await hass.async_block_till_done(wait_background_tasks=True) assert len(mock_async_progress_by_init_data_type.mock_calls) == 1 - assert mock_async_abort.mock_calls[0][1][0] == "mock_flow_id" + assert mock_async_abort.mock_calls == [call("mock_flow_id")] @patch( diff --git a/tests/components/zeroconf/test_init.py b/tests/components/zeroconf/test_init.py index 80998a363303..c1a94f0e0d75 100644 --- a/tests/components/zeroconf/test_init.py +++ b/tests/components/zeroconf/test_init.py @@ -1412,7 +1412,7 @@ async def test_zeroconf_removed(hass: HomeAssistant) -> None: patch.object( hass.config_entries.flow, "async_progress_by_init_data_type", - return_value=[{"flow_id": "mock_flow_id"}], + return_value=[{"flow_id": "mock_flow_id", "context": {}}], ) as mock_async_progress_by_init_data_type, patch.object(hass.config_entries.flow, "async_abort") as mock_async_abort, patch.object( diff --git a/tests/test_config_entries.py b/tests/test_config_entries.py index afe1735c4f14..8efd9996fe08 100644 --- a/tests/test_config_entries.py +++ b/tests/test_config_entries.py @@ -11192,6 +11192,73 @@ async def test_discovery_flow_dismiss_protected_on_configure( assert result["type"] is data_entry_flow.FlowResultType.CREATE_ENTRY +async def test_async_dismiss_discovery_flows( + hass: HomeAssistant, + manager: config_entries.ConfigEntries, +) -> None: + """Test dismissing discovery flows skips the ones the user is working through.""" + mock_integration( + hass, + MockModule("comp", async_setup_entry=AsyncMock(return_value=True)), + ) + mock_platform(hass, "comp.config_flow", None) + + class TestFlow(config_entries.ConfigFlow): + """Test flow.""" + + VERSION = 1 + + async def async_step_zeroconf(self, discovery_info): + """Test zeroconf step.""" + return self.async_show_form(step_id="confirm") + + async def async_step_confirm(self, user_input=None): + """Test confirm step.""" + return self.async_show_form(step_id="confirm") + + def _service_info(name: str) -> ZeroconfServiceInfo: + return ZeroconfServiceInfo( + ip_address=ip_address("192.168.1.1"), + ip_addresses=[ip_address("192.168.1.1")], + hostname="test.local.", + name=name, + port=80, + properties={}, + type="_tcp.local.", + ) + + with mock_config_flow("comp", TestFlow): + untouched = await manager.flow.async_init( + "comp", + context={"source": config_entries.SOURCE_ZEROCONF}, + data=_service_info("other._tcp.local."), + ) + # Matches, and the user has not touched it, so this is the one dismissed + await manager.flow.async_init( + "comp", + context={"source": config_entries.SOURCE_ZEROCONF}, + data=_service_info("test._tcp.local."), + ) + pairing = await manager.flow.async_init( + "comp", + context={"source": config_entries.SOURCE_ZEROCONF}, + data=_service_info("test._tcp.local."), + ) + # Interacting with a flow protects it from being dismissed + await manager.flow.async_configure(pairing["flow_id"]) + assert _get_flow_context(manager, pairing["flow_id"])["dismiss_protected"] + + manager.flow.async_dismiss_discovery_flows( + ZeroconfServiceInfo, + lambda service_info: service_info.name == "test._tcp.local.", + ) + + assert {flow["flow_id"] for flow in manager.flow.async_progress()} == { + untouched["flow_id"], + pairing["flow_id"], + } + + async def test_user_flow_not_dismiss_protected_on_configure( hass: HomeAssistant, manager: config_entries.ConfigEntries,