diff --git a/Changelog.rst b/Changelog.rst index 4de37c94..5e957ebf 100644 --- a/Changelog.rst +++ b/Changelog.rst @@ -7,6 +7,9 @@ Change Log Changes -------- +* Added a native ``SFTPClient`` API with public remote directory, metadata, + mutation and transfer operations, plus remote current working directory + support. * All local file operations now use a thread pool to improve local file I/O performance. This includes loading private key files from a local file path, identity authentication using local files as well as SFTP read/write operations on local files. diff --git a/ci/integration_tests/libssh2_clients/test_sftp_client.py b/ci/integration_tests/libssh2_clients/test_sftp_client.py new file mode 100644 index 00000000..64933325 --- /dev/null +++ b/ci/integration_tests/libssh2_clients/test_sftp_client.py @@ -0,0 +1,43 @@ +# This file is part of parallel-ssh. +# Copyright (C) 2014-2026 Panos Kittenis. +# Copyright (C) 2014-2026 parallel-ssh Contributors. +# +# This library is free software; you can redistribute it and/or +# modify it under the terms of the GNU Lesser General Public +# License as published by the Free Software Foundation, version 2.1. + +import os +import shutil +import tempfile + +from .base_ssh2_case import SSH2TestCase + + +class SFTPClientTest(SSH2TestCase): + + def test_cwd_directory_and_transfer_operations(self): + remote_root = tempfile.mkdtemp(prefix='parallel-ssh-sftp-') + local_root = tempfile.mkdtemp(prefix='parallel-ssh-local-') + local_source = os.path.join(local_root, 'source.txt') + local_copy = os.path.join(local_root, 'copy.txt') + try: + with open(local_source, 'w') as handle: + handle.write('parallel-ssh') + sftp = self.client.open_sftp() + self.assertTrue(sftp.getcwd().startswith('/')) + sftp.chdir(remote_root) + self.assertEqual(sftp.getcwd(), os.path.realpath(remote_root)) + sftp.mkdir('nested') + self.assertIn('nested', sftp.listdir('.')) + sftp.put(local_source, 'nested/remote.txt') + self.assertIn('remote.txt', sftp.listdir('nested')) + sftp.get('nested/remote.txt', local_copy) + with open(local_copy) as handle: + self.assertEqual(handle.read(), 'parallel-ssh') + sftp.rename('nested/remote.txt', 'nested/renamed.txt') + sftp.remove('nested/renamed.txt') + sftp.rmdir('nested') + self.assertNotIn('nested', sftp.listdir('.')) + finally: + shutil.rmtree(remote_root, ignore_errors=True) + shutil.rmtree(local_root, ignore_errors=True) diff --git a/doc/api.rst b/doc/api.rst index 7f45c00a..3c710fa3 100644 --- a/doc/api.rst +++ b/doc/api.rst @@ -6,6 +6,7 @@ API Documentation native_parallel native_single + native_sftp ssh_parallel ssh_single base_parallel diff --git a/doc/native_sftp.rst b/doc/native_sftp.rst new file mode 100644 index 00000000..f2a38f45 --- /dev/null +++ b/doc/native_sftp.rst @@ -0,0 +1,27 @@ +Native SFTP Client +================== + +The native client can open a user-facing SFTP client that owns one reusable +SFTP channel and tracks a remote current working directory. + +.. code-block:: python + + from pssh.clients import SSHClient + + client = SSHClient('localhost') + sftp = client.open_sftp() + sftp.chdir('/srv/uploads') + sftp.mkdir('incoming') + sftp.put('local.txt', 'incoming/remote.txt') + print(sftp.listdir('incoming')) + sftp.get('incoming/remote.txt', 'downloaded.txt') + +Relative remote paths are resolved against ``sftp.getcwd()`` using POSIX path +semantics. The SFTP client is bound to its parent ``SSHClient`` connection. +This API is available for the native ``ssh2-python`` client only; the +``pssh.clients.ssh`` backend does not currently support SFTP. + +.. automodule:: pssh.clients.native.sftp + :members: + :undoc-members: + :member-order: groupwise diff --git a/pssh/clients/native/__init__.py b/pssh/clients/native/__init__.py index 5e5f19ad..6f9ae0c8 100644 --- a/pssh/clients/native/__init__.py +++ b/pssh/clients/native/__init__.py @@ -18,3 +18,4 @@ # flake8: noqa: F401 from .parallel import ParallelSSHClient from .single import SSHClient, logger +from .sftp import SFTPClient diff --git a/pssh/clients/native/sftp.py b/pssh/clients/native/sftp.py new file mode 100644 index 00000000..abf76fe1 --- /dev/null +++ b/pssh/clients/native/sftp.py @@ -0,0 +1,113 @@ +# This file is part of parallel-ssh. +# Copyright (C) 2014-2026 Panos Kittenis. +# Copyright (C) 2014-2026 parallel-ssh Contributors. +# +# This library is free software; you can redistribute it and/or +# modify it under the terms of the GNU Lesser General Public +# License as published by the Free Software Foundation, version 2.1. + +import posixpath + + +class SFTPClient(object): + """User-facing SFTP operations bound to one native SSH client.""" + + __slots__ = ('_client', '_sftp', '_cwd') + + def __init__(self, client, sftp=None): + self._client = client + self._sftp = client.make_sftp_client() if sftp is None else sftp + self._cwd = self._canonical_path('.') + + def _canonical_path(self, path): + return self._client.eagain(self._sftp.realpath, path) + + def _remote_path(self, path): + if not isinstance(path, str): + raise TypeError("Remote path must be a string.") + if not path: + return self._cwd + if posixpath.isabs(path): + return posixpath.normpath(path) + return posixpath.normpath(posixpath.join(self._cwd, path)) + + def getcwd(self): + """Get the current remote working directory.""" + return self._cwd + + def chdir(self, path): + """Change the current remote working directory.""" + target = self._canonical_path(self._remote_path(path)) + with self._client._sftp_openfh(self._sftp.opendir, target): + pass + self._cwd = target + return self._cwd + + def listdir(self, path='.', encoding='utf-8'): + """List names in a remote directory.""" + with self._client._sftp_openfh( + self._sftp.opendir, self._remote_path(path)) as dir_h: + for entry in self._client._sftp_readdir(dir_h): + name = entry.decode(encoding) + if name not in ('.', '..'): + yield name + + def stat(self, path): + """Return attributes for a remote path, following symbolic links.""" + return self._client.eagain(self._sftp.stat, self._remote_path(path)) + + def lstat(self, path): + """Return attributes for a remote path without following links.""" + return self._client.eagain(self._sftp.lstat, self._remote_path(path)) + + def mkdir(self, path): + """Create a remote directory and missing parent directories.""" + return self._client.mkdir(self._sftp, self._remote_path(path)) + + def rmdir(self, path): + """Remove an empty remote directory.""" + return self._client.eagain(self._sftp.rmdir, self._remote_path(path)) + + def rename(self, source, destination): + """Rename a remote path.""" + return self._client.eagain( + self._sftp.rename, + self._remote_path(source), + self._remote_path(destination), + ) + + def remove(self, path): + """Remove a remote file.""" + return self._client.eagain(self._sftp.unlink, self._remote_path(path)) + + unlink = remove + + def get(self, remote_file, local_file): + """Copy one remote file to a local path.""" + return self._client.sftp_get( + self._sftp, self._remote_path(remote_file), local_file) + + def put(self, local_file, remote_file): + """Copy one local file to a remote path.""" + return self._client.sftp_put( + self._sftp, local_file, self._remote_path(remote_file)) + + def copy_file(self, local_file, remote_file, recurse=False): + """Copy a local file or directory to a remote path.""" + return self._client.copy_file( + local_file, + self._remote_path(remote_file), + recurse=recurse, + sftp=self._sftp, + ) + + def copy_remote_file(self, remote_file, local_file, recurse=False, + encoding='utf-8'): + """Copy a remote file or directory to a local path.""" + return self._client.copy_remote_file( + self._remote_path(remote_file), + local_file, + recurse=recurse, + sftp=self._sftp, + encoding=encoding, + ) diff --git a/pssh/clients/native/single.py b/pssh/clients/native/single.py index ee1d9855..92397f46 100644 --- a/pssh/clients/native/single.py +++ b/pssh/clients/native/single.py @@ -35,6 +35,7 @@ LIBSSH2_SFTP_S_IXGRP, LIBSSH2_SFTP_S_IXOTH from .tunnel import FORWARDER +from .sftp import SFTPClient from ..base.single import BaseSSHClient, PollMixIn from ...constants import DEFAULT_RETRIES, RETRY_DELAY from ...exceptions import SessionError, SFTPError, \ @@ -94,6 +95,8 @@ class SSHClient(BaseSSHClient): """ssh2-python (libssh2) based non-blocking SSH client.""" # 2MB buffer _BUF_SIZE = 2048 * 1024 + # Keep SCP receives in bounded chunks instead of requesting the whole file. + _SCP_RECV_BUF_SIZE = 64 * 1024 def __init__(self, host, user=None, password=None, port=None, @@ -434,13 +437,19 @@ def eagain(self, func, *args, **kwargs): def _make_sftp_eagain(self): return self.eagain(self.session.sftp_init) - def _make_sftp(self): + def make_sftp_client(self): try: sftp = self._make_sftp_eagain() except Exception as ex: raise SFTPError(ex) return sftp + _make_sftp = make_sftp_client + + def open_sftp(self): + """Open a user-facing SFTP client bound to this SSH session.""" + return SFTPClient(self) + def _mkdir(self, sftp, directory): """Make directory via SFTP channel. @@ -488,7 +497,7 @@ def copy_file(self, local_file, remote_file, recurse=False, sftp=None): :raises: :py:class:`IOError` on local file IO errors :raises: :py:class:`OSError` on local OS errors like permission denied """ - sftp = self._make_sftp() if sftp is None else sftp + sftp = self.make_sftp_client() if sftp is None else sftp if os.path.isdir(local_file) and recurse: return self._copy_dir(local_file, remote_file, sftp) elif os.path.isdir(local_file) and not recurse: @@ -590,7 +599,7 @@ def copy_remote_file(self, remote_file, local_file, recurse=False, :raises: :py:class:`IOError` on local file IO errors :raises: :py:class:`OSError` on local OS errors like permission denied """ - sftp = self._make_sftp() if sftp is None else sftp + sftp = self.make_sftp_client() if sftp is None else sftp try: self.eagain(sftp.stat, remote_file) except (SFTPHandleError, SFTPProtocolError): @@ -665,7 +674,7 @@ def scp_recv(self, remote_file, local_file, recurse=False, sftp=None, :raises: :py:class:`OSError` on local OS errors like permission denied. """ if recurse: - sftp = self._make_sftp() if sftp is None else sftp + sftp = self.make_sftp_client() if sftp is None else sftp return self._scp_recv_recursive(remote_file, local_file, sftp, encoding=encoding) elif local_file.endswith('/'): remote_filename = remote_file.rsplit('/')[-1] @@ -690,10 +699,16 @@ def _scp_recv(self, remote_file, local_file): try: total = 0 while total < fileinfo.st_size: - size, data = file_chan.read(size=fileinfo.st_size - total) + size = min(self._SCP_RECV_BUF_SIZE, fileinfo.st_size - total) + size, data = file_chan.read(size=size) if size == LIBSSH2_ERROR_EAGAIN: self.poll() continue + if size == 0: + raise SCPError( + "Unexpected EOF while receiving %s (%s of %s bytes)" % ( + remote_file, total, fileinfo.st_size), + remote_file, self.host) total += size local_fh.write(data) finally: @@ -729,7 +744,7 @@ def scp_send(self, local_file, remote_file, recurse=False, sftp=None): :raises: :py:class:`OSError` on local OS errors like permission denied """ if os.path.isdir(local_file) and recurse: - sftp = self._make_sftp() if sftp is None else sftp + sftp = self.make_sftp_client() if sftp is None else sftp return self._scp_send_dir(local_file, remote_file, sftp) elif os.path.isdir(local_file) and not recurse: raise ValueError("Recurse must be True if local_file is a " @@ -737,7 +752,7 @@ def scp_send(self, local_file, remote_file, recurse=False, sftp=None): if recurse: destination = self._remote_paths_split(remote_file) if destination is not None: - sftp = self._make_sftp() if sftp is None else sftp + sftp = self.make_sftp_client() if sftp is None else sftp try: self.eagain(sftp.stat, destination) except (SFTPHandleError, SFTPProtocolError): diff --git a/tests/test_native_sftp.py b/tests/test_native_sftp.py new file mode 100644 index 00000000..44a82467 --- /dev/null +++ b/tests/test_native_sftp.py @@ -0,0 +1,191 @@ +# This file is part of parallel-ssh. +# Copyright (C) 2014-2026 Panos Kittenis. +# Copyright (C) 2014-2026 parallel-ssh Contributors. +# +# This library is free software; you can redistribute it and/or +# modify it under the terms of the GNU Lesser General Public +# License as published by the Free Software Foundation, version 2.1. + +import unittest +from types import GeneratorType + +from pssh.clients.native.sftp import SFTPClient + + +class DirectoryHandle(object): + + def __init__(self): + self.closed = False + + def __enter__(self): + return self + + def __exit__(self, *_args): + self.closed = True + + +class SFTP(object): + + def __init__(self): + self.realpath_calls = [] + self.opendir_calls = [] + self.handles = [] + self.calls = [] + + def realpath(self, path): + self.realpath_calls.append(path) + return '/home/tester' if path == '.' else path + + def opendir(self, path): + self.opendir_calls.append(path) + handle = DirectoryHandle() + self.handles.append(handle) + return handle + + def stat(self, path): + self.calls.append(('stat', path)) + return 'stat-result' + + def lstat(self, path): + self.calls.append(('lstat', path)) + return 'lstat-result' + + def rmdir(self, path): + self.calls.append(('rmdir', path)) + return 0 + + def rename(self, source, destination): + self.calls.append(('rename', source, destination)) + return 0 + + def unlink(self, path): + self.calls.append(('unlink', path)) + return 0 + + +class SSHClient(object): + + def __init__(self, sftp): + self.sftp = sftp + + def make_sftp_client(self): + return self.sftp + + def eagain(self, func, *args): + return func(*args) + + def _sftp_openfh(self, func, *args): + return func(*args) + + def _sftp_readdir(self, _handle): + return iter((b'.', b'..', b'file.txt', b'data')) + + def mkdir(self, sftp, path): + self.calls = getattr(self, 'calls', []) + self.calls.append(('mkdir', sftp, path)) + + def sftp_get(self, sftp, remote_file, local_file): + self.calls = getattr(self, 'calls', []) + self.calls.append(('get', sftp, remote_file, local_file)) + + def sftp_put(self, sftp, local_file, remote_file): + self.calls = getattr(self, 'calls', []) + self.calls.append(('put', sftp, local_file, remote_file)) + + def copy_file(self, local_file, remote_file, recurse=False, sftp=None): + self.calls = getattr(self, 'calls', []) + self.calls.append( + ('copy_file', sftp, local_file, remote_file, recurse)) + + def copy_remote_file(self, remote_file, local_file, recurse=False, + sftp=None, encoding='utf-8'): + self.calls = getattr(self, 'calls', []) + self.calls.append( + ('copy_remote_file', sftp, remote_file, local_file, + recurse, encoding)) + + +class NativeSFTPClientTest(unittest.TestCase): + + def setUp(self): + self.sftp = SFTP() + self.ssh_client = SSHClient(self.sftp) + self.client = SFTPClient(self.ssh_client) + + def test_initial_cwd_uses_server_realpath(self): + self.assertEqual(self.client.getcwd(), '/home/tester') + self.assertEqual(self.sftp.realpath_calls, ['.']) + + def test_remote_path_uses_posix_semantics(self): + self.assertEqual( + self.client._remote_path('../shared/./file'), '/home/shared/file') + self.assertEqual( + self.client._remote_path('/var//data/../log'), '/var/log') + + def test_chdir_canonicalizes_and_verifies_directory(self): + cwd = self.client.chdir('data') + + self.assertEqual(cwd, '/home/tester/data') + self.assertEqual(self.client.getcwd(), '/home/tester/data') + self.assertEqual(self.sftp.opendir_calls, ['/home/tester/data']) + self.assertTrue(self.sftp.handles[0].closed) + + def test_invalid_path_type_fails_before_transport(self): + with self.assertRaises(TypeError): + self.client.chdir(None) + + self.assertEqual(self.sftp.opendir_calls, []) + + def test_listdir_filters_names_and_holds_handle_during_iteration(self): + names = self.client.listdir('data') + + self.assertIsInstance(names, GeneratorType) + self.assertEqual(self.sftp.opendir_calls, []) + self.assertEqual(next(names), 'file.txt') + self.assertEqual(self.sftp.opendir_calls, ['/home/tester/data']) + self.assertFalse(self.sftp.handles[0].closed) + self.assertEqual(list(names), ['data']) + self.assertTrue(self.sftp.handles[0].closed) + + def test_metadata_and_mutations_resolve_remote_paths(self): + self.assertEqual(self.client.stat('file'), 'stat-result') + self.assertEqual(self.client.lstat('../link'), 'lstat-result') + self.client.rmdir('empty') + self.client.rename('old', '../new') + self.client.remove('obsolete') + + self.assertEqual( + self.sftp.calls, + [ + ('stat', '/home/tester/file'), + ('lstat', '/home/link'), + ('rmdir', '/home/tester/empty'), + ('rename', '/home/tester/old', '/home/new'), + ('unlink', '/home/tester/obsolete'), + ], + ) + + def test_transfer_helpers_reuse_bound_channel_and_cwd(self): + self.client.mkdir('new/child') + self.client.get('remote.txt', 'local.txt') + self.client.put('local.bin', '../remote.bin') + self.client.copy_file('tree', 'remote-tree', recurse=True) + self.client.copy_remote_file( + 'remote-tree', 'local-tree', recurse=True, encoding='ascii') + + self.assertEqual( + self.ssh_client.calls, + [ + ('mkdir', self.sftp, '/home/tester/new/child'), + ('get', self.sftp, '/home/tester/remote.txt', 'local.txt'), + ('put', self.sftp, 'local.bin', '/home/remote.bin'), + ('copy_file', self.sftp, 'tree', + '/home/tester/remote-tree', True), + ('copy_remote_file', self.sftp, + '/home/tester/remote-tree', 'local-tree', True, 'ascii'), + ], + ) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_native_single.py b/tests/test_native_single.py index a50e241c..435a38f8 100644 --- a/tests/test_native_single.py +++ b/tests/test_native_single.py @@ -11,8 +11,10 @@ import unittest +from unittest.mock import Mock, call, patch from pssh.clients.native.single import SSHClient +from pssh.exceptions import SCPError, SFTPError from pssh.output import HostOutput @@ -41,6 +43,93 @@ def get_exit_status(self): class NativeSingleClientTest(unittest.TestCase): + @patch('pssh.clients.native.single.FileObjectThread') + def test_scp_recv_rejects_unexpected_eof(self, file_object): + client = object.__new__(SSHClient) + client.host = '127.0.0.1' + client.session = Mock() + client.poll = Mock() + channel = Mock() + channel.read.return_value = (0, b'') + fileinfo = Mock(st_size=4) + client.session.scp_recv2.return_value = (channel, fileinfo) + + with self.assertRaises(SCPError): + client._scp_recv('remote', 'local') + + channel.read.assert_called_once_with(size=4) + client.poll.assert_not_called() + file_object.return_value.write.assert_not_called() + file_object.return_value.flush.assert_called_once_with() + file_object.return_value.close.assert_called_once_with() + channel.close.assert_called_once_with() + + @patch('pssh.clients.native.single.FileObjectThread') + def test_scp_recv_limits_read_size_to_buffer(self, file_object): + client = object.__new__(SSHClient) + client._SCP_RECV_BUF_SIZE = 3 + client.session = Mock() + client.poll = Mock() + channel = Mock() + channel.read.side_effect = [(3, b'one'), (1, b'!')] + fileinfo = Mock(st_size=4) + client.session.scp_recv2.return_value = (channel, fileinfo) + + client._scp_recv('remote', 'local') + + self.assertEqual(channel.read.call_args_list, + [call(size=3), call(size=1)]) + client.poll.assert_not_called() + file_object.return_value.write.assert_has_calls([call(b'one'), call(b'!')]) + + def test_make_sftp_client_returns_channel_and_wraps_errors(self): + client = object.__new__(SSHClient) + sftp = object() + client._make_sftp_eagain = lambda: sftp + + self.assertIs(client.make_sftp_client(), sftp) + + error = RuntimeError('sftp init failed') + + def raise_error(): + raise error + + client._make_sftp_eagain = raise_error + + with self.assertRaises(SFTPError) as raised: + client.make_sftp_client() + + self.assertIs(raised.exception.args[0], error) + + def test_transfer_helpers_create_sftp_client(self): + client = Mock(spec=SSHClient) + client.host = 'host' + sftp = Mock() + client.make_sftp_client.return_value = sftp + client._remote_paths_split.return_value = None + client._sftp_openfh.side_effect = SFTPError + client._scp_recv_recursive.return_value = 'received' + client._scp_send_dir.return_value = 'sent' + client.eagain.side_effect = lambda func, *args: func(*args) + + with patch('pssh.clients.native.single.os.path.isdir') as isdir: + isdir.return_value = False + SSHClient.copy_file(client, 'local', 'remote') + SSHClient.copy_remote_file(client, 'remote', 'local') + self.assertEqual( + SSHClient.scp_recv( + client, 'remote', 'local', recurse=True), 'received') + isdir.return_value = True + self.assertEqual( + SSHClient.scp_send( + client, 'local', 'remote', recurse=True), 'sent') + isdir.return_value = False + client._remote_paths_split.return_value = '/remote' + SSHClient.scp_send( + client, 'local', 'remote/file', recurse=True) + + self.assertEqual(client.make_sftp_client.call_count, 5) + def test_wait_finished_waits_for_close_before_exit_status(self): client = object.__new__(SSHClient) client.eagain = lambda func: func()