diff --git a/asyncssh/__init__.py b/asyncssh/__init__.py index fb316e0e..a1f589dd 100644 --- a/asyncssh/__init__.py +++ b/asyncssh/__init__.py @@ -41,6 +41,7 @@ from .config import ConfigParseError from .forward import SSHForwarder +from .forward import SSHPathForwardTracker, SSHPortForwardTracker from .connection import SSHAcceptor, SSHClientConnection, SSHServerConnection from .connection import SSHClientConnectionOptions, SSHServerConnectionOptions @@ -148,7 +149,8 @@ 'SSHClientChannel', 'SSHClientConnection', 'SSHClientConnectionOptions', 'SSHClientProcess', 'SSHClientSession', 'SSHCompletedProcess', 'SSHForwarder', 'SSHKey', 'SSHKeyPair', 'SSHKnownHosts', - 'SSHLineEditorChannel', 'SSHListener', 'SSHReader', 'SSHServer', + 'SSHLineEditorChannel', 'SSHListener', 'SSHPathForwardTracker', + 'SSHPortForwardTracker', 'SSHReader', 'SSHServer', 'SSHServerChannel', 'SSHServerConnection', 'SSHServerConnectionOptions', 'SSHServerProcess', 'SSHServerProcessFactory', 'SSHServerSession', diff --git a/asyncssh/connection.py b/asyncssh/connection.py index 89fdb165..1e1900c9 100644 --- a/asyncssh/connection.py +++ b/asyncssh/connection.py @@ -85,7 +85,8 @@ from .encryption import encryption_needs_mac from .encryption import get_encryption_params, get_encryption -from .forward import SSHForwarder +from .forward import SSHForwarder, _track_remote_path, _track_remote_port +from .forward import SSHPortForwardTrackerFactory, SSHPathForwardTrackerFactory from .gss import GSSBase, GSSClient, GSSServer, GSSError @@ -3152,7 +3153,8 @@ async def create_unix_connection( raise NotImplementedError async def forward_connection( - self, dest_host: str, dest_port: int) -> SSHForwarder: + self, dest_host: str, dest_port: int, peer_factory: Callable[ + [], SSHForwarder] = SSHForwarder) -> SSHForwarder: """Forward a tunneled TCP connection This method is a coroutine which can be returned by a @@ -3163,15 +3165,18 @@ async def forward_connection( The hostname or address to forward the connections to :param dest_port: The port number to forward the connections to + :param peer_factory: (optional) + Returns the local end's protocol, defaulting to `SSHForwarder` :type dest_host: `str` or `None` :type dest_port: `int` + :type peer_factory: `callable` :returns: :class:`asyncio.BaseProtocol` """ try: - _, peer = await self._loop.create_connection(SSHForwarder, + _, peer = await self._loop.create_connection(peer_factory, dest_host, dest_port) self.logger.info(' Forwarding TCP connection to %s', @@ -3181,7 +3186,9 @@ async def forward_connection( return SSHForwarder(cast(SSHForwarder, peer)) - async def forward_unix_connection(self, dest_path: str) -> SSHForwarder: + async def forward_unix_connection( + self, dest_path: str, peer_factory: Callable[ + [], SSHForwarder] = SSHForwarder) -> SSHForwarder: """Forward a tunneled UNIX domain socket connection This method is a coroutine which can be returned by a @@ -3190,7 +3197,10 @@ async def forward_unix_connection(self, dest_path: str) -> SSHForwarder: :param dest_path: The path to forward the connection to + :param peer_factory: (optional) + Returns the local end's protocol, defaulting to `SSHForwarder` :type dest_path: `str` + :type peer_factory: `callable` :returns: :class:`asyncio.BaseProtocol` @@ -3198,7 +3208,7 @@ async def forward_unix_connection(self, dest_path: str) -> SSHForwarder: try: _, peer = \ - await self._loop.create_unix_connection(SSHForwarder, dest_path) + await self._loop.create_unix_connection(peer_factory, dest_path) self.logger.info(' Forwarding UNIX connection to %s', dest_path) except OSError as exc: @@ -3210,7 +3220,9 @@ async def forward_unix_connection(self, dest_path: str) -> SSHForwarder: async def forward_local_port( self, listen_host: str, listen_port: int, dest_host: str, dest_port: int, - accept_handler: Optional[SSHAcceptHandler] = None) -> SSHListener: + accept_handler: Optional[SSHAcceptHandler] = None, + tracker_factory: + Optional[SSHPortForwardTrackerFactory] = None) -> SSHListener: """Set up local port forwarding This method is a coroutine which attempts to set up port @@ -3233,11 +3245,17 @@ async def forward_local_port( or not to allow connection forwarding, returning `True` to accept the connection and begin forwarding or `False` to reject and close it. + :param tracker_factory: + An optional callable invoked once per accepted connection + which returns a new :class:`SSHPortForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_host: `str` :type listen_port: `int` :type dest_host: `str` :type dest_port: `int` :type accept_handler: `callable` or coroutine + :type tracker_factory: :class:`SSHPortForwardTrackerFactory` :returns: :class:`SSHListener` @@ -3278,10 +3296,9 @@ async def tunnel_connection( (dest_host, dest_port)) try: - listener = await create_tcp_forward_listener(self, self._loop, - tunnel_connection, - listen_host, - listen_port) + listener = await create_tcp_forward_listener( + self, self._loop, tunnel_connection, listen_host, listen_port, + tracker_factory) except OSError as exc: self.logger.debug1('Failed to create local TCP listener: %s', exc) raise @@ -3297,8 +3314,10 @@ async def tunnel_connection( return listener @async_context_manager - async def forward_local_path(self, listen_path: str, - dest_path: str) -> SSHListener: + async def forward_local_path( + self, listen_path: str, dest_path: str, + tracker_factory: + Optional[SSHPathForwardTrackerFactory] = None) -> SSHListener: """Set up local UNIX domain socket forwarding This method is a coroutine which attempts to set up UNIX domain @@ -3311,8 +3330,14 @@ async def forward_local_path(self, listen_path: str, The path on the local host to listen on :param dest_path: The path on the remote host to forward the connections to + :param tracker_factory: + An optional callable invoked once per accepted connection + which returns a new :class:`SSHPathForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_path: `str` :type dest_path: `str` + :type tracker_factory: :class:`SSHPathForwardTrackerFactory` :returns: :class:`SSHListener` @@ -3332,9 +3357,9 @@ async def tunnel_connection( listen_path, dest_path) try: - listener = await create_unix_forward_listener(self, self._loop, - tunnel_connection, - listen_path) + listener = await create_unix_forward_listener( + self, self._loop, tunnel_connection, listen_path, + tracker_factory) except OSError as exc: self.logger.debug1('Failed to create local UNIX listener: %s', exc) raise @@ -5304,7 +5329,9 @@ async def open_tap(self, *args: object, **kwargs: object) -> \ @async_context_manager async def forward_local_port_to_path( self, listen_host: str, listen_port: int, dest_path: str, - accept_handler: Optional[SSHAcceptHandler] = None) -> SSHListener: + accept_handler: Optional[SSHAcceptHandler] = None, + tracker_factory: + Optional[SSHPortForwardTrackerFactory] = None) -> SSHListener: """Set up local TCP port forwarding to a remote UNIX domain socket This method is a coroutine which attempts to set up port @@ -5325,10 +5352,16 @@ async def forward_local_port_to_path( or not to allow connection forwarding, returning `True` to accept the connection and begin forwarding or `False` to reject and close it. + :param tracker_factory: + An optional callable invoked once per accepted connection + which returns a new :class:`SSHPortForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_host: `str` :type listen_port: `int` :type dest_path: `str` :type accept_handler: `callable` or coroutine + :type tracker_factory: :class:`SSHPortForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5362,10 +5395,9 @@ async def tunnel_connection( (listen_host, listen_port), dest_path) try: - listener = await create_tcp_forward_listener(self, self._loop, - tunnel_connection, - listen_host, - listen_port) + listener = await create_tcp_forward_listener( + self, self._loop, tunnel_connection, listen_host, listen_port, + tracker_factory) except OSError as exc: self.logger.debug1('Failed to create local TCP listener: %s', exc) raise @@ -5378,9 +5410,10 @@ async def tunnel_connection( return listener @async_context_manager - async def forward_local_path_to_port(self, listen_path: str, - dest_host: str, - dest_port: int) -> SSHListener: + async def forward_local_path_to_port( + self, listen_path: str, dest_host: str, dest_port: int, + tracker_factory: + Optional[SSHPathForwardTrackerFactory] = None) -> SSHListener: """Set up local UNIX domain socket forwarding to a remote TCP port This method is a coroutine which attempts to set up UNIX domain @@ -5395,9 +5428,15 @@ async def forward_local_path_to_port(self, listen_path: str, The hostname or address to forward the connections to :param dest_port: The port number to forward the connections to + :param tracker_factory: + An optional callable invoked once per accepted connection + which returns a new :class:`SSHPathForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_path: `str` :type dest_host: `str` :type dest_port: `int` + :type tracker_factory: :class:`SSHPathForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5417,9 +5456,9 @@ async def tunnel_connection( listen_path, (dest_host, dest_port)) try: - listener = await create_unix_forward_listener(self, self._loop, - tunnel_connection, - listen_path) + listener = await create_unix_forward_listener( + self, self._loop, tunnel_connection, listen_path, + tracker_factory) except OSError as exc: self.logger.debug1('Failed to create local UNIX listener: %s', exc) raise @@ -5429,9 +5468,10 @@ async def tunnel_connection( return listener @async_context_manager - async def forward_remote_port(self, listen_host: str, - listen_port: int, dest_host: str, - dest_port: int) -> SSHListener: + async def forward_remote_port( + self, listen_host: str, listen_port: int, dest_host: str, + dest_port: int, tracker_factory: Optional[ + SSHPortForwardTrackerFactory] = None) -> SSHListener: """Set up remote port forwarding This method is a coroutine which attempts to set up port @@ -5449,10 +5489,16 @@ async def forward_remote_port(self, listen_host: str, The hostname or address to forward connections to :param dest_port: The port number to forward connections to + :param tracker_factory: + An optional callable invoked once per forwarded connection + which returns a new :class:`SSHPortForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_host: `str` :type listen_port: `int` :type dest_host: `str` :type dest_port: `int` + :type tracker_factory: :class:`SSHPortForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5460,12 +5506,13 @@ async def forward_remote_port(self, listen_host: str, """ - def session_factory(_orig_host: str, - _orig_port: int) -> Awaitable[SSHTCPSession]: + def session_factory(orig_host: str, + orig_port: int) -> Awaitable[SSHTCPSession]: """Return an SSHTCPSession used to do remote port forwarding""" - return cast(Awaitable[SSHTCPSession], - self.forward_connection(dest_host, dest_port)) + return cast(Awaitable[SSHTCPSession], _track_remote_port( + self, tracker_factory, orig_host, orig_port, + self.forward_connection, dest_host, dest_port)) self.logger.info('Creating remote TCP forwarder from %s to %s', (listen_host, listen_port), (dest_host, dest_port)) @@ -5474,8 +5521,9 @@ def session_factory(_orig_host: str, listen_port) @async_context_manager - async def forward_remote_path(self, listen_path: str, - dest_path: str) -> SSHListener: + async def forward_remote_path( + self, listen_path: str, dest_path: str, tracker_factory: Optional[ + SSHPathForwardTrackerFactory] = None) -> SSHListener: """Set up remote UNIX domain socket forwarding This method is a coroutine which attempts to set up UNIX domain @@ -5489,8 +5537,14 @@ async def forward_remote_path(self, listen_path: str, The path on the remote host to listen on :param dest_path: The path on the local host to forward connections to + :param tracker_factory: + An optional callable invoked once per forwarded connection + which returns a new :class:`SSHPathForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_path: `str` :type dest_path: `str` + :type tracker_factory: :class:`SSHPathForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5501,8 +5555,8 @@ async def forward_remote_path(self, listen_path: str, def session_factory() -> Awaitable[SSHUNIXSession[bytes]]: """Return an SSHUNIXSession used to do remote path forwarding""" - return cast(Awaitable[SSHUNIXSession[bytes]], - self.forward_unix_connection(dest_path)) + return cast(Awaitable[SSHUNIXSession[bytes]], _track_remote_path( + self, tracker_factory, self.forward_unix_connection, dest_path)) self.logger.info('Creating remote UNIX forwarder from %s to %s', listen_path, dest_path) @@ -5510,9 +5564,10 @@ def session_factory() -> Awaitable[SSHUNIXSession[bytes]]: return await self.create_unix_server(session_factory, listen_path) @async_context_manager - async def forward_remote_port_to_path(self, listen_host: str, - listen_port: int, - dest_path: str) -> SSHListener: + async def forward_remote_port_to_path( + self, listen_host: str, listen_port: int, dest_path: str, + tracker_factory: Optional[ + SSHPortForwardTrackerFactory] = None) -> SSHListener: """Set up remote TCP port forwarding to a local UNIX domain socket This method is a coroutine which attempts to set up port @@ -5528,9 +5583,15 @@ async def forward_remote_port_to_path(self, listen_host: str, The port number on the remote host to listen on :param dest_path: The path on the local host to forward connections to + :param tracker_factory: + An optional callable invoked once per forwarded connection + which returns a new :class:`SSHPortForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_host: `str` :type listen_port: `int` :type dest_path: `str` + :type tracker_factory: :class:`SSHPortForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5538,12 +5599,13 @@ async def forward_remote_port_to_path(self, listen_host: str, """ - def session_factory(_orig_host: str, - _orig_port: int) -> Awaitable[SSHUNIXSession]: + def session_factory(orig_host: str, + orig_port: int) -> Awaitable[SSHUNIXSession]: """Return an SSHTCPSession used to do remote port forwarding""" - return cast(Awaitable[SSHUNIXSession], - self.forward_unix_connection(dest_path)) + return cast(Awaitable[SSHUNIXSession], _track_remote_port( + self, tracker_factory, orig_host, orig_port, + self.forward_unix_connection, dest_path)) self.logger.info('Creating remote TCP forwarder from %s to %s', (listen_host, listen_port), dest_path) @@ -5552,9 +5614,10 @@ def session_factory(_orig_host: str, listen_port) @async_context_manager - async def forward_remote_path_to_port(self, listen_path: str, - dest_host: str, - dest_port: int) -> SSHListener: + async def forward_remote_path_to_port( + self, listen_path: str, dest_host: str, dest_port: int, + tracker_factory: Optional[ + SSHPathForwardTrackerFactory] = None) -> SSHListener: """Set up remote UNIX domain socket forwarding to a local TCP port This method is a coroutine which attempts to set up UNIX domain @@ -5570,9 +5633,15 @@ async def forward_remote_path_to_port(self, listen_path: str, The hostname or address to forward connections to :param dest_port: The port number to forward connections to + :param tracker_factory: + An optional callable invoked once per forwarded connection + which returns a new :class:`SSHPathForwardTracker` for observing + that connection's lifecycle. `None` (default) disables tracking + with no overhead. :type listen_path: `str` :type dest_host: `str` :type dest_port: `int` + :type tracker_factory: :class:`SSHPathForwardTrackerFactory` :returns: :class:`SSHListener` @@ -5583,8 +5652,9 @@ async def forward_remote_path_to_port(self, listen_path: str, def session_factory() -> Awaitable[SSHTCPSession[bytes]]: """Return an SSHUNIXSession used to do remote path forwarding""" - return cast(Awaitable[SSHTCPSession[bytes]], - self.forward_connection(dest_host, dest_port)) + return cast(Awaitable[SSHTCPSession[bytes]], _track_remote_path( + self, tracker_factory, self.forward_connection, + dest_host, dest_port)) self.logger.info('Creating remote UNIX forwarder from %s to %s', listen_path, (dest_host, dest_port)) diff --git a/asyncssh/forward.py b/asyncssh/forward.py index 8470c000..9ee90629 100644 --- a/asyncssh/forward.py +++ b/asyncssh/forward.py @@ -23,10 +23,11 @@ import asyncio import socket from types import TracebackType -from typing import TYPE_CHECKING, Any, Awaitable, Callable, Dict, Optional -from typing import Type, cast +from typing import TYPE_CHECKING, Any, Awaitable, Callable, Dict, Generic +from typing import Optional, Type, TypeVar, cast from typing_extensions import Self +from .constants import OPEN_ADMINISTRATIVELY_PROHIBITED from .misc import ChannelOpenError, SockAddr @@ -38,6 +39,125 @@ SSHForwarderCoro = Callable[..., Awaitable] +class SSHForwardTracker: + """Base class for observing the lifecycle of a forwarded connection + + A tracker observes a single forwarded connection. A + `tracker_factory` passed to one of the + :meth:`forward_local_port() ` + or :meth:`forward_remote_port() + ` families of methods + is called once per forwarded connection and must return a new + tracker instance, on which asyncssh then calls the hooks below + for the life of that connection. + + All hooks run inside the asyncio event loop and **must not block** + (no I/O, no sleeps). They are pure observers: return values are + ignored and the forwarded data is never altered. Each hook has a + no-op default, so a subclass need only override the ones it cares + about. Exceptions raised by a hook are caught and discarded, so a + buggy tracker can never break forwarding. + + This base class defines the hooks shared by all forward types. + Use :class:`SSHPortForwardTracker` for forwards with a TCP + listener and :class:`SSHPathForwardTracker` for forwards with a + UNIX domain socket listener; they differ only in the signature + of `connection_made`. + + """ + + def connection_lost(self, exc: Optional[Exception]) -> None: + """Called when the forwarded connection has closed + + :param exc: + The exception which caused the connection to close, or + `None` if the connection closed cleanly. + :type exc: :class:`Exception` or `None` + + """ + + def forward_local_bytes(self, data: bytes) -> None: + """Called for data forwarded from the local side into the tunnel + + :param data: + A block of bytes received on the local connection and + about to be sent over the SSH connection. This is called + once per received block, not once per byte. + :type data: `bytes` + + """ + + def forward_remote_bytes(self, data: bytes) -> None: + """Called for data forwarded from the tunnel to the local side + + :param data: + A block of bytes received over the SSH connection and + about to be written to the local connection. This is + called once per received block, not once per byte. + :type data: `bytes` + + """ + + +class SSHPortForwardTracker(SSHForwardTracker): + """Tracker for forwards with a TCP listener + + Used with + :meth:`forward_local_port() `, + :meth:`forward_local_port_to_path() + `, + :meth:`forward_remote_port() ` + and :meth:`forward_remote_port_to_path() + `. + + """ + + def connection_made(self, forwarder: 'SSHForwarder', + orig_host: str, orig_port: int) -> None: + """Called when a new TCP connection is accepted on the listener + + :param forwarder: + The forwarder handling this connection. + :param orig_host: + The originating client host. + :param orig_port: + The originating client port. + :type forwarder: :class:`SSHForwarder` + :type orig_host: `str` + :type orig_port: `int` + + """ + + +class SSHPathForwardTracker(SSHForwardTracker): + """Tracker for forwards with a UNIX domain socket listener + + Used with + :meth:`forward_local_path() `, + :meth:`forward_local_path_to_port() + `, + :meth:`forward_remote_path() ` + and :meth:`forward_remote_path_to_port() + `. + + """ + + def connection_made(self, forwarder: 'SSHForwarder') -> None: + """Called when a new UNIX domain connection is accepted + + :param forwarder: + The forwarder handling this connection. + :type forwarder: :class:`SSHForwarder` + + """ + + +SSHPortForwardTrackerFactory = Callable[[], SSHPortForwardTracker] +SSHPathForwardTrackerFactory = Callable[[], SSHPathForwardTracker] + +_Tracker = TypeVar('_Tracker', bound=SSHForwardTracker) + + class SSHForwarder(asyncio.BaseProtocol): """SSH port forwarding connection handler""" @@ -189,13 +309,87 @@ def close(self) -> None: peer.close() -class SSHLocalForwarder(SSHForwarder): +class SSHLocalForwarder(SSHForwarder, Generic[_Tracker]): """Local forwarding connection handler""" - def __init__(self, conn: 'SSHConnection', coro: SSHForwarderCoro): + def __init__(self, conn: 'SSHConnection', coro: SSHForwarderCoro, + tracker_factory: Optional[Callable[[], _Tracker]] = None): super().__init__() self._conn = conn self._coro = coro + self._tracker: Optional[_Tracker] = None + self._create_tracker(tracker_factory) + + def _create_tracker( + self, tracker_factory: Optional[Callable[[], _Tracker]]) -> None: + """Instantiate this connection's tracker from the factory, if any""" + + if tracker_factory is None: + return + + try: + self._tracker = tracker_factory() + except Exception: # pylint: disable=broad-except + # A buggy factory must not break forwarding; + # self._tracker remains the __init__ default of None. + pass + + @staticmethod + def _notify_tracker(tracker: Optional[_Tracker], + notify: Callable[[_Tracker], None]) -> None: + """Invoke a tracker hook, swallowing exceptions from buggy trackers""" + + if tracker is not None: + try: + notify(tracker) + except Exception: # pylint: disable=broad-except + pass + + def data_received(self, data: bytes, + datatype: Optional[int] = None) -> None: + """Handle incoming data from the local transport""" + + def notify(tracker: _Tracker) -> None: + """Report locally forwarded bytes to the tracker""" + + tracker.forward_local_bytes(data) + + self._notify_tracker(self._tracker, notify) + + super().data_received(data, datatype) + + def write(self, data: bytes) -> None: + """Write tunnel data out to the local transport""" + + def notify(tracker: _Tracker) -> None: + """Report remotely forwarded bytes to the tracker""" + + tracker.forward_remote_bytes(data) + + self._notify_tracker(self._tracker, notify) + + super().write(data) + + def connection_lost(self, exc: Optional[Exception]) -> None: + """Handle a closed local connection + + This is also called manually from `_forward()` on a channel + open failure, so the local transport's eventual close fires + a second `connection_lost(None)` on the protocol. The tracker + reference is cleared on the first call so the hook fires + exactly once per connection. + """ + + tracker, self._tracker = self._tracker, None + + def notify(tracker: _Tracker) -> None: + """Report the closed connection to the tracker""" + + tracker.connection_lost(exc) + + self._notify_tracker(tracker, notify) + + super().connection_lost(exc) async def _forward(self, *args: object) -> None: """Begin local forwarding""" @@ -225,8 +419,38 @@ def forward(self, *args: object) -> None: self._conn.create_task(self._forward(*args)) + async def forward_remote(self, notify: Callable[[_Tracker], None], + *args: object) -> SSHForwarder: + """Open the local end of a remotely forwarded connection + + The tracker is notified once the local socket is connected, + or just before `connection_lost` if it can't be. Closing + the forwarder from `connection_made` rejects the channel. + + """ + + def peer_factory() -> SSHForwarder: + """Return this forwarder as the local socket's protocol""" + + return self + + try: + forwarder = await self._coro(*args, peer_factory) + except ChannelOpenError as exc: + self._notify_tracker(self._tracker, notify) + self.connection_lost(exc) + raise + + self._notify_tracker(self._tracker, notify) -class SSHLocalPortForwarder(SSHLocalForwarder): + if not self._transport: + raise ChannelOpenError(OPEN_ADMINISTRATIVELY_PROHIBITED, + 'Connection forwarding closed') + + return forwarder + + +class SSHLocalPortForwarder(SSHLocalForwarder[SSHPortForwardTracker]): """Local TCP port forwarding connection handler""" def connection_made(self, transport: asyncio.BaseTransport) -> None: @@ -238,15 +462,70 @@ def connection_made(self, transport: asyncio.BaseTransport) -> None: if peername: # pragma: no branch orig_host, orig_port = peername[:2] + else: # pragma: no cover + orig_host, orig_port = '', 0 + + def notify(tracker: SSHPortForwardTracker) -> None: + """Report the new connection to the tracker""" + + tracker.connection_made(self, orig_host, orig_port) + + self._notify_tracker(self._tracker, notify) self.forward(orig_host, orig_port) -class SSHLocalPathForwarder(SSHLocalForwarder): +class SSHLocalPathForwarder(SSHLocalForwarder[SSHPathForwardTracker]): """Local UNIX domain socket forwarding connection handler""" def connection_made(self, transport: asyncio.BaseTransport) -> None: """Handle a newly opened connection""" super().connection_made(transport) + + def notify(tracker: SSHPathForwardTracker) -> None: + """Report the new connection to the tracker""" + + tracker.connection_made(self) + + self._notify_tracker(self._tracker, notify) + self.forward() + + +def _track_remote_port(conn: 'SSHConnection', + tracker_factory: Optional[SSHPortForwardTrackerFactory], + orig_host: str, orig_port: int, coro: SSHForwarderCoro, + *args: object) -> Awaitable[SSHForwarder]: + """Forward a connection from a remote TCP listener, tracking it""" + + if tracker_factory is None: + return coro(*args) + + forwarder = SSHLocalForwarder(conn, coro, tracker_factory) + + def notify(tracker: SSHPortForwardTracker) -> None: + """Report the new connection to the tracker""" + + tracker.connection_made(forwarder, orig_host, orig_port) + + return forwarder.forward_remote(notify, *args) + + +def _track_remote_path(conn: 'SSHConnection', + tracker_factory: Optional[SSHPathForwardTrackerFactory], + coro: SSHForwarderCoro, + *args: object) -> Awaitable[SSHForwarder]: + """Forward a connection from a remote UNIX listener, tracking it""" + + if tracker_factory is None: + return coro(*args) + + forwarder = SSHLocalForwarder(conn, coro, tracker_factory) + + def notify(tracker: SSHPathForwardTracker) -> None: + """Report the new connection to the tracker""" + + tracker.connection_made(forwarder) + + return forwarder.forward_remote(notify, *args) diff --git a/asyncssh/listener.py b/asyncssh/listener.py index e9cc475b..c9d6e483 100644 --- a/asyncssh/listener.py +++ b/asyncssh/listener.py @@ -29,6 +29,7 @@ from typing_extensions import Self from .forward import SSHForwarderCoro +from .forward import SSHPortForwardTrackerFactory, SSHPathForwardTrackerFactory from .forward import SSHLocalPortForwarder, SSHLocalPathForwarder from .misc import HostPort, MaybeAwait from .session import SSHTCPSession, SSHUNIXSession @@ -345,14 +346,16 @@ async def create_tcp_local_listener( async def create_tcp_forward_listener(conn: 'SSHConnection', loop: asyncio.AbstractEventLoop, coro: SSHForwarderCoro, listen_host: str, - listen_port: int) -> \ - 'SSHForwardListener': + listen_port: int, + tracker_factory: + Optional[SSHPortForwardTrackerFactory] = + None) -> 'SSHForwardListener': """Create a listener to forward traffic from a local TCP port over SSH""" def protocol_factory() -> asyncio.BaseProtocol: """Start a port forwarder for each new local connection""" - return SSHLocalPortForwarder(conn, coro) + return SSHLocalPortForwarder(conn, coro, tracker_factory) return await create_tcp_local_listener(conn, loop, protocol_factory, listen_host, listen_port) @@ -361,14 +364,16 @@ def protocol_factory() -> asyncio.BaseProtocol: async def create_unix_forward_listener(conn: 'SSHConnection', loop: asyncio.AbstractEventLoop, coro: SSHForwarderCoro, - listen_path: str) -> \ - 'SSHForwardListener': + listen_path: str, + tracker_factory: + Optional[SSHPathForwardTrackerFactory] = + None) -> 'SSHForwardListener': """Create a listener to forward a local UNIX domain socket over SSH""" def protocol_factory() -> asyncio.BaseProtocol: """Start a path forwarder for each new local connection""" - return SSHLocalPathForwarder(conn, coro) + return SSHLocalPathForwarder(conn, coro, tracker_factory) server = await loop.create_unix_server(protocol_factory, listen_path) diff --git a/docs/api.rst b/docs/api.rst index 046245b3..a19cf24a 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -90,6 +90,10 @@ to a UNIX domain socket or vice-versa can be set up using the functions `. In these cases, data transfer on the channels is managed automatically by AsyncSSH whenever new connections are opened, so custom session objects are not required. +All of these methods also accept an optional ``tracker_factory`` which +creates a per-connection :class:`SSHPortForwardTracker` or +:class:`SSHPathForwardTracker` to observe each forwarded connection, as +described in `Forward Tracker Classes`_. Dynamic TCP port forwarding can be set up by calling :meth:`forward_socks() `. The SOCKS listener set up by @@ -1009,6 +1013,91 @@ Forwarder Classes ============================== = +Forward Tracker Classes +======================= + +The ``forward_local_*`` and ``forward_remote_*`` methods on +:class:`SSHClientConnection` accept an optional ``tracker_factory`` +argument: a zero-argument callable invoked once per forwarded connection +which returns a tracker instance -- +:class:`SSHPortForwardTracker` for TCP listeners or +:class:`SSHPathForwardTracker` for UNIX domain listeners. asyncssh then +calls that instance's hooks for the life of the connection, giving +applications a passive view of per-connection lifecycle and byte flow -- +useful for idle-based auto-shutdown, connection counting, or traffic +metrics. + +The hooks are pure observers: they run inside the asyncio event loop, +must not block, and never alter the forwarded data (return values are +ignored). Every hook has a no-op default, so a subclass overrides only +what it needs, and exceptions raised by a hook or factory are caught and +discarded so a buggy tracker cannot break forwarding. + +Use :class:`SSHPortForwardTracker` with the TCP-listener methods +(:meth:`forward_local_port() `, +:meth:`forward_local_port_to_path() +`, +:meth:`forward_remote_port() ` and +:meth:`forward_remote_port_to_path() +`) and +:class:`SSHPathForwardTracker` with the UNIX-domain-listener methods +(:meth:`forward_local_path() `, +:meth:`forward_local_path_to_port() +`, +:meth:`forward_remote_path() ` and +:meth:`forward_remote_path_to_port() +`). The tracker type +follows the listener, not the destination. The two classes share the +same set of hooks and differ only in the signature of ``connection_made``. + +In both directions, ``forward_local_bytes`` reports bytes generated on +the local host and ``forward_remote_bytes`` reports bytes received over +the SSH connection. For a remote forward, ``connection_made`` is called +once the local destination is connected. If it can't be opened, +``connection_made`` is still called, immediately followed by +``connection_lost`` with the :exc:`ChannelOpenError`. Closing the +forwarder from ``connection_made`` rejects the forwarded connection. + + .. code-block:: python + + class ConnCounter(asyncssh.SSHPortForwardTracker): + def __init__(self, counter): + self._counter = counter + + def connection_made(self, forwarder, orig_host, orig_port): + self._counter.active += 1 + + def connection_lost(self, exc): + self._counter.active -= 1 + + def tracker_factory(): + return ConnCounter(counter) + + listener = await conn.forward_local_port( + '', 0, 'remote-host', 80, tracker_factory=tracker_factory) + + remote_listener = await conn.forward_remote_port( + '', 8080, 'localhost', 80, tracker_factory=tracker_factory) + +.. autoclass:: SSHPortForwardTracker() + + ==================================== = + .. automethod:: connection_made + .. automethod:: connection_lost + .. automethod:: forward_local_bytes + .. automethod:: forward_remote_bytes + ==================================== = + +.. autoclass:: SSHPathForwardTracker() + + ==================================== = + .. automethod:: connection_made + .. automethod:: connection_lost + .. automethod:: forward_local_bytes + .. automethod:: forward_remote_bytes + ==================================== = + + Listener Classes ================ diff --git a/tests/test_forward.py b/tests/test_forward.py index dbfd792e..07e09147 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -30,6 +30,7 @@ from unittest.mock import patch import asyncssh +from asyncssh.constants import OPEN_CONNECT_FAILED from asyncssh.misc import maybe_wait_closed, write_file from asyncssh.packet import String, UInt32 from asyncssh.public_key import CERT_TYPE_USER @@ -85,6 +86,166 @@ async def _async_runtime_error(_reader, _writer): raise RuntimeError('Async internal error') + +_REQUEST = b'request\n' +_RESPONSE = b'a distinctly different response\n' + + +async def _reply(reader, writer): + """Answer a request with a response which is not an echo of it""" + + await reader.readline() + + writer.write(_RESPONSE) + await writer.drain() + + writer.close() + await maybe_wait_closed(writer) + + +class _Recorder: + """Mixin which records a tracker's lifecycle and forwarded bytes""" + + def __init__(self): + self.events = [] + self.local = b'' + self.remote = b'' + self.lost = asyncio.Event() + + def connection_lost(self, exc): + """Record the connection closing""" + + self.events.append(('lost', exc)) + self.lost.set() + + def forward_local_bytes(self, data): + """Record bytes generated on the local host""" + + self.local += data + + def forward_remote_bytes(self, data): + """Record bytes received over SSH""" + + self.remote += data + + def kinds(self): + """Return the kinds of events recorded, in order""" + + return [event[0] for event in self.events] + + +class _PortRecorder(_Recorder, asyncssh.SSHPortForwardTracker): + """Port tracker which records its lifecycle and forwarded bytes""" + + def connection_made(self, forwarder, orig_host, orig_port): + """Record the new connection""" + + self.events.append(('made', forwarder, orig_host, orig_port)) + + +class _PathRecorder(_Recorder, asyncssh.SSHPathForwardTracker): + """Path tracker which records its lifecycle and forwarded bytes""" + + def connection_made(self, forwarder): + """Record the new connection""" + + self.events.append(('made', forwarder)) + + +class _PortCloser(_PortRecorder): + """Port tracker which closes the forwarder as soon as it's made""" + + def connection_made(self, forwarder, orig_host, orig_port): + """Record the new connection and close it""" + + super().connection_made(forwarder, orig_host, orig_port) + forwarder.close() + + +class _PathCloser(_PathRecorder): + """Path tracker which closes the forwarder as soon as it's made""" + + def connection_made(self, forwarder): + """Record the new connection and close it""" + + super().connection_made(forwarder) + forwarder.close() + + +class _Buggy: + """Mixin whose byte and lost hooks record their use and then raise""" + + # pylint: disable=unused-argument + + hooks = set() + + def connection_lost(self, exc): + """Record the hook and raise""" + + self.hooks.add('connection_lost') + raise RuntimeError('lost boom') + + def forward_local_bytes(self, data): + """Record the hook and raise""" + + self.hooks.add('forward_local_bytes') + raise RuntimeError('local boom') + + def forward_remote_bytes(self, data): + """Record the hook and raise""" + + self.hooks.add('forward_remote_bytes') + raise RuntimeError('remote boom') + + +class _BuggyPort(_Buggy, asyncssh.SSHPortForwardTracker): + """Port tracker whose hooks all raise""" + + def connection_made(self, forwarder, orig_host, orig_port): + """Record the hook and raise""" + + self.hooks.add('connection_made') + raise RuntimeError('made boom') + + +class _BuggyPath(_Buggy, asyncssh.SSHPathForwardTracker): + """Path tracker whose hooks all raise""" + + def connection_made(self, forwarder): + """Record the hook and raise""" + + self.hooks.add('connection_made') + raise RuntimeError('made boom') + + +def _recording(cls): + """Return a list of trackers and a factory which adds to it""" + + trackers = [] + + def factory(): + """Create and record a new tracker""" + + tracker = cls() + trackers.append(tracker) + return tracker + + return trackers, factory + + +def _broken_factory(): + """Fail to return a tracker""" + + raise RuntimeError('factory boom') + + +def _closed_port(): + """Return a local TCP port with nothing listening on it""" + + with socket.socket() as sock: + sock.bind(('127.0.0.1', 0)) + return sock.getsockname()[1] + class _ClientConn(asyncssh.SSHClientConnection): """Patched SSH client connection for unit testing""" @@ -307,6 +468,40 @@ async def _check_local_unix_connection(self, listen_path): await self._check_echo_line(reader, writer) + async def _check_reply(self, reader, writer): + """Send a request and check the distinct reply to it""" + + writer.write(_REQUEST) + await writer.drain() + + self.assertEqual((await reader.readline()), _RESPONSE) + + writer.close() + await maybe_wait_closed(writer) + + async def _check_closed(self, reader, writer): + """Check that a connection is closed without any data""" + + self.assertEqual((await asyncio.wait_for(reader.read(), 1)), b'') + + writer.close() + await maybe_wait_closed(writer) + + def _check_made_lost(self, trackers, exc=None): + """Check that one tracker was made and then lost once""" + + self.assertEqual(len(trackers), 1) + self.assertEqual(trackers[0].kinds(), ['made', 'lost']) + self.assertIsInstance(trackers[0].events[0][1], + asyncssh.SSHForwarder) + + if exc is None: + self.assertIsNone(trackers[0].events[1][1]) + else: + self.assertIsInstance(trackers[0].events[1][1], exc) + + return trackers[0] + class _TestTCPForwarding(_CheckForwarding): """Unit tests for AsyncSSH TCP connection forwarding""" @@ -697,6 +892,173 @@ async def accept_handler(_orig_host: str, _orig_port: int) -> bool: writer.close() await maybe_wait_closed(writer) + @asynctest + async def test_port_tracker_made_and_lost(self): + """A port tracker sees connection_made and connection_lost""" + + events = [] + lost = asyncio.Event() + + class _RecordingTracker(asyncssh.SSHPortForwardTracker): + """Tracker which records connection_made and connection_lost""" + + def connection_made(self, forwarder, orig_host, orig_port): + events.append(('made', forwarder, orig_host, orig_port)) + + def connection_lost(self, exc): + events.append(('lost', exc)) + lost.set() + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, + tracker_factory=_RecordingTracker) as listener: + await self._check_local_connection(listener.get_port()) + await asyncio.wait_for(lost.wait(), timeout=1.0) + + kinds = [event[0] for event in events] + self.assertIn('made', kinds) + self.assertIn('lost', kinds) + + made = next(event for event in events if event[0] == 'made') + self.assertIsInstance(made[1], asyncssh.SSHForwarder) + self.assertEqual(made[2], '127.0.0.1') + self.assertIsInstance(made[3], int) + + @asynctest + async def test_tracker_factory_per_connection(self): + """A distinct tracker instance is created for each accepted + connection, and each instance's connection_lost fires once""" + + trackers = [] + lost_events = [] + + class _CountingTracker(asyncssh.SSHPortForwardTracker): + """Tracker which records each instance created by the factory""" + + def __init__(self): + trackers.append(self) + lost_events.append(asyncio.Event()) + + def connection_lost(self, exc): + lost_events[trackers.index(self)].set() + + def factory(): + return _CountingTracker() + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, tracker_factory=factory) as listener: + listen_port = listener.get_port() + await self._check_local_connection(listen_port) + await self._check_local_connection(listen_port) + await asyncio.wait_for( + asyncio.gather(*(e.wait() for e in lost_events)), + timeout=1.0) + + self.assertEqual(len(trackers), 2) + + @asynctest + async def test_port_tracker_byte_hooks(self): + """Byte hooks observe both forwarding directions""" + + local_bytes = bytearray() + remote_bytes = bytearray() + lost = asyncio.Event() + + class _ByteTracker(asyncssh.SSHPortForwardTracker): + """Tracker which records bytes seen in both forwarding directions""" + + def forward_local_bytes(self, data): + local_bytes.extend(data) + + def forward_remote_bytes(self, data): + remote_bytes.extend(data) + + def connection_lost(self, exc): + lost.set() + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, + tracker_factory=_ByteTracker) as listener: + await self._check_local_connection(listener.get_port()) + await asyncio.wait_for(lost.wait(), timeout=1.0) + + line = (str(id(self)) + '\n').encode('utf-8') + self.assertEqual(bytes(local_bytes), line) + self.assertEqual(bytes(remote_bytes), line) + + @asynctest + async def test_port_tracker_factory_exception_swallowed(self): + """A factory that raises does not break forwarding""" + + def factory(): + raise RuntimeError('factory boom') + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, tracker_factory=factory) as listener: + await self._check_local_connection(listener.get_port(), + delay=0.1) + + @asynctest + async def test_port_tracker_hook_exception_swallowed(self): + """A tracker whose hooks raise does not break forwarding""" + + class _BuggyTracker(asyncssh.SSHPortForwardTracker): + """Tracker whose hooks all raise, to verify they're swallowed""" + + def connection_made(self, forwarder, orig_host, orig_port): + raise RuntimeError('made boom') + + def connection_lost(self, exc): + raise RuntimeError('lost boom') + + def forward_local_bytes(self, data): + raise RuntimeError('local boom') + + def forward_remote_bytes(self, data): + raise RuntimeError('remote boom') + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, + tracker_factory=_BuggyTracker) as listener: + await self._check_local_connection(listener.get_port(), + delay=0.1) + + @asynctest + async def test_port_tracker_lost_fires_once(self): + """connection_lost fires once even when ChannelOpenError triggers + a manual notify in _forward() followed by the asyncio close path.""" + + lost_count = 0 + + class _Counting(asyncssh.SSHPortForwardTracker): + """Tracker which counts how many times connection_lost fires""" + + def connection_lost(self, exc): + nonlocal lost_count + lost_count += 1 + + async def deny(_orig_host, _orig_port): + return False + + async with self.connect() as conn: + async with conn.forward_local_port( + '', 0, '', 7, + accept_handler=deny, + tracker_factory=_Counting) as listener: + reader, writer = await asyncio.open_connection( + '127.0.0.1', listener.get_port()) + self.assertEqual((await reader.read()), b'') + writer.close() + await maybe_wait_closed(writer) + await asyncio.sleep(0.1) # bounded: upper-bound for any spurious duplicate + + self.assertEqual(lost_count, 1) + @unittest.skipIf(sys.platform == 'win32', 'skip UNIX domain socket tests on Windows') @asynctest @@ -857,6 +1219,146 @@ async def test_forward_remote_port_to_path(self): try_remove('local') + async def _remote_port(self, handler, factory, client, dest_port=None): + """Run a client over a tracked remote port forward to a handler""" + + server = await asyncio.start_server(handler, '127.0.0.1', 0) + + if dest_port is None: + dest_port = server.sockets[0].getsockname()[1] + + async with self.connect() as conn: + async with conn.forward_remote_port( + '', 0, '127.0.0.1', dest_port, factory) as listener: + reader, writer = await asyncio.open_connection( + '127.0.0.1', listener.get_port()) + await client(reader, writer) + await asyncio.sleep(0.1) + + server.close() + await server.wait_closed() + + @asynctest + async def test_remote_port_tracker(self): + """Test a tracker on a remote port forward""" + + trackers, factory = _recording(_PortRecorder) + + await self._remote_port(echo, factory, self._check_echo_line) + + tracker = self._check_made_lost(trackers) + self.assertEqual(tracker.events[0][2], '127.0.0.1') + self.assertIsInstance(tracker.events[0][3], int) + + @asynctest + async def test_remote_port_bytes(self): + """Test tracker byte hooks on a remote port forward""" + + trackers, factory = _recording(_PortRecorder) + + await self._remote_port(_reply, factory, self._check_reply) + + self.assertEqual(trackers[0].remote, _REQUEST) + self.assertEqual(trackers[0].local, _RESPONSE) + + @asynctest + async def test_remote_port_refused(self): + """Test a tracker on a remote port forward to a closed port""" + + trackers, factory = _recording(_PortRecorder) + + await self._remote_port(echo, factory, self._check_closed, + _closed_port()) + + tracker = self._check_made_lost(trackers, asyncssh.ChannelOpenError) + self.assertEqual(tracker.events[1][1].code, OPEN_CONNECT_FAILED) + + @asynctest + async def test_remote_port_refused_recovery(self): + """Test that a refused remote forward leaves the connection usable""" + + trackers, factory = _recording(_PortRecorder) + + async with self.connect() as conn: + async with conn.forward_remote_port( + '', 0, '127.0.0.1', _closed_port(), factory) as listener: + reader, writer = await asyncio.open_connection( + '127.0.0.1', listener.get_port()) + + await self._check_closed(reader, writer) + + self._check_made_lost(trackers, asyncssh.ChannelOpenError) + + await self._check_connection(conn) + + @asynctest + async def test_remote_port_close(self): + """Test closing a remote port forward from connection_made""" + + dest_closed = asyncio.Event() + + async def wait_eof(reader, writer): + """Wait for the forwarder to close the destination""" + + await reader.read() + dest_closed.set() + writer.close() + + trackers, factory = _recording(_PortCloser) + + await self._remote_port(wait_eof, factory, self._check_closed) + await asyncio.wait_for(dest_closed.wait(), 1) + + self._check_made_lost(trackers) + + @asynctest + async def test_remote_port_buggy(self): + """Test a remote port tracker whose hooks all raise""" + + _BuggyPort.hooks = set() + + await self._remote_port(_reply, _BuggyPort, self._check_reply) + + self.assertEqual(_BuggyPort.hooks, + {'connection_made', 'connection_lost', + 'forward_local_bytes', 'forward_remote_bytes'}) + + @asynctest + async def test_remote_port_factory_error(self): + """Test a remote port tracker factory which raises""" + + await self._remote_port(echo, _broken_factory, self._check_echo_line) + + @unittest.skipIf(sys.platform == 'win32', + 'skip UNIX domain socket tests on Windows') + @asynctest + async def test_remote_port_to_path_tracker(self): + """Test a tracker on a remote port forward to a UNIX socket""" + + trackers, factory = _recording(_PortRecorder) + + # pylint: disable=no-member + server = await asyncio.start_unix_server(_reply, 'local') + # pylint: enable=no-member + + async with self.connect() as conn: + async with conn.forward_remote_port_to_path( + '', 0, 'local', factory) as listener: + reader, writer = await asyncio.open_connection( + '127.0.0.1', listener.get_port()) + await self._check_reply(reader, writer) + await asyncio.wait_for(trackers[0].lost.wait(), 1) + + server.close() + await server.wait_closed() + + try_remove('local') + + tracker = self._check_made_lost(trackers) + self.assertEqual(tracker.events[0][2], '127.0.0.1') + self.assertEqual(tracker.remote, _REQUEST) + self.assertEqual(tracker.local, _RESPONSE) + @asynctest async def test_forward_remote_specific_port(self): """Test forwarding of a specific remote port""" @@ -1149,6 +1651,59 @@ async def test_forward_local_path(self): try_remove('local') + @asynctest + async def test_path_tracker_made_and_lost(self): + """A path tracker sees connection_made (no addr) and connection_lost""" + + events = [] + lost = asyncio.Event() + + class _RecordingTracker(asyncssh.SSHPathForwardTracker): + """Tracker which records connection_made and connection_lost""" + + def connection_made(self, forwarder): + events.append(('made', forwarder)) + + def connection_lost(self, exc): + events.append(('lost', exc)) + lost.set() + + async with self.connect() as conn: + async with conn.forward_local_path( + 'local', '/echo', + tracker_factory=_RecordingTracker): + await self._check_local_unix_connection('local') + await asyncio.wait_for(lost.wait(), timeout=1.0) + + try_remove('local') + + kinds = [event[0] for event in events] + self.assertIn('made', kinds) + self.assertIn('lost', kinds) + + made = next(event for event in events if event[0] == 'made') + self.assertIsInstance(made[1], asyncssh.SSHForwarder) + + @asynctest + async def test_path_tracker_hook_exception_swallowed(self): + """A path tracker whose connection_made raises does not break + forwarding""" + + class _BuggyTracker(asyncssh.SSHPathForwardTracker): + """Tracker whose connection_made hook raises, to verify it's + swallowed""" + + def connection_made(self, forwarder): + raise RuntimeError('made boom') + + async with self.connect() as conn: + async with conn.forward_local_path( + 'local', '/echo', + tracker_factory=_BuggyTracker): + await self._check_local_unix_connection('local') + + try_remove('local') + @asynctest async def test_forward_local_port_to_path_accept_handler(self): """Test forwarding of port to UNIX path with accept handler""" @@ -1248,6 +1803,137 @@ async def test_forward_remote_path_to_port(self): try_remove('echo') + async def _remote_path(self, dest, factory, client): + """Run a client over a tracked remote path forward to dest""" + + path = os.path.abspath('echo') + + async with self.connect() as conn: + if isinstance(dest, int): + listener = await conn.forward_remote_path_to_port( + path, '127.0.0.1', dest, factory) + else: + listener = await conn.forward_remote_path(path, dest, factory) + + async with listener: + # pylint: disable=no-member + reader, writer = await asyncio.open_unix_connection('echo') + # pylint: enable=no-member + + await client(reader, writer) + await asyncio.sleep(0.1) + + try_remove('echo') + + async def _remote_unix_path(self, handler, factory, client): + """Run a client over a tracked remote path forward to a handler""" + + # pylint: disable=no-member + server = await asyncio.start_unix_server(handler, 'local') + # pylint: enable=no-member + + await self._remote_path('local', factory, client) + + server.close() + await server.wait_closed() + + try_remove('local') + + @asynctest + async def test_remote_path_tracker(self): + """Test a tracker on a remote path forward""" + + trackers, factory = _recording(_PathRecorder) + + await self._remote_unix_path(echo, factory, self._check_echo_line) + + tracker = self._check_made_lost(trackers) + self.assertEqual(len(tracker.events[0]), 2) + + @asynctest + async def test_remote_path_bytes(self): + """Test tracker byte hooks on a remote path forward""" + + trackers, factory = _recording(_PathRecorder) + + await self._remote_unix_path(_reply, factory, self._check_reply) + + self.assertEqual(trackers[0].remote, _REQUEST) + self.assertEqual(trackers[0].local, _RESPONSE) + + @asynctest + async def test_remote_path_refused(self): + """Test a tracker on a remote path forward to a missing path""" + + trackers, factory = _recording(_PathRecorder) + + try_remove('missing') + + await self._remote_path('missing', factory, self._check_closed) + + tracker = self._check_made_lost(trackers, asyncssh.ChannelOpenError) + self.assertEqual(tracker.events[1][1].code, OPEN_CONNECT_FAILED) + + @asynctest + async def test_remote_path_close(self): + """Test closing a remote path forward from connection_made""" + + dest_closed = asyncio.Event() + + async def wait_eof(reader, writer): + """Wait for the forwarder to close the destination""" + + await reader.read() + dest_closed.set() + writer.close() + + trackers, factory = _recording(_PathCloser) + + await self._remote_unix_path(wait_eof, factory, self._check_closed) + await asyncio.wait_for(dest_closed.wait(), 1) + + self._check_made_lost(trackers) + + @asynctest + async def test_remote_path_buggy(self): + """Test a remote path tracker whose hooks all raise""" + + _BuggyPath.hooks = set() + + await self._remote_unix_path(_reply, _BuggyPath, self._check_reply) + + self.assertEqual(_BuggyPath.hooks, + {'connection_made', 'connection_lost', + 'forward_local_bytes', 'forward_remote_bytes'}) + + @asynctest + async def test_remote_path_to_port_tracker(self): + """Test a tracker on a remote path forward to a TCP port""" + + trackers, factory = _recording(_PathRecorder) + + server = await asyncio.start_server(_reply, '127.0.0.1', 0) + server_port = server.sockets[0].getsockname()[1] + + await self._remote_path(server_port, factory, self._check_reply) + + server.close() + await server.wait_closed() + + tracker = self._check_made_lost(trackers) + self.assertEqual(tracker.remote, _REQUEST) + self.assertEqual(tracker.local, _RESPONSE) + + @asynctest + async def test_remote_path_to_port_refused(self): + """Test a tracker on a remote path forward to a closed port""" + + trackers, factory = _recording(_PathRecorder) + + await self._remote_path(_closed_port(), factory, self._check_closed) + + self._check_made_lost(trackers, asyncssh.ChannelOpenError) + @asynctest async def test_forward_remote_path_failure(self): """Test failure of forwarding a remote UNIX domain path"""