diff --git a/src/rpcclient/rpcclient/__main__.py b/src/rpcclient/rpcclient/__main__.py index fc5b0bed..77e05f51 100644 --- a/src/rpcclient/rpcclient/__main__.py +++ b/src/rpcclient/rpcclient/__main__.py @@ -7,6 +7,7 @@ from rpcclient.client_manager import ClientManager from rpcclient.clients.darwin.client import DarwinClient from rpcclient.console.console import Console, disable_loggers +from rpcclient.core.client import ClientEvent, CoreClient from rpcclient.core.webdav_mount import ( mount_webdav_volume, reveal_in_file_manager, @@ -99,6 +100,30 @@ def rpcclient( Console(manager).interactive(switch_cid=cid, startup_files=startup_files) +_DISCONNECT_HEARTBEAT_INTERVAL = 5.0 + + +async def _wait_for_disconnect(client: CoreClient) -> None: + """Block until the RPC connection to the target drops. + + The WebDAV server bridges every request to ``client``; if the target goes away (device + reboot, cable pull, server killed) the mount would otherwise keep serving errors until the + user hits Ctrl-C. ``ClientEvent.TERMINATED`` fires the moment an in-flight call fails, and a + periodic liveness probe catches a disconnect that happens while the mount is idle. + """ + loop = asyncio.get_running_loop() + terminated = asyncio.Event() + client.notifier.register_once(ClientEvent.TERMINATED, lambda *_: loop.call_soon_threadsafe(terminated.set)) + while not terminated.is_set(): + try: + await asyncio.wait_for(terminated.wait(), timeout=_DISCONNECT_HEARTBEAT_INTERVAL) + except asyncio.TimeoutError: + try: + await client.symbols.getpid() + except Exception: + return + + def _require_mount_tool(mount: bool) -> None: """Fail early if ``--mount`` was requested but no WebDAV mount tool is available on this host.""" if mount and not webdav_mount_supported(): @@ -141,7 +166,8 @@ async def _setup(): click.echo("Press Ctrl-C to stop.") try: - run_in_loop(asyncio.Event().wait()) + run_in_loop(_wait_for_disconnect(client)) + click.echo(f"Connection to {hostname} lost; stopping.") except KeyboardInterrupt: pass finally: diff --git a/src/rpcclient/tests/test_webdav.py b/src/rpcclient/tests/test_webdav.py index 3242f126..2cc16483 100644 --- a/src/rpcclient/tests/test_webdav.py +++ b/src/rpcclient/tests/test_webdav.py @@ -1,6 +1,7 @@ import asyncio import os from contextlib import suppress +from typing import Any, cast import click import httpx @@ -9,6 +10,45 @@ from tests._types import Client +@pytest.mark.asyncio +async def test_wait_for_disconnect_returns_on_terminated_event() -> None: + # An in-flight WebDAV request that fails fires ClientEvent.TERMINATED; the serve loop must + # unblock immediately instead of waiting for the next heartbeat probe. + from rpcclient.__main__ import _wait_for_disconnect + from rpcclient.core.client import ClientEvent + from rpcclient.event_notifier import EventNotifier + + class _FakeClient: + def __init__(self) -> None: + self.notifier: EventNotifier = EventNotifier() + + client = _FakeClient() + waiter = asyncio.ensure_future(_wait_for_disconnect(cast("Any", client))) + await asyncio.sleep(0) + client.notifier.notify(ClientEvent.TERMINATED, 0) + await asyncio.wait_for(waiter, timeout=1) + + +@pytest.mark.asyncio +async def test_wait_for_disconnect_detects_idle_disconnect_via_heartbeat(monkeypatch) -> None: + # With no WebDAV traffic, a disconnect is noticed only by the periodic liveness probe. + from rpcclient import __main__ as cli + from rpcclient.event_notifier import EventNotifier + + monkeypatch.setattr(cli, "_DISCONNECT_HEARTBEAT_INTERVAL", 0.01) + + class _DeadSymbols: + async def getpid(self) -> None: + raise ConnectionError("device disconnected") + + class _DeadClient: + def __init__(self) -> None: + self.notifier: EventNotifier = EventNotifier() + self.symbols = _DeadSymbols() + + await asyncio.wait_for(cli._wait_for_disconnect(cast("Any", _DeadClient())), timeout=1) + + def test_webdav_cli_subcommand_registered() -> None: from rpcclient.__main__ import rpcclient