Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 69 additions & 6 deletions bellows/ezsp/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand All @@ -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()

Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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:
Expand Down
23 changes: 19 additions & 4 deletions bellows/uart.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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
17 changes: 13 additions & 4 deletions bellows/zigbee/application.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
28 changes: 28 additions & 0 deletions tests/test_application.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
14 changes: 4 additions & 10 deletions tests/test_ash.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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
Expand Down
Loading
Loading