diff --git a/bellows/ezsp/__init__.py b/bellows/ezsp/__init__.py index a158018e..12812ea0 100644 --- a/bellows/ezsp/__init__.py +++ b/bellows/ezsp/__init__.py @@ -78,6 +78,11 @@ def __init__(self, device_config: dict, application: Any | None = None): self._ezsp_version = v4.EZSPv4.VERSION self._xncp_features = FirmwareFeatures.NONE self._gw = None + # Set once the gateway reports that its transport closed. The gateway then has + # nothing left to close and, when it runs in its own thread, its event loop is + # being torn down: nothing may be sent through it and nothing needs to be + # awaited on it. + self._transport_closed = False self._protocol = None self._application = application @@ -143,13 +148,41 @@ async def _startup_reset(self) -> None: async def startup_reset(self) -> None: for attempt in range(RESET_ATTEMPTS): + if self._transport_closed: + await self.disconnect() + raise ConnectionResetError("Connection was lost during startup") + self._protocol = v4.EZSPv4(self.handle_callback, self._gw) try: await self._startup_reset() break + except asyncio.CancelledError: + # See below + await asyncio.sleep(0) + + task = asyncio.current_task() + if not self._transport_closed or ( + task is not None and task.cancelling() + ): + raise + + # A gateway in its own thread cancels the work dispatched to it when its + # transport closes: a lost connection, not our caller cancelling us + await self.disconnect() + raise ConnectionResetError( + "Connection was lost during startup" + ) from None except Exception as exc: - if attempt + 1 < RESET_ATTEMPTS: + # A gateway in its own thread reports a closed transport by queueing + # `connection_lost()` onto this loop. The failure it caused may reach us + # first, without a chance to run anything else: yield once so the + # notification is delivered before deciding what to do. + await asyncio.sleep(0) + + # Retrying is for an NCP that did not answer. Once the transport itself + # is gone, every further attempt would go to a dead gateway. + if attempt + 1 < RESET_ATTEMPTS and not self._transport_closed: LOGGER.debug( "EZSP startup/reset failed, retrying (%d/%d): %r", attempt + 1, @@ -159,10 +192,18 @@ async def startup_reset(self) -> None: continue await self.disconnect() + + if self._transport_closed and not isinstance(exc, ConnectionError): + # What failed is incidental to the transport closing + raise ConnectionResetError( + "Connection was lost during startup" + ) from exc + raise async def connect(self, *, use_thread: bool = True) -> None: assert self._gw is None + self._transport_closed = False self._gw = await bellows.uart.connect(self._config, self, use_thread=use_thread) await self.startup_reset() @@ -227,13 +268,29 @@ async def get_xncp_features(self) -> xncp.FirmwareFeatures: async def disconnect(self): self.stop_ezsp() - if self._gw: - await self._gw.disconnect() + if self._gw is None: + return + + try: + # A gateway whose transport already closed has nothing left to close, and + # its worker loop may be gone: don't call into it + if not self._transport_closed: + await self._gw.disconnect() + finally: self._gw = None async def _command(self, name: str, *args: Any, **kwargs: Any) -> Any: command = getattr(self._protocol, name) + if self._transport_closed: + LOGGER.debug( + "Couldn't send command %s(%s, %s). Connection was lost", + name, + args, + kwargs, + ) + raise EzspError("Connection was lost") + if not self.is_ezsp_running: LOGGER.debug( "Couldn't send command %s(%s, %s). EZSP is not running", @@ -330,12 +387,18 @@ async def leaveNetwork(self, timeout: float | int = NETWORK_OPS_TIMEOUT) -> None def connection_lost(self, exc): """Lost serial connection.""" - if self._application is not None: - self._application.connection_lost(exc) + self._transport_closed = True + self._notify_connection_lost(exc) def enter_failed_state(self, code: t.NcpResetCode) -> None: """UART received reset code.""" - self.connection_lost(NcpFailure(code=code)) + # The transport is still open here: a reset can recover the NCP, and + # `disconnect()` still has a transport to close + self._notify_connection_lost(NcpFailure(code=code)) + + def _notify_connection_lost(self, exc: Exception) -> None: + if self._application is not None: + self._application.connection_lost(exc) def __getattr__(self, name: str) -> Callable: if name not in self._protocol.COMMANDS: diff --git a/bellows/uart.py b/bellows/uart.py index af274dc8..54545637 100644 --- a/bellows/uart.py +++ b/bellows/uart.py @@ -105,11 +105,18 @@ async def reset(self): return await self._reset_future -async def _connect(config, api): +async def _connect(config, api, thread=None): loop = asyncio.get_event_loop() connection_done_future = loop.create_future() + if thread is not None: + # Attach the callback here, on the loop that resolves the future. Attaching + # from the caller's thread once the connection is already lost would queue it + # with a plain `call_soon` that never wakes this loop: the thread would then + # sleep in `select()` until something else happens to be dispatched to it. + connection_done_future.add_done_callback(lambda _: thread.force_stop()) + gateway = Gateway(api, connection_done_future) protocol = AshProtocol(gateway) @@ -139,13 +146,21 @@ async def connect(config, api, use_thread=True): thread = EventLoopThread() await thread.start() try: - protocol, connection_done = await thread.run_coroutine_threadsafe( - _connect(config, api) + protocol, _ = await thread.run_coroutine_threadsafe( + _connect(config, api, thread) ) + except asyncio.CancelledError: + task = asyncio.current_task() + if task is not None and task.cancelling(): + thread.force_stop() + raise + + # Not our caller: the worker stopped itself, cancelling `_connect()`, + # because the connection was lost before it returned + raise ConnectionResetError("Connection was lost while connecting") from None except Exception: thread.force_stop() raise - connection_done.add_done_callback(lambda _: thread.force_stop()) else: protocol, _ = await _connect(config, api) return protocol diff --git a/bellows/zigbee/application.py b/bellows/zigbee/application.py index c16f43eb..f9178fed 100644 --- a/bellows/zigbee/application.py +++ b/bellows/zigbee/application.py @@ -195,8 +195,15 @@ async def connect(self) -> None: await self.register_endpoints() except Exception: if self._ezsp is not None: - await self._ezsp.disconnect() - self._ezsp = None + try: + await self._ezsp.disconnect() + except Exception: + # Don't let cleanup failures mask why connecting failed + LOGGER.warning( + "Failed to disconnect after a connection failure", exc_info=True + ) + finally: + self._ezsp = None raise async def _ensure_network_running(self) -> bool: @@ -612,8 +619,10 @@ async def disconnect(self): # TODO: how do you shut down the stack? self.controller_event.clear() if self._ezsp is not None: - await self._ezsp.disconnect() - self._ezsp = None + try: + await self._ezsp.disconnect() + finally: + self._ezsp = None async def force_remove(self, dev): # This should probably be delivered to the parent device instead diff --git a/tests/test_application.py b/tests/test_application.py index 784f86c8..b1a74703 100644 --- a/tests/test_application.py +++ b/tests/test_application.py @@ -2194,6 +2194,34 @@ async def test_connect_failure(app: ControllerApplication) -> None: assert len(ezsp.disconnect.mock_calls) == 1 +async def test_connect_failure_disconnect_failure( + app: ControllerApplication, caplog +) -> None: + """Test that a failing disconnect after a connection failure doesn't mask it.""" + ezsp = app._ezsp + app._ezsp.write_config = AsyncMock(side_effect=OSError("Connection failed")) + app._ezsp.connect = AsyncMock() + app._ezsp.disconnect = AsyncMock(side_effect=RuntimeError("Uh oh")) + app._ezsp = None + + with patch("bellows.ezsp.EZSP", return_value=ezsp): + with pytest.raises(OSError, match="Connection failed"): + await app.connect() + + assert app._ezsp is None + assert "Failed to disconnect after a connection failure" in caplog.text + + +async def test_disconnect_failure(app: ControllerApplication) -> None: + """Test that EZSP is dropped even when disconnecting fails.""" + app._ezsp.disconnect = AsyncMock(side_effect=RuntimeError("Uh oh")) + + with pytest.raises(RuntimeError): + await app.disconnect() + + assert app._ezsp is None + + async def test_repair_tclk_partner_ieee( app: ControllerApplication, ieee: zigpy_t.EUI64 ) -> None: diff --git a/tests/test_ash.py b/tests/test_ash.py index f97b2c4f..f36bbd0a 100644 --- a/tests/test_ash.py +++ b/tests/test_ash.py @@ -543,11 +543,8 @@ async def test_ash_end_to_end(transport_cls: type[FakeTransport]) -> None: # Let's let a request fail due to a connectivity issue with patch.object(ncp_transport, "paused", True): - send_task = asyncio.create_task(host.send_data(b"host failure")) - await asyncio.sleep(host._t_rx_ack * 15) - - with pytest.raises(TimeoutError): - await send_task + with pytest.raises(TimeoutError): + await host.send_data(b"host failure") ncp_ezsp.data_received.reset_mock() host_ezsp.data_received.reset_mock() @@ -569,11 +566,8 @@ async def test_ash_end_to_end(transport_cls: type[FakeTransport]) -> None: assert ncp._ncp_reset_code is None with patch.object(host_transport, "paused", True): - send_task = asyncio.create_task(ncp.send_data(b"ncp failure")) - await asyncio.sleep(ncp._t_rx_ack * 15) - - with pytest.raises(TimeoutError): - await send_task + with pytest.raises(TimeoutError): + await ncp.send_data(b"ncp failure") assert ( host._ncp_reset_code is t.NcpResetCode.ERROR_EXCEEDED_MAXIMUM_ACK_TIMEOUT_COUNT diff --git a/tests/test_ezsp.py b/tests/test_ezsp.py index e2e26e1d..995bd971 100644 --- a/tests/test_ezsp.py +++ b/tests/test_ezsp.py @@ -4,6 +4,7 @@ from asyncio import timeout as asyncio_timeout import functools import logging +import threading from unittest.mock import ANY, AsyncMock, MagicMock, call, patch import pytest @@ -79,6 +80,54 @@ async def test_disconnect(ezsp_f): gw_disconnect = ezsp_f._gw.disconnect await ezsp_f.disconnect() assert len(gw_disconnect.mock_calls) == 1 + assert ezsp_f._gw is None + + +async def test_disconnect_after_transport_closed(ezsp_f): + """A gateway whose transport closed is not called into: it has nothing to close.""" + ezsp_f._application = MagicMock() + gw_disconnect = ezsp_f._gw.disconnect + exc = ConnectionResetError("Remote server closed connection") + + ezsp_f.connection_lost(exc) + await ezsp_f.disconnect() + + assert gw_disconnect.mock_calls == [] + assert ezsp_f._gw is None + assert ezsp_f._application.connection_lost.mock_calls == [call(exc)] + + +async def test_disconnect_after_failed_state(ezsp_f): + """An NCP failure leaves the transport open, so disconnecting still closes it.""" + ezsp_f._application = MagicMock() + gw_disconnect = ezsp_f._gw.disconnect + + ezsp_f.enter_failed_state(t.NcpResetCode.RESET_SOFTWARE) + await ezsp_f.disconnect() + + assert len(gw_disconnect.mock_calls) == 1 + assert ezsp_f._gw is None + + +async def test_disconnect_drops_gateway_on_failure(ezsp_f): + """The gateway reference is dropped even when closing it fails.""" + ezsp_f._gw.disconnect = AsyncMock(side_effect=RuntimeError("Uh oh")) + + with pytest.raises(RuntimeError): + await ezsp_f.disconnect() + + assert ezsp_f._gw is None + + +async def test_command_after_transport_closed(ezsp_f): + """Commands fail immediately once the transport closed, instead of timing out.""" + ezsp_f._protocol = MagicMock() + ezsp_f.connection_lost(ConnectionResetError("Remote server closed connection")) + + with pytest.raises(EzspError, match="Connection was lost"): + await EZSP._command(ezsp_f, "version") + + assert ezsp_f._protocol.version.mock_calls == [] def test_attr(ezsp_f): @@ -326,6 +375,227 @@ async def startup_reset_mock(): assert disconnect_mock.call_count == 0 +async def test_ezsp_connect_no_retry_after_transport_closed(): + """A startup reset that failed because the transport closed is not retried.""" + exc = ConnectionResetError("Remote server closed connection") + + with patch("bellows.uart.connect") as conn_mock: + ezsp = make_ezsp(version=4) + ezsp._transport_closed = True # stale state from an earlier connection + + async def startup_reset_mock(): + # The gateway reports the loss before the failure reaches the caller + ezsp.connection_lost(exc) + raise exc + + with patch.object( + ezsp, "_startup_reset", side_effect=startup_reset_mock + ) as startup_reset: + with pytest.raises(ConnectionResetError): + await ezsp.connect() + + assert startup_reset.await_count == 1 + assert conn_mock.return_value.disconnect.mock_calls == [] + assert ezsp._gw is None + + +async def test_ezsp_connect_no_retry_with_queued_connection_lost(): + """The transport-closed notification may still be queued when the attempt fails. + + A threaded gateway queues `connection_lost()` onto this loop; a proxy call that + fails synchronously reaches `startup_reset()` before it is delivered. + """ + loop = asyncio.get_running_loop() + exc = ConnectionResetError("Remote server closed connection") + + with patch("bellows.uart.connect") as conn_mock: + ezsp = make_ezsp(version=4) + + async def startup_reset_mock(): + # Queued, not yet delivered + loop.call_soon(ezsp.connection_lost, exc) + raise TypeError("'NoneType' object can't be awaited") + + with patch.object( + ezsp, "_startup_reset", side_effect=startup_reset_mock + ) as startup_reset: + with pytest.raises(ConnectionResetError, match="lost during startup") as e: + await ezsp.connect() + + assert isinstance(e.value.__cause__, TypeError) + assert startup_reset.await_count == 1 + assert conn_mock.return_value.disconnect.mock_calls == [] + assert ezsp._gw is None + + +async def test_ezsp_connect_cancelled_with_queued_connection_lost(): + """Same as above, for a teardown cancellation reaching us before the notification.""" + loop = asyncio.get_running_loop() + exc = ConnectionResetError("Remote server closed connection") + + with patch("bellows.uart.connect") as conn_mock: + ezsp = make_ezsp(version=4) + + async def startup_reset_mock(): + # Queued, not yet delivered + loop.call_soon(ezsp.connection_lost, exc) + raise asyncio.CancelledError() + + with patch.object( + ezsp, "_startup_reset", side_effect=startup_reset_mock + ) as startup_reset: + with pytest.raises(ConnectionResetError, match="lost during startup"): + await ezsp.connect() + + assert startup_reset.await_count == 1 + assert conn_mock.return_value.disconnect.mock_calls == [] + assert ezsp._gw is None + + +async def test_ezsp_connect_transport_closed_before_startup_reset(): + """A transport that closed before the first attempt is not dispatched to at all.""" + with patch("bellows.uart.connect") as conn_mock: + ezsp = make_ezsp(version=4) + + async def connect_mock(*args, **kwargs): + ezsp.connection_lost( + ConnectionResetError("Remote server closed connection") + ) + return conn_mock.return_value + + conn_mock.side_effect = connect_mock + + with patch.object(ezsp, "_startup_reset") as startup_reset: + with pytest.raises(ConnectionResetError, match="lost during startup"): + await ezsp.connect() + + assert startup_reset.await_count == 0 + assert conn_mock.return_value.disconnect.mock_calls == [] + assert ezsp._gw is None + + +async def test_ezsp_connect_cancelled_by_transport_closed(): + """A threaded gateway cancels dispatched work when its transport closes. + + That cancellation is reported as a lost connection, not propagated as-is. + """ + with patch("bellows.uart.connect") as conn_mock: + ezsp = make_ezsp(version=4) + + async def startup_reset_mock(): + ezsp.connection_lost( + ConnectionResetError("Remote server closed connection") + ) + raise asyncio.CancelledError() + + with patch.object( + ezsp, "_startup_reset", side_effect=startup_reset_mock + ) as startup_reset: + with pytest.raises(ConnectionResetError, match="lost during startup"): + await ezsp.connect() + + assert startup_reset.await_count == 1 + assert conn_mock.return_value.disconnect.mock_calls == [] + assert ezsp._gw is None + + +async def test_ezsp_connect_cancelled_without_transport_closed(): + """A cancellation not explained by a closed transport propagates as-is.""" + with patch("bellows.uart.connect"): + ezsp = make_ezsp(version=4) + + with patch.object(ezsp, "_startup_reset", side_effect=asyncio.CancelledError()): + with pytest.raises(asyncio.CancelledError): + await ezsp.connect() + + +async def test_ezsp_connect_cancelled_by_caller_after_transport_closed(): + """The caller cancelling `connect()` still wins, even after the transport closed.""" + loop = asyncio.get_running_loop() + blocked = loop.create_future() + + with patch("bellows.uart.connect"): + ezsp = make_ezsp(version=4) + + async def startup_reset_mock(): + ezsp.connection_lost( + ConnectionResetError("Remote server closed connection") + ) + blocked.set_result(None) + await loop.create_future() + + with patch.object(ezsp, "_startup_reset", side_effect=startup_reset_mock): + connect_task = asyncio.ensure_future(ezsp.connect()) + await blocked + connect_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await connect_task + + +async def test_ezsp_connect_transport_closed_during_startup_reset(): + """The bridge drops the TCP connection while EZSP waits for the startup reset. + + Connecting fails with the transport's error right away: no retry against the + torn-down gateway, no call into its worker loop, and the worker thread exits. + """ + loop = asyncio.get_running_loop() + waiting_for_reset = loop.create_future() + client_connected = loop.create_future() + + async def handle_client(reader, writer): + client_connected.set_result(writer) + + server = await asyncio.start_server(handle_client, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + + wait_for_startup_reset = uart.Gateway.wait_for_startup_reset + + async def wait_for_startup_reset_mock(self): + # Runs on the worker loop + loop.call_soon_threadsafe(waiting_for_reset.set_result, None) + await wait_for_startup_reset(self) + + ezsp = EZSP( + { + **DEVICE_CONFIG, + zigpy.config.CONF_DEVICE_PATH: f"socket://127.0.0.1:{port}", + zigpy.config.CONF_DEVICE_FLOW_CONTROL: None, + }, + application=MagicMock(), + ) + + try: + with patch.object( + uart.Gateway, "wait_for_startup_reset", wait_for_startup_reset_mock + ): + connect_task = asyncio.ensure_future(ezsp.connect(use_thread=True)) + + try: + async with asyncio_timeout(5): + client = await client_connected + await waiting_for_reset + + # The bridge drops the connection + client.close() + + with pytest.raises(ConnectionResetError): + async with asyncio_timeout(1): + await connect_task + finally: + connect_task.cancel() + # Tear down the gateway (and its thread) even if the test failed + await ezsp.disconnect() + finally: + server.close() + + assert ezsp._gw is None + assert ezsp._application.connection_lost.mock_calls == [call(ANY)] + + [t.join(1) for t in threading.enumerate() if "bellows" in t.name] + assert [t for t in threading.enumerate() if "bellows" in t.name] == [] + + async def test_ezsp_newer_version(ezsp_f): """Test newer version of ezsp.""" with patch.object( diff --git a/tests/test_uart.py b/tests/test_uart.py index 7cc44b27..2e4f550d 100644 --- a/tests/test_uart.py +++ b/tests/test_uart.py @@ -1,4 +1,5 @@ import asyncio +from asyncio import timeout as asyncio_timeout import threading from unittest.mock import AsyncMock, MagicMock, call, patch, sentinel @@ -138,6 +139,141 @@ async def mock_connect(loop, protocol_factory, *args, **kwargs): assert len(threads) == 0 +async def test_connect_threaded_cancelled_by_worker_teardown(): + """The worker cancelling `_connect()` itself is reported as a lost connection.""" + + def cancelled(coroutine): + coroutine.close() + raise asyncio.CancelledError() + + with patch.object(uart.EventLoopThread, "start", AsyncMock()), patch.object( + uart.EventLoopThread, "run_coroutine_threadsafe", side_effect=cancelled + ), patch.object(uart.EventLoopThread, "force_stop") as force_stop: + with pytest.raises(ConnectionResetError, match="lost while connecting"): + await uart.connect( + conf.SCHEMA_DEVICE( + { + conf.CONF_DEVICE_PATH: "/dev/serial", + conf.CONF_DEVICE_BAUDRATE: 115200, + } + ), + MagicMock(), + use_thread=True, + ) + + # The worker already stopped itself + assert force_stop.mock_calls == [] + + +async def test_connect_threaded_cancelled_by_caller_stops_thread(): + """A caller cancelling `connect()` still stops the worker thread.""" + loop = asyncio.get_running_loop() + started = loop.create_future() + + async def blocked(coroutine): + coroutine.close() + started.set_result(None) + await loop.create_future() + + with patch.object(uart.EventLoopThread, "start", AsyncMock()), patch.object( + uart.EventLoopThread, "run_coroutine_threadsafe", side_effect=blocked + ), patch.object(uart.EventLoopThread, "force_stop") as force_stop: + connect_task = asyncio.ensure_future( + uart.connect( + conf.SCHEMA_DEVICE( + { + conf.CONF_DEVICE_PATH: "/dev/serial", + conf.CONF_DEVICE_BAUDRATE: 115200, + } + ), + MagicMock(), + use_thread=True, + ) + ) + await started + connect_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await connect_task + + assert len(force_stop.mock_calls) == 1 + + +async def test_connect_threaded_connection_lost_before_connect_returns(): + """The connection is lost before `connect()` resumes: the worker must still stop. + + The callback that stops the worker thread used to be attached once `connect()` + resumed. A future that had already completed by then queued it onto the worker + loop from the wrong thread, without waking it: the thread slept forever. + """ + loop = asyncio.get_running_loop() + lost = loop.create_future() + threads = [] + + thread_init = uart.EventLoopThread.__init__ + + def track_thread(self): + thread_init(self) + threads.append(self) + + async def handle_client(reader, writer): + writer.close() + + server = await asyncio.start_server(handle_client, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + + appmock = MagicMock() + appmock.connection_lost.side_effect = lambda exc: lost.set_result(exc) + + run_coroutine_threadsafe = uart.EventLoopThread.run_coroutine_threadsafe + + def resume_after_loss(self, coroutine): + future = run_coroutine_threadsafe(self, coroutine) + + async def inner(): + result = await future + # The application has been told the connection is lost, so the worker + # loop has already resolved the connection-done future. Give it time to + # go idle in `select()` as well: nothing else will be dispatched to it. + async with asyncio_timeout(1): + await lost + await asyncio.sleep(0.05) + return result + + return inner() + + try: + with patch.object(uart.EventLoopThread, "__init__", track_thread), patch.object( + uart.EventLoopThread, "run_coroutine_threadsafe", resume_after_loss + ): + gw = await uart.connect( + conf.SCHEMA_DEVICE( + { + conf.CONF_DEVICE_PATH: f"socket://127.0.0.1:{port}", + conf.CONF_DEVICE_BAUDRATE: 115200, + } + ), + appmock, + use_thread=True, + ) + finally: + server.close() + + assert gw is not None + + # The worker thread stops on its own + [t.join(1) for t in threading.enumerate() if "bellows" in t.name] + leaked = [t for t in threading.enumerate() if "bellows" in t.name] + + # Never leave a stuck worker behind, even when this fails: it would hang the + # interpreter at exit + (thread,) = threads + thread.force_stop() + [t.join(1) for t in leaked] + + assert leaked == [] + + @pytest.fixture async def gw(): gw = uart.Gateway(MagicMock())